From fc8f8a451d6b0d904ab7112b4eff035bdc9c66f0 Mon Sep 17 00:00:00 2001 From: LjhAUMEM Date: Wed, 30 Sep 2026 01:16:59 +0800 Subject: [PATCH] XDNS finalmask: Refactor and new parameters (#6718) https://github.com/XTLS/Xray-core/pull/6718#issuecomment-5894987590 Fixes https://github.com/XTLS/Xray-core/issues/6692 --- infra/conf/transport_finalmask.go | 100 +- transport/internet/finalmask/udphop/conn.go | 7 - transport/internet/finalmask/xdns/client.go | 664 ++++++------ transport/internet/finalmask/xdns/config.go | 4 +- .../internet/finalmask/xdns/config.pb.go | 230 ++++- .../internet/finalmask/xdns/config.proto | 23 +- transport/internet/finalmask/xdns/dns.go | 581 ----------- transport/internet/finalmask/xdns/dns_test.go | 953 ------------------ transport/internet/finalmask/xdns/domain.go | 215 ++++ transport/internet/finalmask/xdns/frag.go | 171 ++++ .../finalmask/xdns/record_transport.go | 226 ----- transport/internet/finalmask/xdns/resolver.go | 31 + .../internet/finalmask/xdns/resolver_tcp.go | 143 +++ .../internet/finalmask/xdns/resolver_udp.go | 130 +++ transport/internet/finalmask/xdns/resp.go | 392 +++++++ transport/internet/finalmask/xdns/server.go | 703 ++++++------- transport/internet/finalmask/xdns/spec.go | 80 -- .../internet/finalmask/xdns/xdns_test.go | 208 ++++ transport/internet/finalmask/xicmp/client.go | 7 - transport/internet/finalmask/xicmp/server.go | 7 - .../internet/finalmask/xicmp/server_oob.go | 7 - 21 files changed, 2235 insertions(+), 2647 deletions(-) delete mode 100644 transport/internet/finalmask/xdns/dns.go delete mode 100644 transport/internet/finalmask/xdns/dns_test.go create mode 100644 transport/internet/finalmask/xdns/domain.go create mode 100644 transport/internet/finalmask/xdns/frag.go delete mode 100644 transport/internet/finalmask/xdns/record_transport.go create mode 100644 transport/internet/finalmask/xdns/resolver.go create mode 100644 transport/internet/finalmask/xdns/resolver_tcp.go create mode 100644 transport/internet/finalmask/xdns/resolver_udp.go create mode 100644 transport/internet/finalmask/xdns/resp.go delete mode 100644 transport/internet/finalmask/xdns/spec.go create mode 100644 transport/internet/finalmask/xdns/xdns_test.go diff --git a/infra/conf/transport_finalmask.go b/infra/conf/transport_finalmask.go index 62abf3893..784af98a0 100644 --- a/infra/conf/transport_finalmask.go +++ b/infra/conf/transport_finalmask.go @@ -1,6 +1,7 @@ package conf import ( + "context" "crypto/x509" "encoding/base64" "encoding/hex" @@ -14,6 +15,7 @@ import ( googleuuid "github.com/google/uuid" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/serial" "github.com/xtls/xray-core/transport/internet/finalmask/fragment" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom" "github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm" @@ -81,7 +83,7 @@ var ( "noise": func() interface{} { return new(NoiseMask) }, "salamander": func() interface{} { return new(Salamander) }, "sudoku": func() interface{} { return new(Sudoku) }, - "xdns": func() interface{} { return new(Xdns) }, + "xdns": func() interface{} { return new(XDNS) }, "xicmp": func() interface{} { return new(Xicmp) }, "realm": func() interface{} { return new(Realm) }, "udphop": func() interface{} { return new(UDPHop) }, @@ -694,32 +696,88 @@ func (c *Sudoku) Build() (proto.Message, error) { }, nil } -type Xdns struct { - Domain json.RawMessage `json:"domain"` - - Domains []string `json:"domains"` - Resolvers []string `json:"resolvers"` +type XDNSDomain struct { + Name string `json:"name"` + LenLimit int32 `json:"lenLimit"` + LabelLimit int32 `json:"labelLimit"` + Types []int32 `json:"types"` + Edns0 int32 `json:"edns0"` } -func (c *Xdns) Build() (proto.Message, error) { - if c.Domain != nil { - return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)") - } +type XDNSResolverTCP struct { + Addr string `json:"addr"` +} - if len(c.Domains) == 0 && len(c.Resolvers) == 0 { - return nil, errors.New("empty domains & empty resolvers") - } +func (c *XDNSResolverTCP) Build() (proto.Message, error) { + return &xdns.TCPResolverProto{Addr: c.Addr}, nil +} - for _, r := range c.Resolvers { - if !strings.Contains(r, "+udp://") { - return nil, errors.New("invalid resolver ", r) +type XDNSResolverUDP struct { + Addr string `json:"addr"` +} + +func (c *XDNSResolverUDP) Build() (proto.Message, error) { + return &xdns.UDPResolverProto{Addr: c.Addr}, nil +} + +var xdnsLoader = NewJSONConfigLoader(ConfigCreatorCache{ + "tcp": func() interface{} { return new(XDNSResolverTCP) }, + "udp": func() interface{} { return new(XDNSResolverUDP) }, +}, "type", "settings") + +type XDNSResolver struct { + Type string `json:"type"` + Settings json.RawMessage `json:"settings"` +} + +type XDNS struct { + Domains []XDNSDomain `json:"domains"` + Resolvers []XDNSResolver `json:"resolvers"` + ExtraPoll int32 `json:"extraPoll"` +} + +func (c *XDNS) Build() (proto.Message, error) { + var domains []*xdns.DomainProto + var resolvers []*serial.TypedMessage + for i := range c.Domains { + if c.Domains[i].LenLimit == 0 { + c.Domains[i].LenLimit = 255 } + if c.Domains[i].LabelLimit == 0 { + c.Domains[i].LabelLimit = 63 + } + types := make([]uint16, 0, len(c.Domains[i].Types)) + for j := range c.Domains[i].Types { + types = append(types, uint16(c.Domains[i].Types[j])) + } + domain, err := xdns.NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0)) + if err != nil { + return nil, err + } + errors.LogInfo(context.Background(), domain.Show()) + domains = append(domains, &xdns.DomainProto{ + Name: c.Domains[i].Name, + LenLimit: c.Domains[i].LenLimit, + LabelLimit: c.Domains[i].LabelLimit, + Types: c.Domains[i].Types, + Edns0: c.Domains[i].Edns0, + }) } - - return &xdns.Config{ - Domains: c.Domains, - Resolvers: c.Resolvers, - }, nil + for i := range c.Resolvers { + config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type) + if err != nil { + return nil, err + } + pm, err := config.(interface{ Build() (proto.Message, error) }).Build() + if err != nil { + return nil, err + } + resolvers = append(resolvers, serial.ToTypedMessage(pm)) + } + if c.ExtraPoll < 0 || c.ExtraPoll > 3 { + return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3") + } + return &xdns.Config{Domains: domains, Resolvers: resolvers, ExtraPoll: c.ExtraPoll}, nil } type XMC struct { diff --git a/transport/internet/finalmask/udphop/conn.go b/transport/internet/finalmask/udphop/conn.go index a71a40439..8c6578213 100644 --- a/transport/internet/finalmask/udphop/conn.go +++ b/transport/internet/finalmask/udphop/conn.go @@ -223,13 +223,6 @@ func (c *udpHopConn) Close() error { } _ = c.cur.Close() c.wg.Wait() - select { - case packet := <-c.readCh: - if packet.p != nil { - pool.Put(packet.p[:cap(packet.p)]) - } - default: - } close(c.readCh) return nil } diff --git a/transport/internet/finalmask/xdns/client.go b/transport/internet/finalmask/xdns/client.go index 557347f21..b976f154f 100644 --- a/transport/internet/finalmask/xdns/client.go +++ b/transport/internet/finalmask/xdns/client.go @@ -1,417 +1,441 @@ package xdns import ( - "bytes" "context" "crypto/rand" - "encoding/base32" - "encoding/binary" - go_errors "errors" "io" - "net" - "strconv" + mrand "math/rand" "sync" "sync/atomic" "time" "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/transport/internet/finalmask" + "golang.org/x/net/dns/dnsmessage" ) const ( - numPadding = 3 - numPaddingForPoll = 8 initPollDelay = 500 * time.Millisecond maxPollDelay = 10 * time.Second pollDelayMultiplier = 2.0 pollLimit = 16 ) -var base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding) +var pool4K = sync.Pool{ + New: func() any { + return make([]byte, 4096) + }, +} type packet struct { p []byte addr net.Addr } -type xdnsConnClient struct { - net.PacketConn +type xdnsClient struct { + dialer *finalmask.Dialer - resolverAddrs []*net.UDPAddr - resolverTypes []uint16 - resolverIdx uint32 - resolverSend map[string]*atomic.Uint32 + clientID ClientID + fragID atomic.Uint32 + domains []*Domain + extraPoll int32 - clientID []byte - domains []Name + resolvers []Resolver + resolverSends []atomic.Uint32 + resolverIndex atomic.Uint32 - pollChan chan struct{} - readQueue chan *packet - writeQueue chan *packet - - closed bool - mutex sync.Mutex + readCh chan packet + sendCh chan []byte + poolCh chan struct{} + closeCh chan struct{} + wg sync.WaitGroup + mu sync.Mutex } -func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) { + if len(c.Domains) == 0 { + return nil, errors.New("empty domains") + } if len(c.Resolvers) == 0 { return nil, errors.New("empty resolvers") } - - var domains []Name - var servers []string - var resolverTypes []uint16 - for _, rs := range c.Resolvers { - domain, server, resolverType, err := parseResolver(rs) - if err != nil { - return nil, errors.New("invalid resolvers").Base(err) - } - domains = append(domains, domain) - servers = append(servers, server) - resolverTypes = append(resolverTypes, resolverType) + if c.ExtraPoll < 0 || c.ExtraPoll > 3 { + return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3") } - - var resolverAddrs []*net.UDPAddr - resolverSend := make(map[string]*atomic.Uint32) - for _, rs := range servers { - h, p, err := net.SplitHostPort(rs) + domains := make([]*Domain, 0, len(c.Domains)) + for i := range c.Domains { + types := make([]uint16, 0, len(c.Domains[i].Types)) + for j := range c.Domains[i].Types { + types = append(types, uint16(c.Domains[i].Types[j])) + } + domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0)) if err != nil { return nil, err } - ip := net.ParseIP(h) - if ip == nil { - return nil, errors.New("invalid ip address") - } - port, err := strconv.Atoi(p) + domains = append(domains, domain) + } + resolvers := make([]Resolver, 0, len(c.Resolvers)) + for i := range c.Resolvers { + resolver, err := NewResolver(c.Resolvers[i], dialer) if err != nil { - return nil, errors.New("invalid port").Base(err) + return nil, err } - addr := &net.UDPAddr{IP: ip, Port: port} - resolverAddrs = append(resolverAddrs, addr) - resolverSend[addr.String()] = &atomic.Uint32{} + resolvers = append(resolvers, resolver) } + client := &xdnsClient{ + dialer: dialer, - conn := &xdnsConnClient{ - PacketConn: raw, + clientID: NewClientID(), + domains: domains, + extraPoll: c.ExtraPoll, - resolverAddrs: resolverAddrs, - resolverTypes: resolverTypes, - resolverIdx: 0, - resolverSend: resolverSend, + resolvers: resolvers, + resolverSends: make([]atomic.Uint32, len(c.Resolvers)), - clientID: make([]byte, 8), - domains: domains, - - pollChan: make(chan struct{}, pollLimit), - readQueue: make(chan *packet, 256), - writeQueue: make(chan *packet, 256), + readCh: make(chan packet), + sendCh: make(chan []byte, 16), + poolCh: make(chan struct{}, pollLimit), + closeCh: make(chan struct{}), } - - common.Must2(rand.Read(conn.clientID)) - - go conn.recvLoop() - go conn.sendLoop() - - return conn, nil + go client.run() + return client, nil } -func (c *xdnsConnClient) recvLoop() { - var buf [finalmask.UDPSize]byte +func (c *xdnsClient) closed() bool { + select { + case <-c.closeCh: + return true + default: + return false + } +} - for { - if c.closed { +func (c *xdnsClient) read(buf []byte, addr net.Addr) bool { + msg := dnsmessage.Message{} + if err := msg.Unpack(buf); err != nil { + return false + } + if !msg.Header.Response || msg.Header.Truncated || msg.Header.RCode != dnsmessage.RCodeSuccess || len(msg.Questions) != 1 { + return false + } + + var domain *Domain + for i := range c.domains { + if c.domains[i].IsDomain(msg.Questions[0].Name) { + domain = c.domains[i] break } + } + if domain == nil || !domain.HasType(uint16(msg.Questions[0].Type)) { + return false + } - n, addr, err := c.PacketConn.ReadFrom(buf[:]) + edns0 := uint16(0) + for i := range msg.Additionals { + if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT { + edns0 = uint16(msg.Additionals[i].Header.Class) + break + } + } + errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type) + + resp := NewResp(msg, domain, 0) + + p := pool4K.Get().([]byte) + n := resp.Decode(p) + p = p[:n] + + b := p + var bs [][]byte + for len(b) > 1 { + last := b[0]&0xC0 == 0xC0 + length := int(b[0]&0x3F)<<8 | int(b[1]) + b = b[2:] + if length > len(b) { + bs = nil + break + } + packet := make([]byte, length) + copy(packet, b) + bs = append(bs, packet) + if last { + break + } + b = b[length:] + if len(b) < 2 { + bs = nil + } + } + pool4K.Put(p[:cap(p)]) + + for i := range bs { + select { + case <-c.closeCh: + return true + case c.readCh <- packet{p: bs[i], addr: addr}: + } + } + return len(bs) > 0 +} + +func (c *xdnsClient) run() { + for i := range len(c.resolvers) { + c.wg.Add(1) + go c.recv(i) + } + + c.wg.Add(1) + go c.send() + + c.wg.Wait() + close(c.readCh) + close(c.sendCh) + close(c.poolCh) +} + +func (c *xdnsClient) recv(i int) { + defer c.wg.Done() + + var buf [4096]byte + for { + n, err := c.resolvers[i].Read(buf[:]) if err != nil { - if go_errors.Is(err, net.ErrClosed) { - break + if c.closed() { + return } - continue + errors.LogErrorInner(context.Background(), err, "recv err ", i) + return } - - if addr == nil { - continue - } - - send := c.resolverSend[addr.String()] - if send == nil { - continue - } - - resp, err := MessageFromWireFormat(buf[:n]) - if err != nil { - errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err) - continue - } - - payload := dnsResponsePayload(&resp, c.domains) - - r := bytes.NewReader(payload) - anyPacket := false - for { - p, err := nextPacket(r) - if err != nil { - break - } - anyPacket = true - - buf := make([]byte, len(p)) - copy(buf, p) + if c.read(buf[:n], c.resolvers[i].Addr()) { + c.resolverSends[i].Store(0) select { - case c.readQueue <- &packet{ - p: buf, - addr: addr, - }: - default: - errors.LogDebug(context.Background(), addr, " mask read err queue full") - } - } - - if anyPacket { - send.Store(0) - select { - case c.pollChan <- struct{}{}: + case c.poolCh <- struct{}{}: default: } } } - - errors.LogDebug(context.Background(), "xdns closed") - - close(c.pollChan) - close(c.readQueue) - - c.mutex.Lock() - defer c.mutex.Unlock() - - c.closed = true - close(c.writeQueue) } -func (c *xdnsConnClient) sendLoop() { - pollDelay := initPollDelay - pollTimer := time.NewTimer(pollDelay) - for { - var p *packet - pollTimerExpired := false +func (c *xdnsClient) send() { + defer c.wg.Done() - select { - case p = <-c.writeQueue: - default: - select { - case p = <-c.writeQueue: - case <-c.pollChan: - case <-pollTimer.C: - pollTimerExpired = true + var buf [512]byte + var data [255]byte + + sendMsg := func(p []byte, domain *Domain, qtype uint16) { + msg := dnsmessage.Message{ + Header: dnsmessage.Header{ + RecursionDesired: true, + }, + Questions: []dnsmessage.Question{ + { + Name: domain.Encode(p), + Type: dnsmessage.Type(qtype), + Class: dnsmessage.ClassINET, + }, + }, + } + if domain.edns0 > 0 { + msg.Additionals = []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeOPT, + Class: dnsmessage.Class(domain.edns0), + TTL: 0, + }, + Body: &dnsmessage.OPTResource{}, + }, } } + pack := common.Must2(msg.AppendPack(buf[:0])) + common.Must2(rand.Read(pack[:2])) - if p != nil { - select { - case <-c.pollChan: - default: + index := c.resolverIndex.Load() + cur := c.resolverSends[index].Add(1) + i := index + for { + i++ + if i == uint32(len(c.resolvers)) { + i = 0 } - } else { - encoded, _ := encode(nil, c.clientID, c.domains[c.resolverIdx], c.resolverTypes[c.resolverIdx]) - p = &packet{ - p: encoded, + if i == index { + break + } + if cur > c.resolverSends[i].Load() { + break } } + c.resolverIndex.Store(i) + c.resolvers[index].Send(pack) + } - if pollTimerExpired { - pollDelay = time.Duration(float64(pollDelay) * pollDelayMultiplier) - if pollDelay > maxPollDelay { - pollDelay = maxPollDelay - } - } else { - if !pollTimer.Stop() { - <-pollTimer.C - } - pollDelay = initPollDelay - } - pollTimer.Reset(pollDelay) + send := func(p []byte) { + domain := c.domains[mrand.Intn(len(c.domains))] + qtype := domain.types[mrand.Intn(len(domain.types))] - if c.closed { + if len(p) == 0 { + copy(data[:], c.clientID[:]) + data[0] |= TypeMap[qtype] + data[8] = 8 + common.Must2(rand.Read(data[9:17])) + sendMsg(data[:17], domain, qtype) return } - cur := c.resolverIdx - curSend := c.resolverSend[c.resolverAddrs[cur].String()].Add(1) - _, _ = c.PacketConn.WriteTo(p.p, c.resolverAddrs[cur]) - for { - c.resolverIdx += 1 - c.resolverIdx %= uint32(len(c.resolverAddrs)) - if c.resolverIdx == cur { - break + if len(p) <= domain.cap-12 { + copy(data[:], c.clientID[:]) + data[0] |= TypeMap[qtype] + data[8] = 3 + common.Must2(rand.Read(data[9:12])) + copy(data[12:], p) + sendMsg(data[:12+len(p)], domain, qtype) + return + } + + if len(p) <= 255*(domain.cap-15) { + copy(data[:], c.clientID[:]) + data[0] |= TypeMap[qtype] + data[8] = 3 | 0xC0 + common.Must2(rand.Read(data[9:12])) + + fragID := byte(c.fragID.Add(1)) + fragN := len(p) / (domain.cap - 15) + if len(p)%(domain.cap-15) > 0 { + fragN++ } - if c.resolverSend[c.resolverAddrs[c.resolverIdx].String()].Load() < curSend { - break + + for i := range fragN { + data[12] = fragID + data[13] = byte(i) + data[14] = byte(fragN) + size := min(len(p), domain.cap-15) + copy(data[15:], p[:size]) + sendMsg(data[:15+size], domain, qtype) + p = p[size:] + } + return + } + + errors.LogError(context.Background(), "err size ", len(p)) + } + + ticker := time.NewTicker(initPollDelay) + defer ticker.Stop() + delay := initPollDelay + p := []byte(nil) + timeout := false + for { + select { + case <-c.closeCh: + return + default: + select { + case <-c.closeCh: + return + case p = <-c.sendCh: + case <-c.poolCh: + case <-ticker.C: + timeout = true } } + + if len(p) > 0 { + select { + case <-c.poolCh: + default: + } + } + + send(p) + for range c.extraPoll { + send(nil) + } + + if timeout { + delay *= pollDelayMultiplier + if delay > maxPollDelay { + delay = maxPollDelay + } + timeout = false + } else { + delay = initPollDelay + } + ticker.Reset(delay) } } -func (c *xdnsConnClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) { - packet, ok := <-c.readQueue - if !ok { - return 0, nil, net.ErrClosed +func (c *xdnsClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + packet, ok := <-c.readCh + if ok { + return copy(p, packet.p), packet.addr, nil } - if len(p) < len(packet.p) { - errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p)) - return 0, packet.addr, nil - } - copy(p, packet.p) - return len(packet.p), packet.addr, nil + return 0, nil, io.ErrClosedPipe } -func (c *xdnsConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) { - c.mutex.Lock() - defer c.mutex.Unlock() - - if c.closed { +func (c *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed() { return 0, io.ErrClosedPipe } - - idx := c.resolverIdx % uint32(len(c.resolverAddrs)) - encoded, err := encode(p, c.clientID, c.domains[idx], c.resolverTypes[idx]) - if err != nil { - errors.LogDebug(context.Background(), addr, " xdns wireformat err ", err, " ", len(p)) - return 0, nil + if len(p) == 0 || len(p) > 4096 { + errors.LogError(context.Background(), "err size ", len(p)) + return 0, errors.New("err size") } - + b := make([]byte, len(p)) + copy(b, p) select { - case c.writeQueue <- &packet{ - p: encoded, - addr: addr, - }: - return len(p), nil + case c.sendCh <- b: default: - errors.LogDebug(context.Background(), addr, " mask write err queue full") - return 0, nil } + return len(p), nil } -func (c *xdnsConnClient) Close() error { - c.closed = true - return c.PacketConn.Close() -} - -func encode(p []byte, clientID []byte, domain Name, qtype uint16) ([]byte, error) { - var decoded []byte - { - if len(p) >= 224 { - return nil, errors.New("too long") - } - var buf bytes.Buffer - buf.Write(clientID[:]) - n := numPadding - if len(p) == 0 { - n = numPaddingForPoll - } - buf.WriteByte(byte(224 + n)) - _, _ = io.CopyN(&buf, rand.Reader, int64(n)) - if len(p) > 0 { - buf.WriteByte(byte(len(p))) - buf.Write(p) - } - decoded = buf.Bytes() - } - - encoded := make([]byte, base32Encoding.EncodedLen(len(decoded))) - base32Encoding.Encode(encoded, decoded) - encoded = bytes.ToLower(encoded) - labels := chunks(encoded, 63) - labels = append(labels, domain...) - name, err := NewName(labels) - if err != nil { - return nil, err - } - - var id uint16 - _ = binary.Read(rand.Reader, binary.BigEndian, &id) - query := &Message{ - ID: id, - Flags: 0x0100, - Question: []Question{ - { - Name: name, - Type: qtype, - Class: ClassIN, - }, - }, - Additional: []RR{ - { - Name: Name{}, - Type: RRTypeOPT, - Class: 4096, - TTL: 0, - Data: []byte{}, - }, - }, - } - - buf, err := query.WireFormat() - if err != nil { - return nil, err - } - - return buf, nil -} - -func chunks(p []byte, n int) [][]byte { - var result [][]byte - for len(p) > 0 { - sz := len(p) - if sz > n { - sz = n - } - result = append(result, p[:sz]) - p = p[sz:] - } - return result -} - -func nextPacket(r *bytes.Reader) ([]byte, error) { - var n uint16 - err := binary.Read(r, binary.BigEndian, &n) - if err != nil { - return nil, err - } - p := make([]byte, n) - _, err = io.ReadFull(r, p) - if err == io.EOF { - err = io.ErrUnexpectedEOF - } - return p, err -} - -func dnsResponsePayload(resp *Message, domains []Name) []byte { - if resp.Flags&0x8000 != 0x8000 { +func (c *xdnsClient) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed() { return nil } - if resp.Flags&0x000f != RcodeNoError { - return nil + close(c.closeCh) + for i := range c.resolvers { + c.resolvers[i].Close() } + return nil +} - if len(resp.Answer) == 0 { - return nil - } - - for _, answer := range resp.Answer { - var ok bool - for _, domain := range domains { - _, ok = answer.Name.TrimSuffix(domain) - if ok { - break - } - } - if !ok { - return nil - } - } - - return decodeResponsePayload(resp.Answer) +func (c *xdnsClient) LocalAddr() net.Addr { return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} } + +func (c *xdnsClient) SetDeadline(t time.Time) error { return errors.New("not support") } + +func (c *xdnsClient) SetReadDeadline(t time.Time) error { return errors.New("not support") } + +func (c *xdnsClient) SetWriteDeadline(t time.Time) error { return errors.New("not support") } + +type ClientID [8]byte + +func NewClientID() ClientID { + var id ClientID + common.Must2(rand.Read(id[:])) + id[0] &= 0xFC + return id +} + +func ClientIDFromRaw(id [8]byte) ClientID { + id[0] &= 0xFC + return id +} + +func ClientIDFromAddr(addr *net.UDPAddr) ClientID { + return ClientID(addr.IP[8:]) +} + +func (id ClientID) Addr() *net.UDPAddr { + var ip [16]byte + ip[0] = 0xFD + copy(ip[8:], id[:]) + return &net.UDPAddr{IP: ip[:]} } diff --git a/transport/internet/finalmask/xdns/config.go b/transport/internet/finalmask/xdns/config.go index 7bae597ab..d982c0aa1 100644 --- a/transport/internet/finalmask/xdns/config.go +++ b/transport/internet/finalmask/xdns/config.go @@ -6,9 +6,9 @@ import ( ) func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { - return NewConnClient(c, conn) + return NewClient(c, dialer) } func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { - return NewConnServer(c, conn) + return NewServer(c, conn) } diff --git a/transport/internet/finalmask/xdns/config.pb.go b/transport/internet/finalmask/xdns/config.pb.go index e1f06aa93..7db693f98 100644 --- a/transport/internet/finalmask/xdns/config.pb.go +++ b/transport/internet/finalmask/xdns/config.pb.go @@ -7,6 +7,7 @@ package xdns import ( + serial "github.com/xtls/xray-core/common/serial" protoreflect "google.golang.org/protobuf/reflect/protoreflect" protoimpl "google.golang.org/protobuf/runtime/protoimpl" reflect "reflect" @@ -21,17 +22,94 @@ const ( _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) ) +type DomainProto struct { + state protoimpl.MessageState `protogen:"open.v1"` + Name string `protobuf:"bytes,1,opt,name=name,proto3" json:"name,omitempty"` + LenLimit int32 `protobuf:"varint,2,opt,name=len_limit,json=lenLimit,proto3" json:"len_limit,omitempty"` + LabelLimit int32 `protobuf:"varint,3,opt,name=label_limit,json=labelLimit,proto3" json:"label_limit,omitempty"` + Types []int32 `protobuf:"varint,4,rep,packed,name=types,proto3" json:"types,omitempty"` + Edns0 int32 `protobuf:"varint,5,opt,name=edns0,proto3" json:"edns0,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *DomainProto) Reset() { + *x = DomainProto{} + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *DomainProto) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*DomainProto) ProtoMessage() {} + +func (x *DomainProto) ProtoReflect() protoreflect.Message { + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use DomainProto.ProtoReflect.Descriptor instead. +func (*DomainProto) Descriptor() ([]byte, []int) { + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0} +} + +func (x *DomainProto) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *DomainProto) GetLenLimit() int32 { + if x != nil { + return x.LenLimit + } + return 0 +} + +func (x *DomainProto) GetLabelLimit() int32 { + if x != nil { + return x.LabelLimit + } + return 0 +} + +func (x *DomainProto) GetTypes() []int32 { + if x != nil { + return x.Types + } + return nil +} + +func (x *DomainProto) GetEdns0() int32 { + if x != nil { + return x.Edns0 + } + return 0 +} + type Config struct { state protoimpl.MessageState `protogen:"open.v1"` - Domains []string `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"` - Resolvers []string `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"` + Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"` + Resolvers []*serial.TypedMessage `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"` + ExtraPoll int32 `protobuf:"varint,3,opt,name=extra_poll,json=extraPoll,proto3" json:"extra_poll,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } func (x *Config) Reset() { *x = Config{} - mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -43,7 +121,7 @@ func (x *Config) String() string { func (*Config) ProtoMessage() {} func (x *Config) ProtoReflect() protoreflect.Message { - mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -56,31 +134,139 @@ func (x *Config) ProtoReflect() protoreflect.Message { // Deprecated: Use Config.ProtoReflect.Descriptor instead. func (*Config) Descriptor() ([]byte, []int) { - return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0} + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1} } -func (x *Config) GetDomains() []string { +func (x *Config) GetDomains() []*DomainProto { if x != nil { return x.Domains } return nil } -func (x *Config) GetResolvers() []string { +func (x *Config) GetResolvers() []*serial.TypedMessage { if x != nil { return x.Resolvers } return nil } +func (x *Config) GetExtraPoll() int32 { + if x != nil { + return x.ExtraPoll + } + return 0 +} + +type TCPResolverProto struct { + state protoimpl.MessageState `protogen:"open.v1"` + Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *TCPResolverProto) Reset() { + *x = TCPResolverProto{} + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *TCPResolverProto) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*TCPResolverProto) ProtoMessage() {} + +func (x *TCPResolverProto) ProtoReflect() protoreflect.Message { + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use TCPResolverProto.ProtoReflect.Descriptor instead. +func (*TCPResolverProto) Descriptor() ([]byte, []int) { + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2} +} + +func (x *TCPResolverProto) GetAddr() string { + if x != nil { + return x.Addr + } + return "" +} + +type UDPResolverProto struct { + state protoimpl.MessageState `protogen:"open.v1"` + Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *UDPResolverProto) Reset() { + *x = UDPResolverProto{} + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *UDPResolverProto) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*UDPResolverProto) ProtoMessage() {} + +func (x *UDPResolverProto) ProtoReflect() protoreflect.Message { + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3] + if x != nil { + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + if ms.LoadMessageInfo() == nil { + ms.StoreMessageInfo(mi) + } + return ms + } + return mi.MessageOf(x) +} + +// Deprecated: Use UDPResolverProto.ProtoReflect.Descriptor instead. +func (*UDPResolverProto) Descriptor() ([]byte, []int) { + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{3} +} + +func (x *UDPResolverProto) GetAddr() string { + if x != nil { + return x.Addr + } + return "" +} + var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" + "\n" + - ".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\"@\n" + - "\x06Config\x12\x18\n" + - "\adomains\x18\x01 \x03(\tR\adomains\x12\x1c\n" + - "\tresolvers\x18\x02 \x03(\tR\tresolversB\x94\x01\n" + + ".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\x1a!common/serial/typed_message.proto\"\x8b\x01\n" + + "\vDomainProto\x12\x12\n" + + "\x04name\x18\x01 \x01(\tR\x04name\x12\x1b\n" + + "\tlen_limit\x18\x02 \x01(\x05R\blenLimit\x12\x1f\n" + + "\vlabel_limit\x18\x03 \x01(\x05R\n" + + "labelLimit\x12\x14\n" + + "\x05types\x18\x04 \x03(\x05R\x05types\x12\x14\n" + + "\x05edns0\x18\x05 \x01(\x05R\x05edns0\"\xb6\x01\n" + + "\x06Config\x12M\n" + + "\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12>\n" + + "\tresolvers\x18\x02 \x03(\v2 .xray.common.serial.TypedMessageR\tresolvers\x12\x1d\n" + + "\n" + + "extra_poll\x18\x03 \x01(\x05R\textraPoll\"&\n" + + "\x10TCPResolverProto\x12\x12\n" + + "\x04addr\x18\x01 \x01(\tR\x04addr\"&\n" + + "\x10UDPResolverProto\x12\x12\n" + + "\x04addr\x18\x01 \x01(\tR\x04addrB\x94\x01\n" + "*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3" var ( @@ -95,16 +281,22 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte { return file_transport_internet_finalmask_xdns_config_proto_rawDescData } -var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) +var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4) var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{ - (*Config)(nil), // 0: xray.transport.internet.finalmask.xdns.Config + (*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto + (*Config)(nil), // 1: xray.transport.internet.finalmask.xdns.Config + (*TCPResolverProto)(nil), // 2: xray.transport.internet.finalmask.xdns.TCPResolverProto + (*UDPResolverProto)(nil), // 3: xray.transport.internet.finalmask.xdns.UDPResolverProto + (*serial.TypedMessage)(nil), // 4: xray.common.serial.TypedMessage } var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{ - 0, // [0:0] is the sub-list for method output_type - 0, // [0:0] is the sub-list for method input_type - 0, // [0:0] is the sub-list for extension type_name - 0, // [0:0] is the sub-list for extension extendee - 0, // [0:0] is the sub-list for field type_name + 0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto + 4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage + 2, // [2:2] is the sub-list for method output_type + 2, // [2:2] is the sub-list for method input_type + 2, // [2:2] is the sub-list for extension type_name + 2, // [2:2] is the sub-list for extension extendee + 0, // [0:2] is the sub-list for field type_name } func init() { file_transport_internet_finalmask_xdns_config_proto_init() } @@ -118,7 +310,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)), NumEnums: 0, - NumMessages: 1, + NumMessages: 4, NumExtensions: 0, NumServices: 0, }, diff --git a/transport/internet/finalmask/xdns/config.proto b/transport/internet/finalmask/xdns/config.proto index b859b17ae..1e464cb2d 100644 --- a/transport/internet/finalmask/xdns/config.proto +++ b/transport/internet/finalmask/xdns/config.proto @@ -6,7 +6,26 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/xdns option java_package = "com.xray.transport.internet.finalmask.xdns"; option java_multiple_files = true; +import "common/serial/typed_message.proto"; + +message DomainProto { + string name = 1; + int32 len_limit = 2; + int32 label_limit = 3; + repeated int32 types = 4; + int32 edns0 = 5; +} + message Config { - repeated string domains = 1; - repeated string resolvers = 2; + repeated DomainProto domains = 1; + repeated xray.common.serial.TypedMessage resolvers = 2; + int32 extra_poll = 3; +} + +message TCPResolverProto { + string addr = 1; +} + +message UDPResolverProto { + string addr = 1; } \ No newline at end of file diff --git a/transport/internet/finalmask/xdns/dns.go b/transport/internet/finalmask/xdns/dns.go deleted file mode 100644 index 774903857..000000000 --- a/transport/internet/finalmask/xdns/dns.go +++ /dev/null @@ -1,581 +0,0 @@ -// Package dns deals with encoding and decoding DNS wire format. -package xdns - -import ( - "bytes" - "encoding/binary" - "errors" - "fmt" - "io" - "strings" -) - -// The maximum number of DNS name compression pointers we are willing to follow. -// Without something like this, infinite loops are possible. -const compressionPointerLimit = 10 - -var ( - // ErrZeroLengthLabel is the error returned for names that contain a - // zero-length label, like "example..com". - ErrZeroLengthLabel = errors.New("name contains a zero-length label") - - // ErrLabelTooLong is the error returned for labels that are longer than - // 63 octets. - ErrLabelTooLong = errors.New("name contains a label longer than 63 octets") - - // ErrNameTooLong is the error returned for names whose encoded - // representation is longer than 255 octets. - ErrNameTooLong = errors.New("name is longer than 255 octets") - - // ErrReservedLabelType is the error returned when reading a label type - // prefix whose two most significant bits are not 00 or 11. - ErrReservedLabelType = errors.New("reserved label type") - - // ErrTooManyPointers is the error returned when reading a compressed - // name that has too many compression pointers. - ErrTooManyPointers = errors.New("too many compression pointers") - - // ErrTrailingBytes is the error returned when bytes remain in the parse - // buffer after parsing a message. - ErrTrailingBytes = errors.New("trailing bytes after message") - - // ErrIntegerOverflow is the error returned when trying to encode an - // integer greater than 65535 into a 16-bit field. - ErrIntegerOverflow = errors.New("integer overflow") -) - -const ( - // https://tools.ietf.org/html/rfc1035#section-3.2.2 - RRTypeA = 1 - // https://tools.ietf.org/html/rfc1035#section-3.2.2 - RRTypeCNAME = 5 - // https://tools.ietf.org/html/rfc1035#section-3.2.2 - RRTypeTXT = 16 - // https://tools.ietf.org/html/rfc3596#section-2.1 - RRTypeAAAA = 28 - // https://tools.ietf.org/html/rfc6891#section-6.1.1 - RRTypeOPT = 41 - - // https://tools.ietf.org/html/rfc1035#section-3.2.4 - ClassIN = 1 - - // https://tools.ietf.org/html/rfc1035#section-4.1.1 - RcodeNoError = 0 // a.k.a. NOERROR - RcodeFormatError = 1 // a.k.a. FORMERR - RcodeNameError = 3 // a.k.a. NXDOMAIN - RcodeNotImplemented = 4 // a.k.a. NOTIMPL - // https://tools.ietf.org/html/rfc6891#section-9 - ExtendedRcodeBadVers = 16 // a.k.a. BADVERS -) - -// Name represents a domain name, a sequence of labels each of which is 63 -// octets or less in length. -// -// https://tools.ietf.org/html/rfc1035#section-3.1 -type Name [][]byte - -// NewName returns a Name from a slice of labels, after checking the labels for -// validity. Does not include a zero-length label at the end of the slice. -func NewName(labels [][]byte) (Name, error) { - name := Name(labels) - // https://tools.ietf.org/html/rfc1035#section-2.3.4 - // Various objects and parameters in the DNS have size limits. - // labels 63 octets or less - // names 255 octets or less - for _, label := range labels { - if len(label) == 0 { - return nil, ErrZeroLengthLabel - } - if len(label) > 63 { - return nil, ErrLabelTooLong - } - } - // Check the total length. - builder := newMessageBuilder() - builder.WriteName(name) - if len(builder.Bytes()) > 255 { - return nil, ErrNameTooLong - } - return name, nil -} - -// ParseName returns a new Name from a string of labels separated by dots, after -// checking the name for validity. A single dot at the end of the string is -// ignored. -func ParseName(s string) (Name, error) { - b := bytes.TrimSuffix([]byte(s), []byte(".")) - if len(b) == 0 { - // bytes.Split(b, ".") would return [""] in this case - return NewName([][]byte{}) - } else { - return NewName(bytes.Split(b, []byte("."))) - } -} - -// String returns a reversible string representation of name. Labels are -// separated by dots, and any bytes in a label that are outside the set -// [0-9A-Za-z-] are replaced with a \xXX hex escape sequence. -func (name Name) String() string { - if len(name) == 0 { - return "." - } - - var buf strings.Builder - for i, label := range name { - if i > 0 { - buf.WriteByte('.') - } - for _, b := range label { - if b == '-' || - ('0' <= b && b <= '9') || - ('A' <= b && b <= 'Z') || - ('a' <= b && b <= 'z') { - buf.WriteByte(b) - } else { - fmt.Fprintf(&buf, "\\x%02x", b) - } - } - } - return buf.String() -} - -// TrimSuffix returns a Name with the given suffix removed, if it was present. -// The second return value indicates whether the suffix was present. If the -// suffix was not present, the first return value is nil. -func (name Name) TrimSuffix(suffix Name) (Name, bool) { - if len(name) < len(suffix) { - return nil, false - } - split := len(name) - len(suffix) - fore, aft := name[:split], name[split:] - for i := 0; i < len(aft); i++ { - if !bytes.Equal(bytes.ToLower(aft[i]), bytes.ToLower(suffix[i])) { - return nil, false - } - } - return fore, true -} - -// Message represents a DNS message. -// -// https://tools.ietf.org/html/rfc1035#section-4.1 -type Message struct { - ID uint16 - Flags uint16 - - Question []Question - Answer []RR - Authority []RR - Additional []RR -} - -// Opcode extracts the OPCODE part of the Flags field. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.1 -func (message *Message) Opcode() uint16 { - return (message.Flags >> 11) & 0xf -} - -// Rcode extracts the RCODE part of the Flags field. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.1 -func (message *Message) Rcode() uint16 { - return message.Flags & 0x000f -} - -// Question represents an entry in the question section of a message. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.2 -type Question struct { - Name Name - Type uint16 - Class uint16 -} - -// RR represents a resource record. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.3 -type RR struct { - Name Name - Type uint16 - Class uint16 - TTL uint32 - Data []byte -} - -// readName parses a DNS name from r. It leaves r positioned just after the -// parsed name. -func readName(r io.ReadSeeker) (Name, error) { - var labels [][]byte - // We limit the number of compression pointers we are willing to follow. - numPointers := 0 - // If we followed any compression pointers, we must finally seek to just - // past the first pointer. - var seekTo int64 -loop: - for { - var labelType byte - err := binary.Read(r, binary.BigEndian, &labelType) - if err != nil { - return nil, err - } - - switch labelType & 0xc0 { - case 0x00: - // This is an ordinary label. - // https://tools.ietf.org/html/rfc1035#section-3.1 - length := int(labelType & 0x3f) - if length == 0 { - break loop - } - label := make([]byte, length) - _, err := io.ReadFull(r, label) - if err != nil { - return nil, err - } - labels = append(labels, label) - case 0xc0: - // This is a compression pointer. - // https://tools.ietf.org/html/rfc1035#section-4.1.4 - upper := labelType & 0x3f - var lower byte - err := binary.Read(r, binary.BigEndian, &lower) - if err != nil { - return nil, err - } - offset := (uint16(upper) << 8) | uint16(lower) - - if numPointers == 0 { - // The first time we encounter a pointer, - // remember our position so we can seek back to - // it when done. - seekTo, err = r.Seek(0, io.SeekCurrent) - if err != nil { - return nil, err - } - } - numPointers++ - if numPointers > compressionPointerLimit { - return nil, ErrTooManyPointers - } - - // Follow the pointer and continue. - _, err = r.Seek(int64(offset), io.SeekStart) - if err != nil { - return nil, err - } - default: - // "The 10 and 01 combinations are reserved for future - // use." - return nil, ErrReservedLabelType - } - } - // If we followed any pointers, then seek back to just after the first - // one. - if numPointers > 0 { - _, err := r.Seek(seekTo, io.SeekStart) - if err != nil { - return nil, err - } - } - return NewName(labels) -} - -// readQuestion parses one entry from the Question section. It leaves r -// positioned just after the parsed entry. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.2 -func readQuestion(r io.ReadSeeker) (Question, error) { - var question Question - var err error - question.Name, err = readName(r) - if err != nil { - return question, err - } - for _, ptr := range []*uint16{&question.Type, &question.Class} { - err := binary.Read(r, binary.BigEndian, ptr) - if err != nil { - return question, err - } - } - - return question, nil -} - -// readRR parses one resource record. It leaves r positioned just after the -// parsed resource record. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.3 -func readRR(r io.ReadSeeker) (RR, error) { - var rr RR - var err error - rr.Name, err = readName(r) - if err != nil { - return rr, err - } - for _, ptr := range []*uint16{&rr.Type, &rr.Class} { - err := binary.Read(r, binary.BigEndian, ptr) - if err != nil { - return rr, err - } - } - err = binary.Read(r, binary.BigEndian, &rr.TTL) - if err != nil { - return rr, err - } - var rdLength uint16 - err = binary.Read(r, binary.BigEndian, &rdLength) - if err != nil { - return rr, err - } - rr.Data = make([]byte, rdLength) - _, err = io.ReadFull(r, rr.Data) - if err != nil { - return rr, err - } - - return rr, nil -} - -// readMessage parses a complete DNS message. It leaves r positioned just after -// the parsed message. -func readMessage(r io.ReadSeeker) (Message, error) { - var message Message - - // Header section - // https://tools.ietf.org/html/rfc1035#section-4.1.1 - var qdCount, anCount, nsCount, arCount uint16 - for _, ptr := range []*uint16{ - &message.ID, &message.Flags, - &qdCount, &anCount, &nsCount, &arCount, - } { - err := binary.Read(r, binary.BigEndian, ptr) - if err != nil { - return message, err - } - } - - // Question section - // https://tools.ietf.org/html/rfc1035#section-4.1.2 - for i := 0; i < int(qdCount); i++ { - question, err := readQuestion(r) - if err != nil { - return message, err - } - message.Question = append(message.Question, question) - } - - // Answer, Authority, and Additional sections - // https://tools.ietf.org/html/rfc1035#section-4.1.3 - for _, rec := range []struct { - ptr *[]RR - count uint16 - }{ - {&message.Answer, anCount}, - {&message.Authority, nsCount}, - {&message.Additional, arCount}, - } { - for i := 0; i < int(rec.count); i++ { - rr, err := readRR(r) - if err != nil { - return message, err - } - *rec.ptr = append(*rec.ptr, rr) - } - } - - return message, nil -} - -// MessageFromWireFormat parses a message from buf and returns a Message object. -// It returns ErrTrailingBytes if there are bytes remaining in buf after parsing -// is done. -func MessageFromWireFormat(buf []byte) (Message, error) { - r := bytes.NewReader(buf) - message, err := readMessage(r) - if err == io.EOF { - err = io.ErrUnexpectedEOF - } else if err == nil { - // Check for trailing bytes. - _, err = r.ReadByte() - if err == io.EOF { - err = nil - } else if err == nil { - err = ErrTrailingBytes - } - } - return message, err -} - -// messageBuilder manages the state of serializing a DNS message. Its main -// function is to keep track of names already written for the purpose of name -// compression. -type messageBuilder struct { - w bytes.Buffer - nameCache map[string]int -} - -// newMessageBuilder creates a new messageBuilder with an empty name cache. -func newMessageBuilder() *messageBuilder { - return &messageBuilder{ - nameCache: make(map[string]int), - } -} - -// Bytes returns the serialized DNS message as a slice of bytes. -func (builder *messageBuilder) Bytes() []byte { - return builder.w.Bytes() -} - -// WriteName appends name to the in-progress messageBuilder, employing -// compression pointers to previously written names if possible. -func (builder *messageBuilder) WriteName(name Name) { - // https://tools.ietf.org/html/rfc1035#section-3.1 - for i := range name { - // Has this suffix already been encoded in the message? - if ptr, ok := builder.nameCache[name[i:].String()]; ok && ptr&0x3fff == ptr { - // If so, we can write a compression pointer. - binary.Write(&builder.w, binary.BigEndian, uint16(0xc000|ptr)) - return - } - // Not cached; we must encode this label verbatim. Store a cache - // entry pointing to the beginning of it. - builder.nameCache[name[i:].String()] = builder.w.Len() - length := len(name[i]) - if length == 0 || length > 63 { - panic(length) - } - builder.w.WriteByte(byte(length)) - builder.w.Write(name[i]) - } - builder.w.WriteByte(0) -} - -// WriteQuestion appends a Question section entry to the in-progress -// messageBuilder. -func (builder *messageBuilder) WriteQuestion(question *Question) { - // https://tools.ietf.org/html/rfc1035#section-4.1.2 - builder.WriteName(question.Name) - binary.Write(&builder.w, binary.BigEndian, question.Type) - binary.Write(&builder.w, binary.BigEndian, question.Class) -} - -// WriteRR appends a resource record to the in-progress messageBuilder. It -// returns ErrIntegerOverflow if the length of rr.Data does not fit in 16 bits. -func (builder *messageBuilder) WriteRR(rr *RR) error { - // https://tools.ietf.org/html/rfc1035#section-4.1.3 - builder.WriteName(rr.Name) - binary.Write(&builder.w, binary.BigEndian, rr.Type) - binary.Write(&builder.w, binary.BigEndian, rr.Class) - binary.Write(&builder.w, binary.BigEndian, rr.TTL) - rdLength := uint16(len(rr.Data)) - if int(rdLength) != len(rr.Data) { - return ErrIntegerOverflow - } - binary.Write(&builder.w, binary.BigEndian, rdLength) - builder.w.Write(rr.Data) - return nil -} - -// WriteMessage appends a complete DNS message to the in-progress -// messageBuilder. It returns ErrIntegerOverflow if the number of entries in any -// section, or the length of the data in any resource record, does not fit in 16 -// bits. -func (builder *messageBuilder) WriteMessage(message *Message) error { - // Header section - // https://tools.ietf.org/html/rfc1035#section-4.1.1 - binary.Write(&builder.w, binary.BigEndian, message.ID) - binary.Write(&builder.w, binary.BigEndian, message.Flags) - for _, count := range []int{ - len(message.Question), - len(message.Answer), - len(message.Authority), - len(message.Additional), - } { - count16 := uint16(count) - if int(count16) != count { - return ErrIntegerOverflow - } - binary.Write(&builder.w, binary.BigEndian, count16) - } - - // Question section - // https://tools.ietf.org/html/rfc1035#section-4.1.2 - for _, question := range message.Question { - builder.WriteQuestion(&question) - } - - // Answer, Authority, and Additional sections - // https://tools.ietf.org/html/rfc1035#section-4.1.3 - for _, rrs := range [][]RR{message.Answer, message.Authority, message.Additional} { - for _, rr := range rrs { - err := builder.WriteRR(&rr) - if err != nil { - return err - } - } - } - - return nil -} - -// WireFormat encodes a Message as a slice of bytes in DNS wire format. It -// returns ErrIntegerOverflow if the number of entries in any section, or the -// length of the data in any resource record, does not fit in 16 bits. -func (message *Message) WireFormat() ([]byte, error) { - builder := newMessageBuilder() - err := builder.WriteMessage(message) - if err != nil { - return nil, err - } - return builder.Bytes(), nil -} - -// DecodeRDataTXT decodes TXT-DATA (as found in the RDATA for a resource record -// with TYPE=TXT) as a raw byte slice, by concatenating all the -// s it contains. -// -// https://tools.ietf.org/html/rfc1035#section-3.3.14 -func DecodeRDataTXT(p []byte) ([]byte, error) { - var buf bytes.Buffer - for { - if len(p) == 0 { - return nil, io.ErrUnexpectedEOF - } - n := int(p[0]) - p = p[1:] - if len(p) < n { - return nil, io.ErrUnexpectedEOF - } - buf.Write(p[:n]) - p = p[n:] - if len(p) == 0 { - break - } - } - return buf.Bytes(), nil -} - -// EncodeRDataTXT encodes a slice of bytes as TXT-DATA, as appropriate for the -// RDATA of a resource record with TYPE=TXT. No length restriction is enforced -// here; that must be checked at a higher level. -// -// https://tools.ietf.org/html/rfc1035#section-3.3.14 -func EncodeRDataTXT(p []byte) []byte { - // https://tools.ietf.org/html/rfc1035#section-3.3 - // https://tools.ietf.org/html/rfc1035#section-3.3.14 - // TXT data is a sequence of one or more s, where - // is a length octet followed by that number of - // octets. - var buf bytes.Buffer - for len(p) > 255 { - buf.WriteByte(255) - buf.Write(p[:255]) - p = p[255:] - } - // Must write here, even if len(p) == 0, because it's "*one or more* - // s". - buf.WriteByte(byte(len(p))) - buf.Write(p) - return buf.Bytes() -} diff --git a/transport/internet/finalmask/xdns/dns_test.go b/transport/internet/finalmask/xdns/dns_test.go deleted file mode 100644 index 7eac084e6..000000000 --- a/transport/internet/finalmask/xdns/dns_test.go +++ /dev/null @@ -1,953 +0,0 @@ -package xdns - -import ( - "bytes" - "fmt" - "io" - "strconv" - "strings" - "testing" -) - -func namesEqual(a, b Name) bool { - if len(a) != len(b) { - return false - } - for i := 0; i < len(a); i++ { - if !bytes.Equal(a[i], b[i]) { - return false - } - } - return true -} - -func TestName(t *testing.T) { - for _, test := range []struct { - labels [][]byte - err error - s string - }{ - {[][]byte{}, nil, "."}, - {[][]byte{[]byte("test")}, nil, "test"}, - {[][]byte{[]byte("a"), []byte("b"), []byte("c")}, nil, "a.b.c"}, - - {[][]byte{{}}, ErrZeroLengthLabel, ""}, - {[][]byte{[]byte("a"), {}, []byte("c")}, ErrZeroLengthLabel, ""}, - - // 63 octets. - { - [][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE")}, - nil, - "0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE", - }, - // 64 octets. - {[][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDEF")}, ErrLabelTooLong, ""}, - - // 64+64+64+62 octets. - { - [][]byte{ - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC"), - }, - nil, - "0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC", - }, - // 64+64+64+63 octets. - {[][]byte{ - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCD"), - }, ErrNameTooLong, ""}, - // 127 one-octet labels. - { - [][]byte{ - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - }, - nil, - "0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E", - }, - // 128 one-octet labels. - {[][]byte{ - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - }, ErrNameTooLong, ""}, - } { - // Test that NewName returns proper error codes, and otherwise - // returns an equal slice of labels. - name, err := NewName(test.labels) - if err != test.err || (err == nil && !namesEqual(name, test.labels)) { - t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)", - test.labels, name, err, test.labels, test.err) - continue - } - if test.err != nil { - continue - } - - // Test that the string version of the name comes out as - // expected. - s := name.String() - if s != test.s { - t.Errorf("%+q became string %+q, expected %+q", test.labels, s, test.s) - continue - } - - // Test that parsing from a string back to a Name results in the - // original slice of labels. - name, err = ParseName(s) - if err != nil || !namesEqual(name, test.labels) { - t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)", - test.labels, s, name, err, test.labels, nil) - continue - } - // A trailing dot should be ignored. - if !strings.HasSuffix(s, ".") { - dotName, dotErr := ParseName(s + ".") - if dotErr != err || !namesEqual(dotName, name) { - t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)", - test.labels, s+".", dotName, dotErr, name, err) - continue - } - } - } -} - -func TestParseName(t *testing.T) { - for _, test := range []struct { - s string - name Name - err error - }{ - // This case can't be tested by TestName above because String - // will never produce "" (it produces "." instead). - {"", [][]byte{}, nil}, - } { - name, err := ParseName(test.s) - if err != test.err || (err == nil && !namesEqual(name, test.name)) { - t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)", - test.s, name, err, test.name, test.err) - continue - } - } -} - -func unescapeString(s string) ([][]byte, error) { - if s == "." { - return [][]byte{}, nil - } - - var result [][]byte - for _, label := range strings.Split(s, ".") { - var buf bytes.Buffer - i := 0 - for i < len(label) { - switch label[i] { - case '\\': - if i+3 >= len(label) { - return nil, fmt.Errorf("truncated escape sequence at index %v", i) - } - if label[i+1] != 'x' { - return nil, fmt.Errorf("malformed escape sequence at index %v", i) - } - b, err := strconv.ParseUint(string(label[i+2:i+4]), 16, 8) - if err != nil { - return nil, fmt.Errorf("malformed hex sequence at index %v", i+2) - } - buf.WriteByte(byte(b)) - i += 4 - default: - buf.WriteByte(label[i]) - i++ - } - } - result = append(result, buf.Bytes()) - } - return result, nil -} - -func TestNameString(t *testing.T) { - for _, test := range []struct { - name Name - s string - }{ - {[][]byte{}, "."}, - {[][]byte{[]byte("\x00"), []byte("a.b"), []byte("c\nd\\")}, "\\x00.a\\x2eb.c\\x0ad\\x5c"}, - {[][]byte{ - []byte("\x00\x01\x02\x03\x04\x05\x06\x07\x08\t\n\x0b\x0c\r\x0e\x0f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f !\"#$%&'()*+,-./0123456789:;<=>"), - []byte("?@ABCDEFGHIJKLMNOPQRSTUVWXYZ[\\]^_`abcdefghijklmnopqrstuvwxyz{|}"), - []byte("~\x7f\x80\x81\x82\x83\x84\x85\x86\x87\x88\x89\x8a\x8b\x8c\x8d\x8e\x8f\x90\x91\x92\x93\x94\x95\x96\x97\x98\x99\x9a\x9b\x9c\x9d\x9e\x9f\xa0\xa1\xa2\xa3\xa4\xa5\xa6\xa7\xa8\xa9\xaa\xab\xac\xad\xae\xaf\xb0\xb1\xb2\xb3\xb4\xb5\xb6\xb7\xb8\xb9\xba\xbb\xbc"), - []byte("\xbd\xbe\xbf\xc0\xc1\xc2\xc3\xc4\xc5\xc6\xc7\xc8\xc9\xca\xcb\xcc\xcd\xce\xcf\xd0\xd1\xd2\xd3\xd4\xd5\xd6\xd7\xd8\xd9\xda\xdb\xdc\xdd\xde\xdf\xe0\xe1\xe2\xe3\xe4\xe5\xe6\xe7\xe8\xe9\xea\xeb\xec\xed\xee\xef\xf0\xf1\xf2\xf3\xf4\xf5\xf6\xf7\xf8\xf9\xfa\xfb"), - []byte("\xfc\xfd\xfe\xff"), - }, "\\x00\\x01\\x02\\x03\\x04\\x05\\x06\\x07\\x08\\x09\\x0a\\x0b\\x0c\\x0d\\x0e\\x0f\\x10\\x11\\x12\\x13\\x14\\x15\\x16\\x17\\x18\\x19\\x1a\\x1b\\x1c\\x1d\\x1e\\x1f\\x20\\x21\\x22\\x23\\x24\\x25\\x26\\x27\\x28\\x29\\x2a\\x2b\\x2c-\\x2e\\x2f0123456789\\x3a\\x3b\\x3c\\x3d\\x3e.\\x3f\\x40ABCDEFGHIJKLMNOPQRSTUVWXYZ\\x5b\\x5c\\x5d\\x5e\\x5f\\x60abcdefghijklmnopqrstuvwxyz\\x7b\\x7c\\x7d.\\x7e\\x7f\\x80\\x81\\x82\\x83\\x84\\x85\\x86\\x87\\x88\\x89\\x8a\\x8b\\x8c\\x8d\\x8e\\x8f\\x90\\x91\\x92\\x93\\x94\\x95\\x96\\x97\\x98\\x99\\x9a\\x9b\\x9c\\x9d\\x9e\\x9f\\xa0\\xa1\\xa2\\xa3\\xa4\\xa5\\xa6\\xa7\\xa8\\xa9\\xaa\\xab\\xac\\xad\\xae\\xaf\\xb0\\xb1\\xb2\\xb3\\xb4\\xb5\\xb6\\xb7\\xb8\\xb9\\xba\\xbb\\xbc.\\xbd\\xbe\\xbf\\xc0\\xc1\\xc2\\xc3\\xc4\\xc5\\xc6\\xc7\\xc8\\xc9\\xca\\xcb\\xcc\\xcd\\xce\\xcf\\xd0\\xd1\\xd2\\xd3\\xd4\\xd5\\xd6\\xd7\\xd8\\xd9\\xda\\xdb\\xdc\\xdd\\xde\\xdf\\xe0\\xe1\\xe2\\xe3\\xe4\\xe5\\xe6\\xe7\\xe8\\xe9\\xea\\xeb\\xec\\xed\\xee\\xef\\xf0\\xf1\\xf2\\xf3\\xf4\\xf5\\xf6\\xf7\\xf8\\xf9\\xfa\\xfb.\\xfc\\xfd\\xfe\\xff"}, - } { - s := test.name.String() - if s != test.s { - t.Errorf("%+q escaped to %+q, expected %+q", test.name, s, test.s) - continue - } - unescaped, err := unescapeString(s) - if err != nil { - t.Errorf("%+q unescaping %+q resulted in error %v", test.name, s, err) - continue - } - if !namesEqual(Name(unescaped), test.name) { - t.Errorf("%+q roundtripped through %+q to %+q", test.name, s, unescaped) - continue - } - } -} - -func TestNameTrimSuffix(t *testing.T) { - for _, test := range []struct { - name, suffix string - trimmed string - ok bool - }{ - {"", "", ".", true}, - {".", ".", ".", true}, - {"abc", "", "abc", true}, - {"abc", ".", "abc", true}, - {"", "abc", ".", false}, - {".", "abc", ".", false}, - {"example.com", "com", "example", true}, - {"example.com", "net", ".", false}, - {"example.com", "example.com", ".", true}, - {"example.com", "test.com", ".", false}, - {"example.com", "xample.com", ".", false}, - {"example.com", "example", ".", false}, - {"example.com", "COM", "example", true}, - {"EXAMPLE.COM", "com", "EXAMPLE", true}, - } { - tmp, ok := mustParseName(test.name).TrimSuffix(mustParseName(test.suffix)) - trimmed := tmp.String() - if ok != test.ok || trimmed != test.trimmed { - t.Errorf("TrimSuffix %+q %+q returned (%+q, %v), expected (%+q, %v)", - test.name, test.suffix, trimmed, ok, test.trimmed, test.ok) - continue - } - } -} - -func TestReadName(t *testing.T) { - // Good tests. - for _, test := range []struct { - start int64 - end int64 - input string - s string - }{ - // Empty name. - {0, 1, "\x00abcd", "."}, - // No pointers. - {12, 25, "AAAABBBBCCCC\x07example\x03com\x00", "example.com"}, - // Backward pointer. - {25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c", "sub.example.com"}, - // Forward pointer. - {0, 4, "\x01a\xc0\x04\x03bcd\x00", "a.bcd"}, - // Two backwards pointers. - {31, 38, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c\x04sub2\xc0\x19", "sub2.sub.example.com"}, - // Forward then backward pointer. - {25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x1f\x04sub2\xc0\x0c", "sub.sub2.example.com"}, - // Overlapping codons. - {0, 4, "\x01a\xc0\x03bcd\x00", "a.bcd"}, - // Pointer to empty label. - {0, 10, "\x07example\xc0\x0a\x00", "example"}, - {1, 11, "\x00\x07example\xc0\x00", "example"}, - // Pointer to pointer to empty label. - {0, 10, "\x07example\xc0\x0a\xc0\x0c\x00", "example"}, - {1, 11, "\x00\x07example\xc0\x0c\xc0\x00", "example"}, - } { - r := bytes.NewReader([]byte(test.input)) - _, err := r.Seek(test.start, io.SeekStart) - if err != nil { - panic(err) - } - name, err := readName(r) - if err != nil { - t.Errorf("%+q returned error %s", test.input, err) - continue - } - s := name.String() - if s != test.s { - t.Errorf("%+q returned %+q, expected %+q", test.input, s, test.s) - continue - } - cur, _ := r.Seek(0, io.SeekCurrent) - if cur != test.end { - t.Errorf("%+q left offset %d, expected %d", test.input, cur, test.end) - continue - } - } - - // Bad tests. - for _, test := range []struct { - start int64 - input string - err error - }{ - {0, "", io.ErrUnexpectedEOF}, - // Reserved label type. - {0, "\x80example", ErrReservedLabelType}, - // Reserved label type. - {0, "\x40example", ErrReservedLabelType}, - // No Terminating empty label. - {0, "\x07example\x03com", io.ErrUnexpectedEOF}, - // Pointer past end of buffer. - {0, "\x07example\xc0\xff", io.ErrUnexpectedEOF}, - // Pointer to self. - {0, "\x07example\x03com\xc0\x0c", ErrTooManyPointers}, - // Pointer to self with intermediate label. - {0, "\x07example\x03com\xc0\x08", ErrTooManyPointers}, - // Two pointers that point to each other. - {0, "\xc0\x02\xc0\x00", ErrTooManyPointers}, - // Two pointers that point to each other, with intermediate labels. - {0, "\x01a\xc0\x04\x01b\xc0\x00", ErrTooManyPointers}, - // EOF while reading label. - {0, "\x0aexample", io.ErrUnexpectedEOF}, - // EOF before second byte of pointer. - {0, "\xc0", io.ErrUnexpectedEOF}, - {0, "\x07example\xc0", io.ErrUnexpectedEOF}, - } { - r := bytes.NewReader([]byte(test.input)) - _, err := r.Seek(test.start, io.SeekStart) - if err != nil { - panic(err) - } - name, err := readName(r) - if err == io.EOF { - err = io.ErrUnexpectedEOF - } - if err != test.err { - t.Errorf("%+q returned (%+q, %v), expected %v", test.input, name, err, test.err) - continue - } - } -} - -func mustParseName(s string) Name { - name, err := ParseName(s) - if err != nil { - panic(err) - } - return name -} - -func questionsEqual(a, b *Question) bool { - if !namesEqual(a.Name, b.Name) { - return false - } - if a.Type != b.Type || a.Class != b.Class { - return false - } - return true -} - -func rrsEqual(a, b *RR) bool { - if !namesEqual(a.Name, b.Name) { - return false - } - if a.Type != b.Type || a.Class != b.Class || a.TTL != b.TTL { - return false - } - if !bytes.Equal(a.Data, b.Data) { - return false - } - return true -} - -func messagesEqual(a, b *Message) bool { - if a.ID != b.ID || a.Flags != b.Flags { - return false - } - if len(a.Question) != len(b.Question) { - return false - } - for i := 0; i < len(a.Question); i++ { - if !questionsEqual(&a.Question[i], &b.Question[i]) { - return false - } - } - for _, rec := range []struct{ rrA, rrB []RR }{ - {a.Answer, b.Answer}, - {a.Authority, b.Authority}, - {a.Additional, b.Additional}, - } { - if len(rec.rrA) != len(rec.rrB) { - return false - } - for i := 0; i < len(rec.rrA); i++ { - if !rrsEqual(&rec.rrA[i], &rec.rrB[i]) { - return false - } - } - } - return true -} - -func TestMessageFromWireFormat(t *testing.T) { - for _, test := range []struct { - buf string - expected Message - err error - }{ - { - "\x12\x34", - Message{}, - io.ErrUnexpectedEOF, - }, - { - "\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01", - Message{ - ID: 0x1234, - Flags: 0x0100, - Question: []Question{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - }, - }, - Answer: []RR{}, - Authority: []RR{}, - Additional: []RR{}, - }, - nil, - }, - { - "\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01X", - Message{}, - ErrTrailingBytes, - }, - { - "\x12\x34\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01\x03www\x07example\x03com\x00\x00\x01\x00\x01\x00\x00\x00\x80\x00\x04\xc0\x00\x02\x01", - Message{ - ID: 0x1234, - Flags: 0x8180, - Question: []Question{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - }, - }, - Answer: []RR{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - TTL: 128, - Data: []byte{192, 0, 2, 1}, - }, - }, - Authority: []RR{}, - Additional: []RR{}, - }, - nil, - }, - } { - message, err := MessageFromWireFormat([]byte(test.buf)) - if err != test.err || (err == nil && !messagesEqual(&message, &test.expected)) { - t.Errorf("%+q\nreturned (%+v, %v)\nexpected (%+v, %v)", - test.buf, message, err, test.expected, test.err) - continue - } - } -} - -func TestMessageWireFormatRoundTrip(t *testing.T) { - for _, message := range []Message{ - { - ID: 0x1234, - Flags: 0x0100, - Question: []Question{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - }, - { - Name: mustParseName("www2.example.com"), - Type: 2, - Class: 2, - }, - }, - Answer: []RR{ - { - Name: mustParseName("abc"), - Type: 2, - Class: 3, - TTL: 0xffffffff, - Data: []byte{1}, - }, - { - Name: mustParseName("xyz"), - Type: 2, - Class: 3, - TTL: 255, - Data: []byte{}, - }, - }, - Authority: []RR{ - { - Name: mustParseName("."), - Type: 65535, - Class: 65535, - TTL: 0, - Data: []byte("XXXXXXXXXXXXXXXXXXX"), - }, - }, - Additional: []RR{}, - }, - } { - buf, err := message.WireFormat() - if err != nil { - t.Errorf("%+v cannot make wire format: %v", message, err) - continue - } - message2, err := MessageFromWireFormat(buf) - if err != nil { - t.Errorf("%+q cannot parse wire format: %v", buf, err) - continue - } - if !messagesEqual(&message, &message2) { - t.Errorf("messages unequal\nbefore: %+v\n after: %+v", message, message2) - continue - } - } -} - -func TestDecodeRDataTXT(t *testing.T) { - for _, test := range []struct { - p []byte - decoded []byte - err error - }{ - {[]byte{}, nil, io.ErrUnexpectedEOF}, - {[]byte("\x00"), []byte{}, nil}, - {[]byte("\x01"), nil, io.ErrUnexpectedEOF}, - } { - decoded, err := DecodeRDataTXT(test.p) - if err != test.err || (err == nil && !bytes.Equal(decoded, test.decoded)) { - t.Errorf("%+q\nreturned (%+q, %v)\nexpected (%+q, %v)", - test.p, decoded, err, test.decoded, test.err) - continue - } - } -} - -func TestEncodeRDataTXT(t *testing.T) { - // Encoding 0 bytes needs to return at least a single length octet of - // zero, not an empty slice. - p := make([]byte, 0) - encoded := EncodeRDataTXT(p) - if len(encoded) < 0 { - t.Errorf("EncodeRDataTXT(%v) returned %v", p, encoded) - } - - // 255 bytes should be able to be encoded into 256 bytes. - p = make([]byte, 255) - encoded = EncodeRDataTXT(p) - if len(encoded) > 256 { - t.Errorf("EncodeRDataTXT(%d bytes) returned %d bytes", len(p), len(encoded)) - } - - fmt.Println(EncodeRDataTXT(nil)) - fmt.Println(computeMaxEncodedPayload(maxUDPPayload)) -} - -func TestRDataTXTRoundTrip(t *testing.T) { - for _, p := range [][]byte{ - {}, - []byte("\x00"), - { - 0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f, - 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f, - 0x20, 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, - 0x30, 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, - 0x40, 0x41, 0x42, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48, 0x49, 0x4a, 0x4b, 0x4c, 0x4d, 0x4e, 0x4f, - 0x50, 0x51, 0x52, 0x53, 0x54, 0x55, 0x56, 0x57, 0x58, 0x59, 0x5a, 0x5b, 0x5c, 0x5d, 0x5e, 0x5f, - 0x60, 0x61, 0x62, 0x63, 0x64, 0x65, 0x66, 0x67, 0x68, 0x69, 0x6a, 0x6b, 0x6c, 0x6d, 0x6e, 0x6f, - 0x70, 0x71, 0x72, 0x73, 0x74, 0x75, 0x76, 0x77, 0x78, 0x79, 0x7a, 0x7b, 0x7c, 0x7d, 0x7e, 0x7f, - 0x80, 0x81, 0x82, 0x83, 0x84, 0x85, 0x86, 0x87, 0x88, 0x89, 0x8a, 0x8b, 0x8c, 0x8d, 0x8e, 0x8f, - 0x90, 0x91, 0x92, 0x93, 0x94, 0x95, 0x96, 0x97, 0x98, 0x99, 0x9a, 0x9b, 0x9c, 0x9d, 0x9e, 0x9f, - 0xa0, 0xa1, 0xa2, 0xa3, 0xa4, 0xa5, 0xa6, 0xa7, 0xa8, 0xa9, 0xaa, 0xab, 0xac, 0xad, 0xae, 0xaf, - 0xb0, 0xb1, 0xb2, 0xb3, 0xb4, 0xb5, 0xb6, 0xb7, 0xb8, 0xb9, 0xba, 0xbb, 0xbc, 0xbd, 0xbe, 0xbf, - 0xc0, 0xc1, 0xc2, 0xc3, 0xc4, 0xc5, 0xc6, 0xc7, 0xc8, 0xc9, 0xca, 0xcb, 0xcc, 0xcd, 0xce, 0xcf, - 0xd0, 0xd1, 0xd2, 0xd3, 0xd4, 0xd5, 0xd6, 0xd7, 0xd8, 0xd9, 0xda, 0xdb, 0xdc, 0xdd, 0xde, 0xdf, - 0xe0, 0xe1, 0xe2, 0xe3, 0xe4, 0xe5, 0xe6, 0xe7, 0xe8, 0xe9, 0xea, 0xeb, 0xec, 0xed, 0xee, 0xef, - 0xf0, 0xf1, 0xf2, 0xf3, 0xf4, 0xf5, 0xf6, 0xf7, 0xf8, 0xf9, 0xfa, 0xfb, 0xfc, 0xfd, 0xfe, 0xff, - }, - } { - rdata := EncodeRDataTXT(p) - decoded, err := DecodeRDataTXT(rdata) - if err != nil || !bytes.Equal(decoded, p) { - t.Errorf("%+q returned (%+q, %v)", p, decoded, err) - continue - } - } -} - -func TestIPAnswerPayloadRoundTrip(t *testing.T) { - for _, rrType := range []uint16{RRTypeA, RRTypeAAAA} { - for _, payload := range [][]byte{ - {}, - {0x01}, - []byte("hello world"), - bytes.Repeat([]byte{0xab}, payloadChunkSizeForType(rrType)*3+1), - } { - question := Question{ - Name: mustParseName("example.com"), - Type: rrType, - Class: ClassIN, - } - answers, err := answersForPayload(question, responseTTL, payload) - if err != nil { - t.Fatalf("answersForPayload(%d) err = %v", rrType, err) - } - - if len(answers) > 1 { - answers[0], answers[len(answers)-1] = answers[len(answers)-1], answers[0] - } - - decoded := decodeResponsePayload(answers) - if !bytes.Equal(decoded, payload) { - t.Fatalf("rrType=%d decoded %x want %x", rrType, decoded, payload) - } - } - } -} - -func TestParseResolver(t *testing.T) { - tests := []struct { - resolver string - rrType uint16 - }{ - {"example.com+udp://1.1.1.1:53", RRTypeTXT}, - {"example.com:txt+udp://1.1.1.1:53", RRTypeTXT}, - {"example.com:a+udp://1.1.1.1:53", RRTypeA}, - {"example.com:aaaa+udp://1.1.1.1:53", RRTypeAAAA}, - } - - for _, test := range tests { - domain, server, rrType, err := parseResolver(test.resolver) - if err != nil { - t.Fatalf("parseResolver(%q) err = %v", test.resolver, err) - } - if domain.String() != "example.com" || server != "1.1.1.1:53" || rrType != test.rrType { - t.Fatalf("parseResolver(%q) = (%q, %q, %d)", test.resolver, domain.String(), server, rrType) - } - } -} - -func TestParseDomainSpec(t *testing.T) { - tests := []struct { - spec string - def string - rrType uint16 - wantErr bool - }{ - {"example.com", "", 0, false}, - {"example.com", "txt", RRTypeTXT, false}, - {"example.com:a", "", RRTypeA, false}, - {"example.com:aaaa", "", RRTypeAAAA, false}, - {"example.com:doh", "", 0, true}, - } - - for _, test := range tests { - got, err := parseDomainSpec(test.spec, test.def) - if test.wantErr { - if err == nil { - t.Fatalf("parseDomainSpec(%q, %q) err = nil", test.spec, test.def) - } - continue - } - if err != nil { - t.Fatalf("parseDomainSpec(%q, %q) err = %v", test.spec, test.def, err) - } - if got.name.String() != "example.com" || got.rrType != test.rrType { - t.Fatalf("parseDomainSpec(%q, %q) = (%q, %d)", test.spec, test.def, got.name.String(), got.rrType) - } - } -} - -func TestResponseForMethodRestriction(t *testing.T) { - query := &Message{ - ID: 1, - Flags: 0x0100, - Question: []Question{{ - Name: mustParseName("abc.example.com"), - Type: RRTypeTXT, - Class: ClassIN, - }}, - Additional: []RR{{ - Name: Name{}, - Type: RRTypeOPT, - Class: 4096, - }}, - } - - resp, _ := responseFor(query, []domainSpec{{name: mustParseName("example.com"), rrType: RRTypeA}}) - if resp == nil || resp.Rcode() != RcodeNameError { - t.Fatalf("responseFor method restriction rcode = %v", resp) - } - - resp, _ = responseFor(query, []domainSpec{{name: mustParseName("example.com")}}) - if resp == nil || resp.Rcode() != RcodeNoError { - t.Fatalf("responseFor unrestricted rcode = %v", resp) - } -} diff --git a/transport/internet/finalmask/xdns/domain.go b/transport/internet/finalmask/xdns/domain.go new file mode 100644 index 000000000..c7ca40870 --- /dev/null +++ b/transport/internet/finalmask/xdns/domain.go @@ -0,0 +1,215 @@ +package xdns + +import ( + "encoding/base32" + "errors" + "fmt" + "strings" + + "golang.org/x/net/dns/dnsmessage" + "golang.org/x/net/idna" +) + +func Lower(c byte) byte { + if c >= 'A' && c <= 'Z' { + return c + ('a' - 'A') + } + return c +} + +func ToUpper(b []byte) { + for i, c := range b { + if c >= 'a' && c <= 'z' { + b[i] = c - 'a' + 'A' + } + } +} + +func ToLower(b []byte) { + for i, c := range b { + if c >= 'A' && c <= 'Z' { + b[i] = c - 'A' + 'a' + } + } +} + +func NewTable() ([256]int, [256]int) { + var t, t_ [256]int + for i := range t { + t[i] = base32Encoding.DecodedLen(i) + } + for i := range t_ { + t_[i] = base32Encoding.EncodedLen(i) + } + return t, t_ +} + +const ( + TypeA uint16 = 1 + TypeCNAME uint16 = 5 + TypeTXT uint16 = 16 + TypeAAAA uint16 = 28 +) + +var ( + base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding) + table, table_ = NewTable() + TypeMap = map[uint16]byte{ + TypeA: 0, + TypeCNAME: 1, + TypeTXT: 2, + TypeAAAA: 3, + } + TypeMap_ = map[byte]uint16{ + 0: TypeA, + 1: TypeCNAME, + 2: TypeTXT, + 3: TypeAAAA, + } +) + +type Domain struct { + name dnsmessage.Name + lenLimit int + labelLimit int + types []uint16 + edns0 uint16 + + cap int + lenMax int +} + +func NewDomain(domain string, lenLimit int, labelLimit int, types []uint16, edns0 uint16) (*Domain, error) { + if strings.Contains(domain, "..") { + return nil, errors.New("invalid domain") + } + if lenLimit < 0 || lenLimit > 255 { + return nil, errors.New("lenLimit < 0 || lenLimit > 255") + } + if labelLimit < 0 || labelLimit > 63 { + return nil, errors.New("labelLimit < 0 || labelLimit > 63") + } + if len(types) == 0 { + return nil, errors.New("empty types") + } + for i := range types { + switch types[i] { + case uint16(dnsmessage.TypeA), uint16(dnsmessage.TypeCNAME), uint16(dnsmessage.TypeTXT), uint16(dnsmessage.TypeAAAA): + default: + return nil, errors.New("unknown types") + } + } + if edns0 != 0 && (edns0 < 512 || edns0 > 4096) { + return nil, errors.New("edns0 != 0 && (edns0 < 512 || edns0 > 4096)") + } + + ascii, err := idna.ToASCII(domain) + if err != nil { + return nil, err + } + ascii = strings.Trim(ascii, ".") + + name, err := dnsmessage.NewName(domain + ".") + if err != nil { + return nil, err + } + + if lenLimit < int(name.Length)+1 { + return nil, errors.New("lenLimit < int(name.Length)+1") + } + n := (lenLimit - int(name.Length) - 1) / (labelLimit + 1) + left := (lenLimit - int(name.Length) - 1) % (labelLimit + 1) + total := n * labelLimit + if left > 1 { + total += left - 1 + } + cap := table[total] + if cap < 17 { + return nil, errors.New("cap < 17") + } + total = table_[cap] + lenMax := int(name.Length) + 1 + total + total/labelLimit + if total%labelLimit > 0 { + lenMax += 1 + } + return &Domain{ + name: name, + lenLimit: lenLimit, + labelLimit: labelLimit, + types: types, + edns0: edns0, + + cap: cap, + lenMax: lenMax, + }, nil +} + +func (d *Domain) Show() string { + return fmt.Sprint(d.name, d.cap) +} + +func (d *Domain) IsDomain(name dnsmessage.Name) bool { + if d.name.Length >= name.Length { + return false + } + i := d.name.Length + j := name.Length + for i > 0 { + i-- + j-- + if Lower(d.name.Data[i]) != Lower(name.Data[j]) { + return false + } + } + return true +} + +func (d *Domain) HasType(qtype uint16) bool { + for i := range d.types { + if d.types[i] == qtype { + return true + } + } + return false +} + +func (d *Domain) Encode(data []byte) dnsmessage.Name { + var name dnsmessage.Name + var encoded [255]byte + base32Encoding.Encode(encoded[:], data) + ToLower(encoded[:table_[len(data)]]) + b1 := name.Data[:0] + b2 := encoded[:table_[len(data)]] + for len(b2) > 0 { + size := min(len(b2), d.labelLimit) + b1 = append(b1, b2[:size]...) + b1 = append(b1, '.') + b2 = b2[size:] + } + b1 = append(b1, d.name.Data[:d.name.Length]...) + if len(b1) > 254 { + panic("len(b1) > 254") + } + name.Length = byte(len(b1)) + return name +} + +func (d *Domain) Decode(decoded *[255]byte, name dnsmessage.Name) int { + if !d.IsDomain(name) { + return 0 + } + var encoded [255]byte + b1 := encoded[:0] + b2 := name.Data[:name.Length-d.name.Length] + for i := range b2 { + if b2[i] != '.' { + b1 = append(b1, b2[i]) + } + } + ToUpper(b1) + n, err := base32Encoding.Decode(decoded[:], b1) + if err != nil { + return 0 + } + return n +} diff --git a/transport/internet/finalmask/xdns/frag.go b/transport/internet/finalmask/xdns/frag.go new file mode 100644 index 000000000..444f405af --- /dev/null +++ b/transport/internet/finalmask/xdns/frag.go @@ -0,0 +1,171 @@ +package xdns + +import ( + "sync" + "time" +) + +const ( + fragTTL = 8 * time.Second + fragSize = 4096 + fragClientIDSize = 16384 + fragCount = 4096 +) + +type FragKey struct { + clientID ClientID + fragID byte +} + +type FragEntry struct { + data [][]byte + size int + len int + total byte + deadline time.Time +} + +type FragManager struct { + m map[FragKey]*FragEntry + sizem map[ClientID]int + ch chan struct{} + mu sync.Mutex +} + +func NewFragManager() *FragManager { + m := &FragManager{ + m: make(map[FragKey]*FragEntry), + sizem: make(map[ClientID]int), + ch: make(chan struct{}), + } + go m.gc() + return m +} + +func (m *FragManager) closed() bool { + select { + case <-m.ch: + return true + default: + return false + } +} + +func (m *FragManager) removeEntey(k FragKey, e *FragEntry) { + m.sizem[k.clientID] -= e.size + delete(m.m, k) +} + +func (m *FragManager) tryRemove() { + if len(m.m) < fragCount { + return + } + var key FragKey + var entry *FragEntry + first := true + for k, e := range m.m { + if first || e.deadline.Before(entry.deadline) { + key = k + entry = e + first = false + } + } + m.removeEntey(key, entry) +} + +func (m *FragManager) gc() { + ticker := time.NewTicker(fragTTL / 2) + defer ticker.Stop() + for { + select { + case <-m.ch: + return + case now := <-ticker.C: + m.mu.Lock() + for k, e := range m.m { + if now.After(e.deadline) { + m.removeEntey(k, e) + } + } + m.mu.Unlock() + } + } +} + +func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []byte) int { + m.mu.Lock() + defer m.mu.Unlock() + if m.closed() { + return 0 + } + + if fragN < 2 { + return 0 + } + + now := time.Now() + entry := m.m[key] + if entry == nil || now.After(entry.deadline) { + if entry == nil { + m.tryRemove() + } else { + m.removeEntey(key, entry) + } + entry = &FragEntry{ + data: make([][]byte, fragN), + total: fragN, + deadline: now.Add(fragTTL), + } + m.m[key] = entry + } + + if fragN != entry.total { + return 0 + } + if fragIdx >= entry.total { + return 0 + } + if entry.data[fragIdx] != nil { + return 0 + } + if entry.size+len(data) > fragSize { + return 0 + } + if entry.len < int(entry.total)-1 { + if m.sizem[key.clientID]+len(data) > fragClientIDSize { + return 0 + } + } + + cp := make([]byte, len(data)) + copy(cp, data) + + entry.data[fragIdx] = cp + entry.size += len(data) + entry.len++ + entry.deadline = now.Add(fragTTL) + m.sizem[key.clientID] += len(data) + + if entry.len < int(entry.total) { + return 0 + } + + out = out[:0] + for i := range entry.data { + out = append(out, entry.data[i]...) + } + m.removeEntey(key, entry) + return len(out) +} + +func (m *FragManager) Close() { + m.mu.Lock() + defer m.mu.Unlock() + if m.closed() { + return + } + close(m.ch) + for k := range m.m { + delete(m.m, k) + } +} diff --git a/transport/internet/finalmask/xdns/record_transport.go b/transport/internet/finalmask/xdns/record_transport.go deleted file mode 100644 index 8428baa40..000000000 --- a/transport/internet/finalmask/xdns/record_transport.go +++ /dev/null @@ -1,226 +0,0 @@ -package xdns - -import "bytes" - -const ipRecordHeaderSize = 2 - -func maxEncodedPayloadForType(rrType uint16) int { - switch rrType { - case RRTypeA: - return maxEncodedPayloadA - case RRTypeAAAA: - return maxEncodedPayloadAAAA - default: - return maxEncodedPayloadTXT - } -} - -func rrDataSizeForType(rrType uint16) int { - switch rrType { - case RRTypeA: - return 4 - case RRTypeAAAA: - return 16 - default: - return 0 - } -} - -func payloadChunkSizeForType(rrType uint16) int { - size := rrDataSizeForType(rrType) - if size <= ipRecordHeaderSize { - return 0 - } - return size - ipRecordHeaderSize -} - -func answersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) { - switch question.Type { - case RRTypeTXT: - return []RR{ - { - Name: question.Name, - Type: question.Type, - Class: question.Class, - TTL: ttl, - Data: EncodeRDataTXT(payload), - }, - }, nil - case RRTypeA, RRTypeAAAA: - return ipAnswersForPayload(question, ttl, payload) - default: - return nil, ErrIntegerOverflow - } -} - -func ipAnswersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) { - chunkSize := payloadChunkSizeForType(question.Type) - rrDataSize := rrDataSizeForType(question.Type) - if chunkSize == 0 || rrDataSize == 0 { - return nil, ErrIntegerOverflow - } - - numRecords := 1 - if len(payload) > 0 { - numRecords = (len(payload) + chunkSize - 1) / chunkSize - } - if numRecords > 256 { - return nil, ErrIntegerOverflow - } - - answers := make([]RR, 0, numRecords) - for i := 0; i < numRecords; i++ { - offset := i * chunkSize - n := len(payload) - offset - if n < 0 { - n = 0 - } - if n > chunkSize { - n = chunkSize - } - - data := make([]byte, rrDataSize) - data[0] = byte(i) - data[1] = byte(n) - copy(data[ipRecordHeaderSize:], payload[offset:offset+n]) - - answers = append(answers, RR{ - Name: question.Name, - Type: question.Type, - Class: question.Class, - TTL: ttl, - Data: data, - }) - } - - return answers, nil -} - -func decodeResponsePayload(answers []RR) []byte { - if len(answers) == 0 { - return nil - } - - switch answers[0].Type { - case RRTypeTXT: - if len(answers) != 1 { - return nil - } - payload, err := DecodeRDataTXT(answers[0].Data) - if err != nil { - return nil - } - return payload - case RRTypeA, RRTypeAAAA: - return decodeIPAnswerPayload(answers, answers[0].Type) - default: - return nil - } -} - -func decodeIPAnswerPayload(answers []RR, rrType uint16) []byte { - chunkSize := payloadChunkSizeForType(rrType) - rrDataSize := rrDataSizeForType(rrType) - if chunkSize == 0 || rrDataSize == 0 || len(answers) > 256 { - return nil - } - - parts := make([][]byte, len(answers)) - for _, answer := range answers { - if answer.Type != rrType || len(answer.Data) != rrDataSize { - return nil - } - idx := int(answer.Data[0]) - n := int(answer.Data[1]) - if idx >= len(answers) || n > chunkSize || parts[idx] != nil { - return nil - } - - part := make([]byte, n) - copy(part, answer.Data[ipRecordHeaderSize:ipRecordHeaderSize+n]) - parts[idx] = part - } - - var payload bytes.Buffer - for _, part := range parts { - if part == nil { - return nil - } - payload.Write(part) - } - return payload.Bytes() -} - -func computeMaxEncodedPayload(limit int) int { - return computeMaxEncodedPayloadForType(limit, RRTypeTXT) -} - -func computeMaxEncodedPayloadForType(limit int, rrType uint16) int { - maxLengthName, err := NewName([][]byte{ - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - }) - if err != nil { - panic(err) - } - { - n := 0 - for _, label := range maxLengthName { - n += len(label) + 1 - } - n += 1 - if n != 255 { - panic("computeMaxEncodedPayload n != 255") - } - } - - queryLimit := uint16(limit) - if int(queryLimit) != limit { - queryLimit = 0xffff - } - query := &Message{ - Question: []Question{ - { - Name: maxLengthName, - Type: rrType, - Class: ClassIN, - }, - }, - Additional: []RR{ - { - Name: Name{}, - Type: RRTypeOPT, - Class: queryLimit, - TTL: 0, - Data: []byte{}, - }, - }, - } - resp, _ := responseFor(query, []domainSpec{{name: Name{[]byte{}}}}) - - low := 0 - high := 32768 - if chunkSize := payloadChunkSizeForType(rrType); chunkSize > 0 { - high = 256*chunkSize + 1 - } - for low+1 < high { - mid := (low + high) / 2 - resp.Answer, err = answersForPayload(query.Question[0], responseTTL, make([]byte, mid)) - if err != nil { - panic(err) - } - buf, err := resp.WireFormat() - if err != nil { - panic(err) - } - if len(buf) <= limit { - low = mid - } else { - high = mid - } - } - - return low -} diff --git a/transport/internet/finalmask/xdns/resolver.go b/transport/internet/finalmask/xdns/resolver.go new file mode 100644 index 000000000..279af5053 --- /dev/null +++ b/transport/internet/finalmask/xdns/resolver.go @@ -0,0 +1,31 @@ +package xdns + +import ( + "errors" + "net" + + "github.com/xtls/xray-core/common/serial" + "github.com/xtls/xray-core/transport/internet/finalmask" +) + +type Resolver interface { + Addr() *net.UDPAddr + Read(p []byte) (int, error) + Send(p []byte) + Close() +} + +func NewResolver(proto *serial.TypedMessage, dialer *finalmask.Dialer) (Resolver, error) { + config, err := proto.GetInstance() + if err != nil { + return nil, err + } + switch v := config.(type) { + case *TCPResolverProto: + return NewTCPResolver(v, dialer) + case *UDPResolverProto: + return NewUDPResolver(v, dialer) + default: + return nil, errors.New("unknown proto") + } +} diff --git a/transport/internet/finalmask/xdns/resolver_tcp.go b/transport/internet/finalmask/xdns/resolver_tcp.go new file mode 100644 index 000000000..89f2b0ea6 --- /dev/null +++ b/transport/internet/finalmask/xdns/resolver_tcp.go @@ -0,0 +1,143 @@ +package xdns + +import ( + "encoding/binary" + "errors" + "io" + "sync" + + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" +) + +type TCPResolver struct { + dest net.Destination + dialer *finalmask.Dialer + + conn net.Conn + tcpAddr *net.TCPAddr + udpAddr *net.UDPAddr + + readCh chan []byte + closeCh chan struct{} + wg sync.WaitGroup + mu sync.Mutex +} + +func NewTCPResolver(config *TCPResolverProto, dialer *finalmask.Dialer) (Resolver, error) { + dest, err := net.ParseDestination("tcp:" + config.Addr) + if err != nil { + return nil, err + } + r := &TCPResolver{ + dest: dest, + dialer: dialer, + readCh: make(chan []byte), + closeCh: make(chan struct{}), + } + if err := r.dial(); err != nil { + r.Close() + return nil, err + } + return r, nil +} + +func (r *TCPResolver) closed() bool { + select { + case <-r.closeCh: + return true + default: + return false + } +} + +func (r *TCPResolver) dial() error { + if r.closed() { + return errors.New("closed") + } + if r.conn != nil { + return nil + } + conn, err := r.dialer.DialTCP(r.dest) + if err != nil { + return err + } + r.conn = conn + r.tcpAddr = conn.RemoteAddr().(*net.TCPAddr) + r.udpAddr = &net.UDPAddr{IP: r.tcpAddr.IP, Port: r.tcpAddr.Port} + r.wg.Add(1) + go r.recv(conn) + return nil +} + +func (r *TCPResolver) recv(conn net.Conn) { + defer r.wg.Done() + + var buf [4096]byte + for { + _, err := io.ReadFull(conn, buf[:2]) + if err != nil { + break + } + n := binary.BigEndian.Uint16(buf[:2]) + if n == 0 || n > 4096 { + io.CopyN(io.Discard, conn, int64(n)) + continue + } + _, err = io.ReadFull(conn, buf[:n]) + if err != nil { + break + } + p := pool4K.Get().([]byte) + copy(p, buf[:n]) + select { + case <-r.closeCh: + pool4K.Put(p[:cap(p)]) + case r.readCh <- p[:n]: + } + } + + r.mu.Lock() + defer r.mu.Unlock() + + _ = conn.Close() + r.conn = nil +} + +func (r *TCPResolver) Addr() *net.UDPAddr { + return r.udpAddr +} + +func (r *TCPResolver) Read(p []byte) (n int, err error) { + packet, ok := <-r.readCh + if ok { + n = copy(p, packet) + pool4K.Put(packet[:cap(packet)]) + return n, nil + } + return 0, io.ErrClosedPipe +} + +func (r *TCPResolver) Send(p []byte) { + r.mu.Lock() + defer r.mu.Unlock() + if r.dial() != nil { + return + } + _ = binary.Write(r.conn, binary.BigEndian, len(p)) + _, _ = r.conn.Write(p) +} + +func (r *TCPResolver) Close() { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed() { + return + } + close(r.closeCh) + if r.conn != nil { + _ = r.conn.Close() + } + r.wg.Wait() + close(r.readCh) +} diff --git a/transport/internet/finalmask/xdns/resolver_udp.go b/transport/internet/finalmask/xdns/resolver_udp.go new file mode 100644 index 000000000..92bd67423 --- /dev/null +++ b/transport/internet/finalmask/xdns/resolver_udp.go @@ -0,0 +1,130 @@ +package xdns + +import ( + "errors" + "io" + "sync" + + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" +) + +type UDPResolver struct { + dest net.Destination + dialer *finalmask.Dialer + + conn net.PacketConn + udpAddr *net.UDPAddr + + readCh chan []byte + closeCh chan struct{} + wg sync.WaitGroup + mu sync.Mutex +} + +func NewUDPResolver(config *UDPResolverProto, dialer *finalmask.Dialer) (Resolver, error) { + dest, err := net.ParseDestination("udp:" + config.Addr) + if err != nil { + return nil, err + } + r := &UDPResolver{ + dest: dest, + dialer: dialer, + readCh: make(chan []byte), + closeCh: make(chan struct{}), + } + if err := r.dial(); err != nil { + r.Close() + return nil, err + } + return r, nil +} + +func (r *UDPResolver) closed() bool { + select { + case <-r.closeCh: + return true + default: + return false + } +} + +func (r *UDPResolver) dial() error { + if r.closed() { + return errors.New("closed") + } + if r.conn != nil { + return nil + } + conn, err := r.dialer.DialUDP(r.dest) + if err != nil { + return err + } + r.conn = conn.(*finalmask.PacketConnWrapper).PacketConn + r.udpAddr = conn.RemoteAddr().(*net.UDPAddr) + r.wg.Add(1) + go r.recv(conn.(*finalmask.PacketConnWrapper).PacketConn) + return nil +} + +func (r *UDPResolver) recv(conn net.PacketConn) { + defer r.wg.Done() + + var buf [4096]byte + for { + n, _, err := conn.ReadFrom(buf[:]) + if err != nil { + break + } + p := pool4K.Get().([]byte) + copy(p, buf[:n]) + select { + case <-r.closeCh: + pool4K.Put(p[:cap(p)]) + case r.readCh <- p[:n]: + } + } + + r.mu.Lock() + defer r.mu.Unlock() + + _ = conn.Close() + r.conn = nil +} + +func (r *UDPResolver) Addr() *net.UDPAddr { + return r.udpAddr +} + +func (r *UDPResolver) Read(p []byte) (n int, err error) { + packet, ok := <-r.readCh + if ok { + n = copy(p, packet) + pool4K.Put(packet[:cap(packet)]) + return n, nil + } + return 0, io.ErrClosedPipe +} + +func (r *UDPResolver) Send(p []byte) { + r.mu.Lock() + defer r.mu.Unlock() + if err := r.dial(); err != nil { + return + } + _, _ = r.conn.WriteTo(p, r.udpAddr) +} + +func (r *UDPResolver) Close() { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed() { + return + } + close(r.closeCh) + if r.conn != nil { + _ = r.conn.Close() + } + r.wg.Wait() + close(r.readCh) +} diff --git a/transport/internet/finalmask/xdns/resp.go b/transport/internet/finalmask/xdns/resp.go new file mode 100644 index 000000000..93ff8a8c3 --- /dev/null +++ b/transport/internet/finalmask/xdns/resp.go @@ -0,0 +1,392 @@ +package xdns + +import ( + "sort" + "sync" + "time" + + "github.com/xtls/xray-core/common" + "golang.org/x/net/dns/dnsmessage" +) + +const ( + sendTTL = 4 * time.Second +) + +type Resp struct { + msg dnsmessage.Message + domain *Domain + edns0 uint16 + + cap int +} + +func NewResp(msg dnsmessage.Message, domain *Domain, edns0 uint16) *Resp { + if msg.Header.Response { + return &Resp{ + msg: msg, + domain: domain, + } + } + + size := min(max(int(edns0), 512), max(int(domain.edns0), 512)) + + left := size - 12 - int(msg.Questions[0].Name.Length) - 1 - 2 - 2 + if edns0 > 0 { + left -= 1 + 2 + 2 + 4 + 2 + 0 + } + cap := 0 + switch msg.Questions[0].Type { + case dnsmessage.TypeA: + single := 2 + 2 + 2 + 4 + 2 + 4 + n := left / single + if n > 255 { + n = 255 + } + cap = 4*n - n - 1 + case dnsmessage.TypeCNAME: + single := 2 + 2 + 2 + 4 + 2 + domain.lenMax + n := left / single + if n > 255 { + n = 255 + } + cap = domain.cap*n - n - 1 + case dnsmessage.TypeTXT: + left -= 2 + 2 + 2 + 4 + 2 + single := 255 + n := left / single + m := left % single + cap = 255*n - n + if m > 1 { + cap += m - 1 + } + case dnsmessage.TypeAAAA: + single := 2 + 2 + 2 + 4 + 2 + 16 + n := left / single + if n > 255 { + n = 255 + } + cap = 16*n - n - 1 + } + + return &Resp{ + msg: msg, + domain: domain, + edns0: edns0, + + cap: cap, + } +} + +func (r *Resp) Encode(encoded []byte, data []byte) []byte { + msg := r.msg + msg.Header = dnsmessage.Header{ + ID: msg.Header.ID, + Response: true, + Authoritative: true, + RCode: dnsmessage.RCodeSuccess, + } + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + switch msg.Questions[0].Type { + case dnsmessage.TypeA: + fragN := 0 + if len(data) > 0 { + fragN = 1 + } + if (len(data) - (4 - 2)) > 0 { + fragN += (len(data) - (4 - 2)) / (4 - 1) + if (len(data)-(4-2))%(4-1) > 0 { + fragN++ + } + } + + for i := range fragN { + A := [4]byte{byte(i)} + if i == 0 { + A[1] = byte(fragN) + n := copy(A[2:], data) + data = data[n:] + } else { + n := copy(A[1:], data) + data = data[n:] + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AResource{A: A}, + }) + } + case dnsmessage.TypeCNAME: + fragN := 0 + if len(data) > 0 { + fragN = 1 + } + if (len(data) - (r.domain.cap - 2)) > 0 { + fragN += (len(data) - (r.domain.cap - 2)) / (r.domain.cap - 1) + if (len(data)-(r.domain.cap-2))%(r.domain.cap-1) > 0 { + fragN++ + } + } + + DATA := make([]byte, r.domain.cap) + for i := range fragN { + DATA[0] = byte(i) + if i == 0 { + DATA[1] = byte(fragN) + n := copy(DATA[2:], data) + data = data[n:] + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:2+n])}, + }) + } else { + n := copy(DATA[1:], data) + data = data[n:] + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:1+n])}, + }) + } + } + case dnsmessage.TypeTXT: + var txt []string + for len(data) > 0 { + size := min(len(data), 255) + txt = append(txt, string(data[:size])) + data = data[size:] + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.TXTResource{TXT: txt}, + }) + case dnsmessage.TypeAAAA: + fragN := 0 + if len(data) > 0 { + fragN = 1 + } + if (len(data) - (16 - 2)) > 0 { + fragN += (len(data) - (16 - 2)) / (16 - 1) + if (len(data)-(16-2))%(16-1) > 0 { + fragN++ + } + } + + for i := range fragN { + AAAA := [16]byte{byte(i)} + if i == 0 { + AAAA[1] = byte(fragN) + n := copy(AAAA[2:], data) + data = data[n:] + } else { + n := copy(AAAA[1:], data) + data = data[n:] + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AAAAResource{AAAA: AAAA}, + }) + } + } + if r.edns0 > 0 { + msg.Additionals = append(msg.Additionals, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeOPT, + Class: dnsmessage.Class(r.edns0), + TTL: 0, + }, + Body: &dnsmessage.OPTResource{}, + }) + } + return common.Must2(msg.AppendPack(encoded[:0])) +} + +func (r *Resp) Decode(decoded []byte) int { + decoded = decoded[:0] + msg := r.msg + if msg.Questions[0].Type == dnsmessage.TypeTXT { + if len(msg.Answers) == 1 && r.domain.IsDomain(msg.Answers[0].Header.Name) && msg.Answers[0].Header.Type == dnsmessage.TypeTXT { + for i := range msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT { + decoded = append(decoded, msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i]...) + } + } + return len(decoded) + } else { + var frags [][]byte + for i := range msg.Answers { + if !r.domain.IsDomain(msg.Answers[i].Header.Name) || msg.Answers[i].Header.Type != msg.Questions[0].Type { + continue + } + switch msg.Questions[0].Type { + case dnsmessage.TypeA: + frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AResource).A[:]) + case dnsmessage.TypeCNAME: + var decoded [255]byte + n := r.domain.Decode(&decoded, msg.Answers[i].Body.(*dnsmessage.CNAMEResource).CNAME) + if n == 0 { + continue + } + frags = append(frags, decoded[:n]) + case dnsmessage.TypeAAAA: + frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AAAAResource).AAAA[:]) + } + } + sort.Slice(frags, func(i, j int) bool { + return frags[i][0] < frags[j][0] + }) + if len(frags) < 1 || len(frags[0]) < 2 || int(frags[0][1]) > len(frags) { + return 0 + } + decoded = append(decoded, frags[0][2:]...) + for i := range frags { + if i > 0 { + if frags[i][0] == frags[i-1][0] { + return 0 + } + decoded = append(decoded, frags[i][1:]...) + } + } + return len(decoded) + } +} + +type SendInfo struct { + stash chan []byte + ch chan []byte + deadline time.Time +} + +type SendManager struct { + m map[ClientID]*SendInfo + ch chan struct{} + mu sync.Mutex +} + +func NewSendManager() *SendManager { + m := &SendManager{ + m: make(map[ClientID]*SendInfo), + ch: make(chan struct{}), + } + go m.gc() + return m +} + +func (m *SendManager) closed() bool { + select { + case <-m.ch: + return true + default: + return false + } +} + +func (m *SendManager) gc() { + ticker := time.NewTicker(sendTTL) + defer ticker.Stop() + for { + select { + case <-m.ch: + return + case now := <-ticker.C: + m.mu.Lock() + for key, info := range m.m { + if now.After(info.deadline) { + close(info.stash) + close(info.ch) + delete(m.m, key) + } + } + m.mu.Unlock() + ticker.Reset(sendTTL) + } + } +} + +func (m *SendManager) Push(clientID ClientID, p []byte) { + m.mu.Lock() + defer m.mu.Unlock() + info := m.m[clientID] + if info == nil { + info = &SendInfo{ + stash: make(chan []byte, 1), + ch: make(chan []byte, 128), + deadline: time.Now().Add(sendTTL), + } + m.m[clientID] = info + } + b := make([]byte, len(p)) + copy(b, p) + select { + case info.ch <- b: + default: + } +} + +func (m *SendManager) Stash(clientID ClientID, p []byte) { + m.mu.Lock() + defer m.mu.Unlock() + info := m.m[clientID] + if info == nil { + return + } + info.deadline = time.Now().Add(sendTTL) + select { + case info.stash <- p: + default: + } +} + +func (m *SendManager) Pop(clientID ClientID) (chan []byte, chan []byte) { + m.mu.Lock() + defer m.mu.Unlock() + info := m.m[clientID] + if info == nil { + info = &SendInfo{ + stash: make(chan []byte, 1), + ch: make(chan []byte, 128), + } + m.m[clientID] = info + } + info.deadline = time.Now().Add(sendTTL) + return info.ch, info.stash +} + +func (m *SendManager) Close() { + m.mu.Lock() + defer m.mu.Unlock() + if m.closed() { + return + } + close(m.ch) + for key, info := range m.m { + close(info.stash) + close(info.ch) + delete(m.m, key) + } +} diff --git a/transport/internet/finalmask/xdns/server.go b/transport/internet/finalmask/xdns/server.go index 654f7fdba..c7506f5f3 100644 --- a/transport/internet/finalmask/xdns/server.go +++ b/transport/internet/finalmask/xdns/server.go @@ -1,512 +1,385 @@ package xdns import ( - "bytes" "context" - "encoding/binary" - go_errors "errors" "io" - "net" "sync" "time" + "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/errors" - "github.com/xtls/xray-core/transport/internet/finalmask" + "github.com/xtls/xray-core/common/net" + "golang.org/x/net/dns/dnsmessage" ) const ( - idleTimeout = 10 * time.Second - responseTTL = 60 - maxResponseDelay = 1 * time.Second + maxResponseDelay = time.Second ) -var ( - maxUDPPayload = 1280 - 40 - 8 - maxEncodedPayloadTXT = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeTXT) - maxEncodedPayloadA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeA) - maxEncodedPayloadAAAA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeAAAA) -) - -func clientIDToAddr(clientID [8]byte) *net.UDPAddr { - ip := make(net.IP, 16) - - copy(ip, []byte{0xfd, 0x00, 0, 0, 0, 0, 0, 0}) - copy(ip[8:], clientID[:]) - - return &net.UDPAddr{ - IP: ip, - } +type resp struct { + msg dnsmessage.Message + addr net.Addr } -type record struct { - Resp *Message - Addr net.Addr - // ClientID [8]byte - ClientAddr net.Addr +type Rec struct { + resp *Resp + clientID ClientID + addr net.Addr } -type queue struct { - last time.Time - rrType uint16 - queue chan []byte - stash chan []byte -} - -type xdnsConnServer struct { +type xdnsServer struct { net.PacketConn - domains []domainSpec + domains []*Domain + fragManager *FragManager + sendManager *SendManager - ch chan *record - readQueue chan *packet - writeQueueMap map[string]*queue - - closed bool - mutex sync.Mutex + readCh chan packet + recCh chan *Rec + drCh chan resp + closeCh chan struct{} + wg sync.WaitGroup + mu sync.RWMutex } -func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { if len(c.Domains) == 0 { return nil, errors.New("empty domains") } - domains := make([]domainSpec, 0, len(c.Domains)) - for _, domain := range c.Domains { - domain, err := parseDomainSpec(domain, "") + domains := make([]*Domain, 0, len(c.Domains)) + for i := range c.Domains { + types := make([]uint16, 0, len(c.Domains[i].Types)) + for j := range c.Domains[i].Types { + types = append(types, uint16(c.Domains[i].Types[j])) + } + domain, err := NewDomain(c.Domains[i].Name, int(c.Domains[i].LenLimit), int(c.Domains[i].LabelLimit), types, uint16(c.Domains[i].Edns0)) if err != nil { return nil, err } domains = append(domains, domain) } - - conn := &xdnsConnServer{ + server := &xdnsServer{ PacketConn: raw, - domains: domains, + domains: domains, + fragManager: NewFragManager(), + sendManager: NewSendManager(), - ch: make(chan *record, 500), - readQueue: make(chan *packet, 512), - writeQueueMap: make(map[string]*queue), + readCh: make(chan packet), + recCh: make(chan *Rec, 255), + drCh: make(chan resp), + closeCh: make(chan struct{}), } - - go conn.clean() - go conn.recvLoop() - go conn.sendLoop() - - return conn, nil + go server.run() + return server, nil } -func (c *xdnsConnServer) clean() { - f := func() bool { - c.mutex.Lock() - defer c.mutex.Unlock() - - if c.closed { - return true - } - - now := time.Now() - - for key, q := range c.writeQueueMap { - if now.Sub(q.last) >= idleTimeout { - close(q.queue) - close(q.stash) - delete(c.writeQueueMap, key) - } - } - +func (c *xdnsServer) closed() bool { + select { + case <-c.closeCh: + return true + default: return false } - - for { - time.Sleep(idleTimeout / 2) - if f() { - return - } - } } -func (c *xdnsConnServer) ensureQueue(addr net.Addr) *queue { - if c.closed { - return nil - } - - q, ok := c.writeQueueMap[addr.String()] - if !ok { - q = &queue{ - queue: make(chan []byte, 512), - stash: make(chan []byte, 1), - } - c.writeQueueMap[addr.String()] = q - } - q.last = time.Now() - - return q -} - -func (c *xdnsConnServer) stash(queue *queue, p []byte) { - c.mutex.Lock() - defer c.mutex.Unlock() - - if c.closed { - return - } - +func (c *xdnsServer) decref(msg dnsmessage.Message, addr net.Addr) { select { - case queue.stash <- p: + case c.drCh <- resp{msg: msg, addr: addr}: default: } } -func (c *xdnsConnServer) recvLoop() { - var buf [finalmask.UDPSize]byte +func (c *xdnsServer) read(buf []byte, addr net.Addr) { + msg := dnsmessage.Message{} + if err := msg.Unpack(buf); err != nil { + return + } + if msg.Header.Response { + return + } - for { - if c.closed { - break - } + if msg.Header.OpCode != 0 { + msg.Header.Response = true + msg.Header.RCode = dnsmessage.RCodeNotImplemented + c.decref(msg, addr) + return + } - n, addr, err := c.PacketConn.ReadFrom(buf[:]) - if err != nil { - if go_errors.Is(err, net.ErrClosed) { - break + if len(msg.Questions) != 1 { + msg.Header.Response = true + msg.Header.RCode = dnsmessage.RCodeFormatError + c.decref(msg, addr) + return + } + + opt := false + edns0 := uint16(0) + for i := range msg.Additionals { + if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT { + if opt { + msg.Header.RCode = dnsmessage.RCodeFormatError + c.decref(msg, addr) + return } - continue - } - - query, err := MessageFromWireFormat(buf[:n]) - if err != nil { - errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err) - continue - } - - resp, payload := responseFor(&query, c.domains) - - var clientID [8]byte - n = copy(clientID[:], payload) - payload = payload[n:] - if n == len(clientID) { - r := bytes.NewReader(payload) - for { - p, err := nextPacketServer(r) - if err != nil { - break - } - - buf := make([]byte, len(p)) - copy(buf, p) - select { - case c.readQueue <- &packet{ - p: buf, - addr: clientIDToAddr(clientID), - }: - default: - errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err queue full") - } - } - } else { - if resp != nil && resp.Rcode() == RcodeNoError { - resp.Flags |= RcodeNameError - } - } - - if resp != nil { - select { - case c.ch <- &record{resp, addr, clientIDToAddr(clientID)}: - default: - errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err record queue full") + opt = true + edns0 = uint16(msg.Additionals[i].Header.Class) + if ver := (msg.Additionals[i].Header.TTL >> 16) & 0xFF; ver != 0 { + msg.Header.RCode = dnsmessage.RCodeSuccess + msg.Additionals[i].Header.TTL = 1 << 24 + c.decref(msg, addr) + return } } } + if opt { + if edns0 < 512 { + edns0 = 512 + } + if edns0 > 4096 { + edns0 = 4096 + } + } + errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type) - errors.LogDebug(context.Background(), "xdns closed") + var domain *Domain + for i := range c.domains { + if c.domains[i].IsDomain(msg.Questions[0].Name) { + domain = c.domains[i] + break + } + } + if domain == nil { + msg.Header.Response = true + msg.Header.RCode = dnsmessage.RCodeNameError + c.decref(msg, addr) + return + } + if !domain.HasType(uint16(msg.Questions[0].Type)) { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } - close(c.ch) - close(c.readQueue) + var decoded [255]byte + n := domain.Decode(&decoded, msg.Questions[0].Name) + if n < 9 { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } + if TypeMap_[decoded[0]&3] != uint16(msg.Questions[0].Type) || (decoded[8]&0x3F != 3 && decoded[8]&0x3F != 8) || (decoded[8]&0x3F == 3 && n < 9+3+1) || (decoded[8]&0x3F == 8 && n != 9+8) { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } + clientID := ClientIDFromRaw([8]byte(decoded[:8])) - c.mutex.Lock() - defer c.mutex.Unlock() + r := NewResp(msg, domain, edns0) + if r == nil { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } + select { + case c.recCh <- &Rec{resp: r, clientID: clientID, addr: addr}: + default: + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + } - c.closed = true - for key, q := range c.writeQueueMap { - close(q.queue) - close(q.stash) - delete(c.writeQueueMap, key) + if decoded[8]&0x3F == 8 { + return + } + p := pool4K.Get().([]byte) + p = p[:0] + if decoded[8]&0xC0 == 0xC0 { + out := pool4K.Get().([]byte) + n := c.fragManager.Feed(out, FragKey{clientID: clientID, fragID: decoded[12]}, decoded[13], decoded[14], decoded[15:n]) + pool4K.Put(p[:cap(p)]) + if n > 0 { + p = out[:n] + } else { + pool4K.Put(out[:cap(out)]) + return + } + } else { + p = append(p, decoded[12:n]...) + } + select { + case <-c.closeCh: + pool4K.Put(p[:cap(p)]) + return + case c.readCh <- packet{p: p, addr: clientID.Addr()}: + return } } -func (c *xdnsConnServer) sendLoop() { - var nextRec *record +func (c *xdnsServer) run() { + c.wg.Add(1) + go c.recv() + + c.wg.Add(1) + go c.send() + + c.wg.Add(1) + go c.dr() + + c.wg.Wait() + close(c.readCh) + close(c.recCh) + close(c.drCh) + c.fragManager.Close() + c.sendManager.Close() +} + +func (c *xdnsServer) recv() { + defer c.wg.Done() + + var buf [512]byte + for { + n, addr, err := c.PacketConn.ReadFrom(buf[:]) + if err != nil { + if c.closed() { + return + } + errors.LogErrorInner(context.Background(), err, "recv err") + return + } + c.read(buf[:n], addr) + } +} + +func (c *xdnsServer) send() { + defer c.wg.Done() + + timer := time.NewTimer(maxResponseDelay) + timer.Stop() + var buf [4096]byte + var data [4096]byte + var nextRec *Rec for { - var err error rec := nextRec nextRec = nil if rec == nil { - var ok bool - rec, ok = <-c.ch - if !ok { - break + select { + case rec = <-c.recCh: + case <-c.closeCh: + return } } - if rec.Resp.Rcode() == RcodeNoError && len(rec.Resp.Question) == 1 { - var payload bytes.Buffer - limit := maxEncodedPayloadForType(rec.Resp.Question[0].Type) - timer := time.NewTimer(maxResponseDelay) - - for { - c.mutex.Lock() - q := c.ensureQueue(rec.ClientAddr) - if q == nil { - c.mutex.Unlock() - return - } - q.rrType = rec.Resp.Question[0].Type - c.mutex.Unlock() - - var p []byte - + ch, stash := c.sendManager.Pop(rec.clientID) + left := rec.resp.cap + timer.Reset(maxResponseDelay) + var ps [][]byte + for { + var p []byte + select { + case p = <-stash: + default: select { - case p = <-q.stash: + case p = <-stash: + case p = <-ch: default: select { - case p = <-q.stash: - case p = <-q.queue: - default: - select { - case p = <-q.stash: - case p = <-q.queue: - case <-timer.C: - case nextRec = <-c.ch: - } + case p = <-stash: + case p = <-ch: + case <-timer.C: + case nextRec = <-c.recCh: } } - - timer.Reset(0) - - if len(p) == 0 { + } + if len(p) == 0 { + break + } + timer.Reset(0) + left -= 2 + len(p) + if left < 0 { + if len(ps) == 0 { + errors.LogError(context.Background(), "err size ", len(p)) break } - - limit -= 2 + len(p) - if limit < 0 { - if payload.Len() == 0 { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns payload too large for rrtype ", rec.Resp.Question[0].Type, " ", len(p)) - continue - } - c.stash(q, p) - break - } - - // if len(p) > 65535 { - // panic(len(p)) - // } - - _ = binary.Write(&payload, binary.BigEndian, uint16(len(p))) - payload.Write(p) + c.sendManager.Stash(rec.clientID, p) + break } + ps = append(ps, p) + } + timer.Stop() - timer.Stop() - rec.Resp.Answer, err = answersForPayload(rec.Resp.Question[0], responseTTL, payload.Bytes()) - if err != nil { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns encode err ", err) - continue + d := data[:0] + for i := range ps { + l := len(ps[i]) + if i == len(ps)-1 { + l |= 0xC000 } + d = append(d, []byte{byte(l >> 8), byte(l)}...) + d = append(d, ps[i]...) } + _, _ = c.PacketConn.WriteTo(rec.resp.Encode(buf[:0], d), rec.addr) + } +} - buf, err := rec.Resp.WireFormat() - if err != nil { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns wireformat err ", err) - continue - } +func (c *xdnsServer) dr() { + defer c.wg.Done() - if len(buf) > maxUDPPayload { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns truncate ", len(buf)) - buf = buf[:maxUDPPayload] - buf[2] |= 0x02 - } - - if c.closed { + var buf [512]byte + for { + select { + case <-c.closeCh: return - } - - _, err = c.PacketConn.WriteTo(buf, rec.Addr) - if go_errors.Is(err, net.ErrClosed) { - c.closed = true - break + case r := <-c.drCh: + _, _ = c.PacketConn.WriteTo(common.Must2(r.msg.AppendPack(buf[:0])), r.addr) } } } -func (c *xdnsConnServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) { - packet, ok := <-c.readQueue - if !ok { - return 0, nil, net.ErrClosed +func (c *xdnsServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + packet, ok := <-c.readCh + if ok { + n = copy(p, packet.p) + pool4K.Put(packet.p[:cap(packet.p)]) + return n, packet.addr, nil } - if len(p) < len(packet.p) { - errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p)) - return 0, packet.addr, nil - } - copy(p, packet.p) - return len(packet.p), packet.addr, nil + return 0, nil, io.ErrClosedPipe } -func (c *xdnsConnServer) WriteTo(p []byte, addr net.Addr) (n int, err error) { - c.mutex.Lock() - defer c.mutex.Unlock() - - q := c.ensureQueue(addr) - if q == nil { +func (c *xdnsServer) WriteTo(p []byte, addr net.Addr) (n int, err error) { + if c.closed() { return 0, io.ErrClosedPipe } - limit := maxEncodedPayloadForType(q.rrType) - if q.rrType == 0 { - limit = maxEncodedPayloadTXT - } - if len(p)+2 > limit { - errors.LogDebug(context.Background(), addr, " mask write err short write ", len(p), "+2 > ", limit) - return 0, nil - } - - buf := make([]byte, len(p)) - copy(buf, p) - - select { - case q.queue <- buf: - return len(p), nil - default: - // errors.LogDebug(context.Background(), addr, " mask write err queue full") - return 0, nil + if len(p) == 0 || len(p) > 4096 { + errors.LogError(context.Background(), "err size ", len(p)) + return 0, errors.New("err size") } + c.sendManager.Push(ClientIDFromAddr(addr.(*net.UDPAddr)), p) + return len(p), nil } -func (c *xdnsConnServer) Close() error { - c.closed = true - return c.PacketConn.Close() +func (c *xdnsServer) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed() { + return nil + } + close(c.closeCh) + _ = c.PacketConn.Close() + return nil } -func nextPacketServer(r *bytes.Reader) ([]byte, error) { - eof := func(err error) error { - if err == io.EOF { - err = io.ErrUnexpectedEOF - } - return err - } +func (c *xdnsServer) SetDeadline(t time.Time) error { return errors.New("not support") } - for { - prefix, err := r.ReadByte() - if err != nil { - return nil, err - } - if prefix >= 224 { - paddingLen := prefix - 224 - _, err := io.CopyN(io.Discard, r, int64(paddingLen)) - if err != nil { - return nil, eof(err) - } - } else { - p := make([]byte, int(prefix)) - _, err = io.ReadFull(r, p) - return p, eof(err) - } - } -} +func (c *xdnsServer) SetReadDeadline(t time.Time) error { return errors.New("not support") } -func responseFor(query *Message, domains []domainSpec) (*Message, []byte) { - resp := &Message{ - ID: query.ID, - Flags: 0x8000, - Question: query.Question, - } - - if query.Flags&0x8000 != 0 { - return nil, nil - } - - payloadSize := 0 - for _, rr := range query.Additional { - if rr.Type != RRTypeOPT { - continue - } - if len(resp.Additional) != 0 { - resp.Flags |= RcodeFormatError - return resp, nil - } - resp.Additional = append(resp.Additional, RR{ - Name: Name{}, - Type: RRTypeOPT, - Class: 4096, - TTL: 0, - Data: []byte{}, - }) - additional := &resp.Additional[0] - - version := (rr.TTL >> 16) & 0xff - if version != 0 { - resp.Flags |= ExtendedRcodeBadVers & 0xf - additional.TTL = (ExtendedRcodeBadVers >> 4) << 24 - return resp, nil - } - - payloadSize = int(rr.Class) - } - if payloadSize < 512 { - payloadSize = 512 - } - - if len(query.Question) != 1 { - resp.Flags |= RcodeFormatError - return resp, nil - } - question := query.Question[0] - - var ( - prefix Name - ok bool - match domainSpec - ) - for _, domain := range domains { - prefix, ok = question.Name.TrimSuffix(domain.name) - if ok { - match = domain - break - } - } - if !ok { - resp.Flags |= RcodeNameError - return resp, nil - } - resp.Flags |= 0x0400 - - if query.Opcode() != 0 { - resp.Flags |= RcodeNotImplemented - return resp, nil - } - - switch question.Type { - case RRTypeTXT, RRTypeA, RRTypeAAAA: - default: - resp.Flags |= RcodeNameError - return resp, nil - } - if match.rrType != 0 && question.Type != match.rrType { - resp.Flags |= RcodeNameError - return resp, nil - } - - encoded := bytes.ToUpper(bytes.Join(prefix, nil)) - payload := make([]byte, base32Encoding.DecodedLen(len(encoded))) - n, err := base32Encoding.Decode(payload, encoded) - if err != nil { - resp.Flags |= RcodeNameError - return resp, nil - } - payload = payload[:n] - - if payloadSize < maxUDPPayload { - resp.Flags |= RcodeFormatError - return resp, nil - } - - return resp, payload -} +func (c *xdnsServer) SetWriteDeadline(t time.Time) error { return errors.New("not support") } diff --git a/transport/internet/finalmask/xdns/spec.go b/transport/internet/finalmask/xdns/spec.go deleted file mode 100644 index 28461569f..000000000 --- a/transport/internet/finalmask/xdns/spec.go +++ /dev/null @@ -1,80 +0,0 @@ -package xdns - -import ( - "strings" - - "github.com/xtls/xray-core/common/errors" -) - -type domainSpec struct { - name Name - rrType uint16 -} - -func rrTypeFromMethod(method string) (uint16, error) { - switch strings.ToLower(method) { - case "", "txt": - return RRTypeTXT, nil - case "a": - return RRTypeA, nil - case "aaaa": - return RRTypeAAAA, nil - default: - return 0, errors.New("unsupported method") - } -} - -func parseDomainSpec(s string, defaultMethod string) (domainSpec, error) { - domainPart := s - method := "" - hasMethod := false - - if i := strings.LastIndex(s, ":"); i >= 0 { - domainPart = s[:i] - method = s[i+1:] - hasMethod = true - } else if defaultMethod != "" { - method = defaultMethod - hasMethod = true - } - - if domainPart == "" { - return domainSpec{}, errors.New("empty domain") - } - - name, err := ParseName(domainPart) - if err != nil { - return domainSpec{}, err - } - - rrType := uint16(0) - if hasMethod { - var err error - rrType, err = rrTypeFromMethod(method) - if err != nil { - return domainSpec{}, err - } - } - - return domainSpec{ - name: name, - rrType: rrType, - }, nil -} - -func parseResolver(s string) (Name, string, uint16, error) { - head, server, ok := strings.Cut(s, "+udp://") - if !ok { - return nil, "", 0, errors.New("invalid resolver scheme") - } - if server == "" { - return nil, "", 0, errors.New("empty resolver server") - } - - spec, err := parseDomainSpec(head, "txt") - if err != nil { - return nil, "", 0, err - } - - return spec.name, server, spec.rrType, nil -} diff --git a/transport/internet/finalmask/xdns/xdns_test.go b/transport/internet/finalmask/xdns/xdns_test.go new file mode 100644 index 000000000..d2a4d1c1b --- /dev/null +++ b/transport/internet/finalmask/xdns/xdns_test.go @@ -0,0 +1,208 @@ +package xdns + +import ( + "bytes" + "crypto/rand" + "fmt" + mrand "math/rand" + "testing" + + "github.com/xtls/xray-core/common" + "golang.org/x/net/dns/dnsmessage" +) + +func TestXxx(t *testing.T) { + m1 := dnsmessage.Message{ + Questions: []dnsmessage.Question{ + { + Name: dnsmessage.MustNewName("a.example.com."), + }, + }, + Answers: []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeA, + Class: dnsmessage.ClassINET, + TTL: 60, + Length: 16, + }, + Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}}, + }, + }, + Additionals: []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeOPT, + Class: 255, + TTL: 0, + Length: 16, + }, + Body: &dnsmessage.OPTResource{}, + }, + }, + } + p1, e1 := m1.Pack() + if e1 != nil { + t.Fatal(e1) + } + if !bytes.Equal(p1, []byte{ + 0, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 1, + 1, 97, 7, 101, 120, 97, 109, 112, 108, 101, 3, 99, 111, 109, 0, + 0, 0, + 0, 0, + 192, 12, + 0, 1, + 0, 1, + 0, 0, 0, 60, + 0, 4, + 127, 0, 0, 1, + 0, + 0, 41, + 0, 255, + 0, 0, 0, 0, + 0, 0, + }) { + t.Fatal("!bytes.Equal") + } + + domain, _ := NewDomain("a.example.com", 200, 1, []uint16{1}, 0) + fmt.Println(domain.cap, domain.lenMax) + lenMax := domain.lenMax + data := make([]byte, domain.cap) + msg := dnsmessage.Message{} + msg.Unpack(p1) + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := mrand.Intn(255) + for range n { + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeA, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}}, + }) + } + if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+4) { + t.Fatal("fatal a") + } + } + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := mrand.Intn(255) + for range n { + common.Must2(rand.Read(data)) + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeCNAME, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.CNAMEResource{ + CNAME: domain.Encode(data), + }, + }) + } + if len(common.Must2(msg.Pack())) > 12+15+2+2+n*(2+2+2+4+2+lenMax) { + t.Fatal("fatal cname") + } + } + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := (mrand.Intn(2048) + 1024) % 2048 + a := n / 255 + b := n % 255 + c := 0 + var d [255]byte + var s []string + for range a { + s = append(s, string(d[:])) + } + if b > 0 { + c = 1 + s = append(s, string(d[:b])) + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeTXT, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.TXTResource{TXT: s}, + }) + if len(common.Must2(msg.Pack())) != 12+15+2+2+(2+2+2+4+2+n+n/255+c) { + t.Fatal("fatal txt") + } + } + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := mrand.Intn(255) + for range n { + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeAAAA, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AAAAResource{AAAA: [16]byte{}}, + }) + } + if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+16) { + t.Fatal("fatal aaaa") + } + } +} + +func TestTXT(t *testing.T) { + txt := [][]byte{{}, {}} + for i := range 255 { + txt[0] = append(txt[0], byte(i)) + } + txt[1] = []byte{255} + str := []string{} + for i := range txt { + str = append(str, string(txt[i])) + } + m1 := dnsmessage.Message{ + Answers: []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeTXT, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.TXTResource{ + TXT: str, + }, + }, + }, + } + p1 := common.Must2(m1.Pack()) + + m2 := dnsmessage.Message{} + common.Must(m2.Unpack(p1)) + if len(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT) != len(txt) { + t.Fatal("fatal txt") + } + for i := range txt { + if !bytes.Equal(txt[i], []byte(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i])) { + t.Fatal("fatal txt") + } + } +} diff --git a/transport/internet/finalmask/xicmp/client.go b/transport/internet/finalmask/xicmp/client.go index cabdb1942..702eefb9b 100644 --- a/transport/internet/finalmask/xicmp/client.go +++ b/transport/internet/finalmask/xicmp/client.go @@ -310,13 +310,6 @@ func (c *xicmpConnClient) Close() error { _ = c.icmp4.Close() _ = c.icmp6.Close() c.wg.Wait() - select { - case p := <-c.readCh: - if p.p != nil { - pool.Put(p.p) - } - default: - } close(c.readCh) return nil } diff --git a/transport/internet/finalmask/xicmp/server.go b/transport/internet/finalmask/xicmp/server.go index 117638923..1d99d03b7 100644 --- a/transport/internet/finalmask/xicmp/server.go +++ b/transport/internet/finalmask/xicmp/server.go @@ -329,13 +329,6 @@ func (c *xicmpConnServer) Close() error { _ = c.icmp4.Close() _ = c.icmp6.Close() c.wg.Wait() - select { - case p := <-c.readCh: - if p.p != nil { - pool.Put(p.p) - } - default: - } close(c.readCh) return nil } diff --git a/transport/internet/finalmask/xicmp/server_oob.go b/transport/internet/finalmask/xicmp/server_oob.go index 8d1c26db8..8eebf43bb 100644 --- a/transport/internet/finalmask/xicmp/server_oob.go +++ b/transport/internet/finalmask/xicmp/server_oob.go @@ -340,13 +340,6 @@ func (c *xicmpConnServer) Close() error { _ = c.icmp4.Close() _ = c.icmp6.Close() c.wg.Wait() - select { - case p := <-c.readCh: - if p.p != nil { - pool.Put(p.p) - } - default: - } close(c.readCh) return nil }