diff --git a/README.md b/README.md index 7259542..b6e75a4 100644 --- a/README.md +++ b/README.md @@ -84,6 +84,18 @@ Using DoT: -pubkey-file server.pub -domain t.example.com -listen 127.0.0.1:7000 ``` +Using multiple resolvers: + +```sh +./vaydns-client \ + -doh https://dns.google/dns-query \ + -doh https://cloudflare-dns.com/dns-query \ + -dot one.one.one.one:853 \ + -pubkey-file server.pub -domain t.example.com -listen 127.0.0.1:7000 +``` + +Repeat `-doh`, `-dot`, and `-udp` to enable multi-resolver mode. The client spreads queries across healthy resolvers, routes around resolvers that stop returning usable responses, and probes unhealthy resolvers for recovery. Use `-log-level debug` to see the per-resolver health table. + ### 5. Test ```sh @@ -143,13 +155,15 @@ sudo ip6tables -t nat -I PREROUTING -i eth0 -p udp --dport 53 -j REDIRECT --to-p ### Client flags -#### Transport (pick one) +#### Transport (one or more) | Flag | Description | | ----------- | ------------------------------------------------ | -| `-doh URL` | Use DNS over HTTPS with the given resolver URL | -| `-dot ADDR` | Use DNS over TLS with the given resolver address | -| `-udp ADDR` | Use plaintext UDP DNS (no covertness) | +| `-doh URL` | Use DNS over HTTPS with the given resolver URL. Repeatable | +| `-dot ADDR` | Use DNS over TLS with the given resolver address. Repeatable | +| `-udp ADDR` | Use plaintext UDP DNS (no covertness). Repeatable | + +When more than one resolver is configured, VayDNS uses multi-resolver mode with round-robin selection, health-based routing, and recovery probes. Transport types may be mixed in one client command. #### Required diff --git a/client/client.go b/client/client.go index 6bc6acc..377a4d8 100644 --- a/client/client.go +++ b/client/client.go @@ -16,11 +16,17 @@ // t.InitiateSmuxSession() // stream, _ := t.OpenStream() // returns net.Conn // defer t.Close() +// +// Multi-resolver usage (spread queries across multiple DNS resolvers): +// +// r1, _ := client.NewResolver(client.ResolverTypeUDP, "8.8.8.8:53") +// r2, _ := client.NewResolver(client.ResolverTypeDOH, "https://1.1.1.1/dns-query") +// ts, _ := client.NewTunnelServer("t.example.com", "pubkey-hex") +// t, _ := client.NewTunnelMulti([]client.Resolver{r1, r2}, ts) +// t.ListenAndServe("127.0.0.1:7000") package client import ( - "context" - "crypto/tls" "errors" "fmt" "io" @@ -190,7 +196,7 @@ func (ts *TunnelServer) effectiveMaxQnameLen() int { // either call the step-by-step Initiate* methods (for embedding in frameworks // like xray-core) or call ListenAndServe for a fully managed session. type Tunnel struct { - Resolver Resolver + Resolvers []Resolver TunnelServer TunnelServer // Session configuration. Zero values use defaults. @@ -217,11 +223,22 @@ type Tunnel struct { remoteAddr net.Addr } -// NewTunnel creates a Tunnel with the given resolver and server configuration. -// Zero-value fields use sensible defaults. +// NewTunnel creates a Tunnel with a single resolver and server configuration. +// For multiple resolvers, use NewTunnelMulti. func NewTunnel(resolver Resolver, tunnelServer TunnelServer) (*Tunnel, error) { + return NewTunnelMulti([]Resolver{resolver}, tunnelServer) +} + +// NewTunnelMulti creates a Tunnel with multiple resolvers and server +// configuration. When more than one resolver is provided, the client +// multiplexes queries across them with health-based routing. +// Zero-value fields use sensible defaults. +func NewTunnelMulti(resolvers []Resolver, tunnelServer TunnelServer) (*Tunnel, error) { + if len(resolvers) == 0 { + return nil, fmt.Errorf("at least one resolver is required") + } t := &Tunnel{ - Resolver: resolver, + Resolvers: resolvers, TunnelServer: tunnelServer, } t.wireConfig = tunnelServer.wireConfig() @@ -293,79 +310,28 @@ func (t *Tunnel) effectiveKCPWindowSize() int { // InitiateResolverConnection creates the underlying transport connection // based on the Resolver configuration. func (t *Tunnel) InitiateResolverConnection() error { - r := t.Resolver - switch r.ResolverType { - case ResolverTypeUDP: - addr, err := net.ResolveUDPAddr("udp", r.ResolverAddr) - if err != nil { - return err - } - t.remoteAddr = addr - if r.UDPSharedSocket { - lc := net.ListenConfig{Control: r.DialerControl} - conn, err := lc.ListenPacket(context.Background(), "udp", ":0") - if err != nil { - return err - } - t.resolverConn = conn - } else { - workers := r.UDPWorkers - if workers <= 0 { - workers = DefaultUDPWorkers - } - timeout := r.UDPTimeout - if timeout <= 0 { - timeout = DefaultUDPResponseTimeout - } - conn, forgedStats, err := NewUDPPacketConn(addr, r.DialerControl, workers, timeout, !r.UDPAcceptErrors, t.effectivePacketQueueSize(), t.effectiveQueueOverflowMode()) - if err != nil { - return err - } - t.forgedStats = forgedStats - t.resolverConn = conn - } - return nil - - case ResolverTypeDOH: - t.remoteAddr = turbotunnel.DummyAddr{} - var rt http.RoundTripper - if r.RoundTripper != nil { - rt = r.RoundTripper - } else if r.UTLSClientHelloID != nil { - rt = NewUTLSRoundTripper(nil, r.UTLSClientHelloID) - } else { - rt = http.DefaultTransport - } - conn, err := NewHTTPPacketConn(rt, r.ResolverAddr, 8, t.effectivePacketQueueSize(), t.effectiveQueueOverflowMode()) + if len(t.Resolvers) > 1 { + conn, err := NewMultiResolver(t.Resolvers, t.effectivePacketQueueSize(), t.effectiveQueueOverflowMode()) if err != nil { return err } t.resolverConn = conn - return nil - - case ResolverTypeDOT: t.remoteAddr = turbotunnel.DummyAddr{} - var dialTLSContext func(ctx context.Context, network, addr string) (net.Conn, error) - if r.UTLSClientHelloID != nil { - id := r.UTLSClientHelloID - dialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) { - return UTLSDialContext(ctx, network, addr, nil, id) - } - } else { - dialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) { - return tls.DialWithDialer(&net.Dialer{}, network, addr, nil) - } - } - conn, err := NewTLSPacketConn(r.ResolverAddr, dialTLSContext, t.effectivePacketQueueSize(), t.effectiveQueueOverflowMode()) - if err != nil { - return err - } - t.resolverConn = conn + // In multi-resolver mode, each entry owns a per-resolver + // ForgedStats. t.forgedStats stays nil so DNSPacketConn + // creates an unlabeled catch-all (which should rarely fire + // since MultiResolver filters forged responses upstream). return nil - - default: - return fmt.Errorf("unsupported resolver type: %s", r.ResolverType) } + r := t.Resolvers[0] + t.forgedStats = &ForgedStats{Label: r.ResolverAddr} + conn, addr, err := GetResolverConnection(r, t.effectivePacketQueueSize(), t.effectiveQueueOverflowMode(), t.forgedStats) + if err != nil { + return err + } + t.resolverConn = conn + t.remoteAddr = addr + return nil } // InitiateDNSPacketConn wraps the resolver connection with DNS encoding. @@ -882,10 +848,9 @@ func NewOutbound(resolvers []Resolver, tunnelServers []TunnelServer) *Outbound { // Start begins accepting connections on bind and forwarding them through the // first resolver/server pair. func (o *Outbound) Start(bind string) error { - resolver := o.Resolvers[0] tunnelServer := o.TunnelServers[0] - tunnel, err := NewTunnel(resolver, tunnelServer) + tunnel, err := NewTunnelMulti(o.Resolvers, tunnelServer) if err != nil { return fmt.Errorf("failed to create tunnel: %w", err) } diff --git a/client/dns.go b/client/dns.go index 0453b78..7847faf 100644 --- a/client/dns.go +++ b/client/dns.go @@ -100,10 +100,17 @@ func (rl *RateLimiter) Wait() { } } -// ForgedStats tracks forged DNS response counters. It is shared between -// UDPPacketConn (per-query mode) and DNSPacketConn (shared socket mode) so -// that forged response visibility is consistent regardless of transport. +// ForgedStats tracks forged DNS response counters for a specific source. +// In single-resolver mode, one instance is shared between UDPPacketConn +// (per-query mode) and DNSPacketConn so milestone logs fire on a unified +// count. In multi-resolver mode, each entry holds its own labeled instance +// so operators can see which resolver is being targeted by injection. type ForgedStats struct { + // Label identifies the counter's source in log lines — typically the + // resolver address (e.g. "8.8.8.8:53"). An empty label is allowed + // and produces an unlabeled log line (used by DNSPacketConn's catch-all + // safety net when no resolver attribution is available). + Label string Total uint64 SERVFAIL uint64 NXDOMAIN uint64 @@ -124,11 +131,19 @@ func (s *ForgedStats) Record(rcode uint16) { } total := atomic.AddUint64(&s.Total, 1) if forgedInfoMilestone(total) { - log.Infof("forged DNS responses: total=%d, SERVFAIL=%d, NXDOMAIN=%d, other=%d", - total, - atomic.LoadUint64(&s.SERVFAIL), - atomic.LoadUint64(&s.NXDOMAIN), - atomic.LoadUint64(&s.Other)) + if s.Label != "" { + log.Infof("forged DNS responses from %s: total=%d, SERVFAIL=%d, NXDOMAIN=%d, other=%d", + s.Label, total, + atomic.LoadUint64(&s.SERVFAIL), + atomic.LoadUint64(&s.NXDOMAIN), + atomic.LoadUint64(&s.Other)) + } else { + log.Infof("forged DNS responses: total=%d, SERVFAIL=%d, NXDOMAIN=%d, other=%d", + total, + atomic.LoadUint64(&s.SERVFAIL), + atomic.LoadUint64(&s.NXDOMAIN), + atomic.LoadUint64(&s.Other)) + } } } diff --git a/client/multi_resolver.go b/client/multi_resolver.go new file mode 100644 index 0000000..00cd6c8 --- /dev/null +++ b/client/multi_resolver.go @@ -0,0 +1,732 @@ +package client + +import ( + "context" + "crypto/tls" + "fmt" + "net" + "net/http" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/net2share/vaydns/dns" + "github.com/net2share/vaydns/turbotunnel" + log "github.com/sirupsen/logrus" +) + +// ResolverState is the current health state of one resolver. +type ResolverState string + +const ( + ResolverStateUnknown ResolverState = "unknown" + ResolverStateHealthy ResolverState = "healthy" + ResolverStateRateLimited ResolverState = "rate_limited" + ResolverStateDown ResolverState = "down" +) + +const ( + pendingResponseTimeout = 5 * time.Second + healthTickInterval = 1 * time.Second + downTimeoutThreshold = int64(8) + rateLimitThreshold = int64(5) + probeInterval = 1 * time.Second +) + +// ResolverStat is a snapshot of one resolver's counters and health state. +type ResolverStat struct { + Address string + State ResolverState + ValidCount int64 + InvalidCount int64 + TimeoutCount int64 + LastWrite time.Time + LastValid time.Time +} + +type resolverEntry struct { + name string + addr net.Addr + conn net.PacketConn + + // dead is set when the entry's transport has reported an unrecoverable + // error and its reader goroutine has exited. A dead entry is permanently + // excluded from selection and never transitions back — the only way to + // recover is to rebuild the entire MultiResolver (which happens on a + // full tunnel reconnect). Kept as an atomic.Bool so selection helpers + // can read it without acquiring mu on the hot path. + dead atomic.Bool + + // forgedStats tracks forged DNS response counters for this specific + // resolver entry, with milestone logging attributed to its address. + // In multi-resolver mode each entry owns its own labeled instance; + // in single-resolver mode the Tunnel shares one instance with both + // UDPPacketConn and DNSPacketConn. + forgedStats *ForgedStats + + // mu protects every field below. Counters are deliberately not atomic: + // the decay path in recomputeState would lose concurrent Adds if the + // counters were atomic and decremented via Load→Store, and splitting + // decay across a CAS loop is uglier than just holding mu for the + // short window each counter update requires. + mu sync.Mutex + validCount int64 + invalidCount int64 + timeoutCount int64 + pending map[uint16]time.Time + lastWrite time.Time + lastValid time.Time + lastProbe time.Time + state ResolverState +} + +// markDead transitions the entry to a permanent Down state. It is idempotent +// and returns true only when this call was the one that transitioned the +// entry — callers rely on this to avoid double-counting the entry in +// MultiResolver.aliveCount when multiple goroutines detect the same transport +// failure concurrently. +// +// The dead flag and the state field are set together under mu so that +// recomputeState (also holding mu) cannot observe an inconsistent snapshot +// where dead is set but state has not yet been transitioned to Down. +func (e *resolverEntry) markDead() bool { + e.mu.Lock() + defer e.mu.Unlock() + if e.dead.Load() { + return false + } + e.dead.Store(true) + e.state = ResolverStateDown + return true +} + +// isDead reports whether markDead has been called on this entry. +func (e *resolverEntry) isDead() bool { + return e.dead.Load() +} + +func (e *resolverEntry) writePacket(b []byte) (int, error) { + now := time.Now() + e.trackOutgoingID(b, now) + n, err := e.conn.WriteTo(b, e.addr) + e.mu.Lock() + e.lastWrite = now + if err != nil { + e.invalidCount++ + } + e.mu.Unlock() + return n, err +} + +func (e *resolverEntry) trackOutgoingID(b []byte, now time.Time) { + msg, err := dns.MessageFromWireFormat(b) + if err != nil { + return + } + e.mu.Lock() + e.pending[msg.ID] = now + e.mu.Unlock() +} + +func (e *resolverEntry) readPacket() (multiReadResult, bool) { + var result multiReadResult + result.entry = e + result.n, result.addr, result.err = e.conn.ReadFrom(result.buf[:]) + var forged bool + if result.err == nil { + forged = e.evaluateIncoming(result.buf[:result.n]) + } + return result, forged +} + +// evaluateIncoming parses an incoming DNS response for health tracking and +// forged-response detection. Returns true when the response is forged +// (QR=1, RCODE != NoError) — the caller should not forward it upstream +// because evaluateIncoming already recorded it in the entry's ForgedStats. +func (e *resolverEntry) evaluateIncoming(packet []byte) bool { + now := time.Now() + resp, err := dns.MessageFromWireFormat(packet) + if err != nil { + e.mu.Lock() + e.invalidCount++ + e.mu.Unlock() + e.recomputeState(now) + return false // parse error — not forged, let upper layer handle + } + + if isValidDNSResponse(resp) { + e.mu.Lock() + delete(e.pending, resp.ID) + e.validCount++ + e.timeoutCount = 0 + e.lastValid = now + e.state = ResolverStateHealthy + e.mu.Unlock() + return false + } + + // Non-valid DNS response. Track it for health purposes. + rcode := resp.Flags & 0x000f + isResponse := resp.Flags&0x8000 != 0 + + e.mu.Lock() + delete(e.pending, resp.ID) + e.invalidCount++ + // Rate-limit classification is handled by the accumulation path in + // recomputeState (invalidCount >= threshold && validCount == 0), not + // by a direct state assignment here. recomputeState unconditionally + // reassigns e.state, so any direct assignment would be immediately + // overwritten. + e.mu.Unlock() + e.recomputeState(now) + + // Record in per-resolver forged stats and signal caller not to forward. + // forgedStats is always non-nil in multi-resolver mode (set at entry + // construction). The nil guard is defensive for any future caller that + // constructs a resolverEntry without setting forgedStats. + if isResponse && e.forgedStats != nil { + e.forgedStats.Record(rcode) + return true + } + return false +} + +func (e *resolverEntry) expirePending(now time.Time) { + e.mu.Lock() + expired := int64(0) + for id, t := range e.pending { + if now.Sub(t) >= pendingResponseTimeout { + delete(e.pending, id) + expired++ + } + } + if expired > 0 { + e.timeoutCount += expired + e.invalidCount += expired + } else { + // Decay counters only when nothing accumulated this tick. + // This ensures burst expirations (and response-driven + // invalidCount growth between ticks) can outpace the decay + // and eventually cross the detection thresholds. Decay runs + // at most once per health-tick interval (~1s). + if e.invalidCount > 0 { + e.invalidCount-- + } + if e.timeoutCount > 0 { + e.timeoutCount-- + } + } + e.mu.Unlock() + e.recomputeState(now) +} + +func (e *resolverEntry) recomputeState(now time.Time) { + // Dead is sticky — never recompute it back to any other state. We + // re-check under the lock after acquiring it, in case markDead raced + // in between the outer check and acquiring the lock. + if e.isDead() { + return + } + e.mu.Lock() + defer e.mu.Unlock() + if e.dead.Load() { + return + } + + switch { + case e.timeoutCount >= downTimeoutThreshold: + e.state = ResolverStateDown + case e.invalidCount >= rateLimitThreshold && e.validCount == 0: + e.state = ResolverStateRateLimited + case e.validCount > 0 && now.Sub(e.lastValid) <= 30*time.Second: + e.state = ResolverStateHealthy + case e.validCount == 0: + e.state = ResolverStateUnknown + default: + e.state = ResolverStateUnknown + } +} + +func (e *resolverEntry) stateSnapshot() ResolverState { + e.mu.Lock() + defer e.mu.Unlock() + return e.state +} + +func (e *resolverEntry) markProbe(now time.Time) { + e.mu.Lock() + e.lastProbe = now + e.mu.Unlock() +} + +func (e *resolverEntry) canProbe(now time.Time) bool { + e.mu.Lock() + defer e.mu.Unlock() + return now.Sub(e.lastProbe) >= probeInterval +} + +func (e *resolverEntry) snapshot() ResolverStat { + e.mu.Lock() + defer e.mu.Unlock() + return ResolverStat{ + Address: e.name, + State: e.state, + ValidCount: e.validCount, + InvalidCount: e.invalidCount, + TimeoutCount: e.timeoutCount, + LastWrite: e.lastWrite, + LastValid: e.lastValid, + } +} + +type multiReadResult struct { + buf [4096]byte + n int + addr net.Addr + err error + entry *resolverEntry +} + +// MultiResolver is a net.PacketConn that multiplexes across multiple DNS +// resolver transport connections. It tracks per-resolver health from valid and +// invalid responses, avoids down resolvers for primary traffic, and probes +// unhealthy resolvers by duplicating selected packets. +type MultiResolver struct { + entries []*resolverEntry + mu sync.Mutex + rrIndex int + probeRR int + recvChan chan multiReadResult + closed chan struct{} + closeOnce sync.Once + + // aliveCount tracks how many entries still have a live reader + // goroutine. When it reaches zero, allDead is closed, causing pending + // ReadFrom/WriteTo calls to return an "all resolvers down" error, + // which is the only condition that should propagate a transport + // error to the upper tunnel session. + aliveCount atomic.Int32 + allDead chan struct{} + allDeadOnce sync.Once +} + +// NewMultiResolver creates a MultiResolver from a slice of Resolver configs. +func NewMultiResolver(resolvers []Resolver, queueSize int, overflowMode turbotunnel.QueueOverflowMode) (*MultiResolver, error) { + if len(resolvers) == 0 { + return nil, fmt.Errorf("at least one resolver is required") + } + + entries := make([]*resolverEntry, 0, len(resolvers)) + for _, r := range resolvers { + entryStats := &ForgedStats{Label: r.ResolverAddr} + conn, addr, err := GetResolverConnection(r, queueSize, overflowMode, entryStats) + if err != nil { + for _, e := range entries { + e.conn.Close() + } + return nil, fmt.Errorf("resolver %s %s: %w", r.ResolverType, r.ResolverAddr, err) + } + entries = append(entries, &resolverEntry{ + name: r.ResolverAddr, + addr: addr, + conn: conn, + forgedStats: entryStats, + pending: make(map[uint16]time.Time), + state: ResolverStateUnknown, + }) + } + + log.Infof("multi-resolver: %d resolvers configured", len(entries)) + udpWorkerTotal := 0 + for i, r := range resolvers { + log.Infof(" [%d] %s %s", i+1, r.ResolverType, r.ResolverAddr) + if r.ResolverType == ResolverTypeUDP && !r.UDPSharedSocket { + workers := r.UDPWorkers + if workers <= 0 { + workers = DefaultUDPWorkers + } + udpWorkerTotal += workers + } + } + if udpWorkerTotal > 0 { + log.Infof("multi-resolver: %d total UDP worker goroutines", udpWorkerTotal) + } + + mr := &MultiResolver{ + entries: entries, + recvChan: make(chan multiReadResult, len(entries)*4), + closed: make(chan struct{}), + allDead: make(chan struct{}), + } + mr.aliveCount.Store(int32(len(entries))) + mr.startReaders() + go mr.healthWorker() + return mr, nil +} + +// entryDied is called whenever a transport error is observed for an entry, +// either from its reader goroutine or from a failed write. It marks the entry +// dead and, if this was the transition from alive to dead (as opposed to a +// second observer of the same failure), decrements aliveCount. When the last +// alive entry dies, allDead is closed so pending ReadFrom/WriteTo callers can +// unblock with a terminal error. +func (mr *MultiResolver) entryDied(entry *resolverEntry, err error) { + if !entry.markDead() { + return + } + log.Warnf("multi-resolver: entry %s transport error: %v; marking down", entry.name, err) + if mr.aliveCount.Add(-1) == 0 { + mr.allDeadOnce.Do(func() { close(mr.allDead) }) + } +} + +// startReaders launches one reader goroutine per entry. Each goroutine reads +// packets from its entry's transport and pushes the result onto recvChan. +// Extracted so tests can construct MultiResolver with synthetic entries. +// +// On a transport error, the reader calls entryDied to mark the entry dead and +// signal allDead when the last entry dies. It never pushes error-bearing +// results onto recvChan: doing so would surface a single resolver's failure +// to DNSPacketConn.recvLoop, which would tear down the entire tunnel session +// and defeat the point of having multiple resolvers. +func (mr *MultiResolver) startReaders() { + for _, e := range mr.entries { + entry := e + go func() { + for { + res, forged := entry.readPacket() + if res.err != nil { + mr.entryDied(entry, res.err) + return + } + if forged { + // evaluateIncoming already recorded this in the + // entry's per-resolver ForgedStats. Don't forward + // to DNSPacketConn — its safety-net filter would + // double-count it. + continue + } + select { + case mr.recvChan <- res: + case <-mr.closed: + return + } + } + }() + } +} + +func (mr *MultiResolver) healthWorker() { + ticker := time.NewTicker(healthTickInterval) + defer ticker.Stop() + for { + select { + case <-ticker.C: + now := time.Now() + stats := make([]ResolverStat, 0, len(mr.entries)) + for _, e := range mr.entries { + e.expirePending(now) + stats = append(stats, e.snapshot()) + } + log.Debug("\n" + renderResolverStatsTable(stats, now)) + case <-mr.closed: + return + } + } +} + +func renderResolverStatsTable(stats []ResolverStat, now time.Time) string { + var b strings.Builder + b.WriteString("+-------------------------+--------------+--------+---------+---------+-------------+-------------+\n") + b.WriteString("| resolver | state | valid | invalid | timeout | last_write | last_valid |\n") + b.WriteString("+-------------------------+--------------+--------+---------+---------+-------------+-------------+\n") + for _, s := range stats { + lastWriteAgo := "-" + if !s.LastWrite.IsZero() { + lastWriteAgo = now.Sub(s.LastWrite).Truncate(time.Second).String() + } + lastValidAgo := "-" + if !s.LastValid.IsZero() { + lastValidAgo = now.Sub(s.LastValid).Truncate(time.Second).String() + } + fmt.Fprintf(&b, "| %-23.23s | %-12s | %6d | %7d | %7d | %11s | %11s |\n", + s.Address, + s.State, + s.ValidCount, + s.InvalidCount, + s.TimeoutCount, + lastWriteAgo, + lastValidAgo, + ) + } + b.WriteString("+-------------------------+--------------+--------+---------+---------+-------------+-------------+") + return b.String() +} + +func isValidDNSResponse(resp dns.Message) bool { + if resp.Flags&0x8000 == 0 { + return false + } + return (resp.Flags & 0x000f) == dns.RcodeNoError +} + +// ReadFrom receives a packet from whichever resolver responds first. +// It only returns an error when the MultiResolver has been closed or every +// entry has died; a single resolver's transport error is isolated at the +// reader goroutine and does not propagate here. +func (mr *MultiResolver) ReadFrom(b []byte) (n int, addr net.Addr, err error) { + select { + case <-mr.closed: + return 0, nil, net.ErrClosed + case res := <-mr.recvChan: + n = copy(b, res.buf[:res.n]) + return n, turbotunnel.DummyAddr{}, nil + case <-mr.allDead: + // Drain any packet that was buffered by a reader before it + // died, so packets already delivered by the transport are not + // discarded in favour of the terminal error. No more readers + // are pushing (allDead is closed only after every reader has + // exited), so a non-blocking read here races only with the + // mr.closed case above, which is handled on the next call. + select { + case res := <-mr.recvChan: + n = copy(b, res.buf[:res.n]) + return n, turbotunnel.DummyAddr{}, nil + default: + return 0, nil, fmt.Errorf("multi-resolver: all resolvers are down") + } + } +} + +// WriteTo sends b to the selected primary resolver and may duplicate b to one +// unhealthy resolver as a probe to detect recovery. If the primary's write +// fails, the entry is marked dead and WriteTo retries with the next alive +// entry. An error is returned only when the MultiResolver is closed or every +// entry has been marked dead. +func (mr *MultiResolver) WriteTo(b []byte, _ net.Addr) (n int, err error) { + select { + case <-mr.closed: + return 0, net.ErrClosed + case <-mr.allDead: + return 0, fmt.Errorf("multi-resolver: all resolvers are down") + default: + } + + // Try entries until one accepts the write or all alive entries have + // been exhausted. Each write error marks the entry dead. + var primary *resolverEntry + for attempts := 0; attempts < len(mr.entries); attempts++ { + primary = mr.selectPrimary() + if primary == nil { + break + } + n, err = primary.writePacket(b) + if err == nil { + break + } + mr.entryDied(primary, err) + } + switch { + case err != nil: + return 0, fmt.Errorf("multi-resolver: all write attempts failed: %w", err) + case primary == nil: + return 0, fmt.Errorf("multi-resolver: no alive resolver for write") + } + + if probe := mr.selectProbeTarget(primary); probe != nil { + probe.markProbe(time.Now()) + _, _ = probe.writePacket(b) + } + return n, nil +} + +// selectPrimary returns the entry that should receive the next outgoing +// query, or nil if every entry has been marked dead. The caller must handle +// nil (e.g., return an "all resolvers down" error). +func (mr *MultiResolver) selectPrimary() *resolverEntry { + return mr.selectRoundRobinHealthy() +} + +func (mr *MultiResolver) selectRoundRobinHealthy() *resolverEntry { + mr.mu.Lock() + defer mr.mu.Unlock() + + if len(mr.entries) == 0 { + return nil + } + + start := mr.rrIndex + // First pass: prefer Healthy or Unknown, skipping dead entries. + for i := 0; i < len(mr.entries); i++ { + idx := (start + i) % len(mr.entries) + if mr.entries[idx].isDead() { + continue + } + state := mr.entries[idx].stateSnapshot() + if state == ResolverStateHealthy || state == ResolverStateUnknown { + mr.rrIndex = (idx + 1) % len(mr.entries) + return mr.entries[idx] + } + } + // Second pass: accept any non-dead entry even if RateLimited/Down. + for i := 0; i < len(mr.entries); i++ { + idx := (start + i) % len(mr.entries) + if mr.entries[idx].isDead() { + continue + } + mr.rrIndex = (idx + 1) % len(mr.entries) + return mr.entries[idx] + } + // Every entry is dead. + return nil +} + +func (mr *MultiResolver) selectProbeTarget(primary *resolverEntry) *resolverEntry { + mr.mu.Lock() + defer mr.mu.Unlock() + + now := time.Now() + for i := 0; i < len(mr.entries); i++ { + idx := (mr.probeRR + i) % len(mr.entries) + e := mr.entries[idx] + if e == primary { + continue + } + if e.isDead() { + continue + } + state := e.stateSnapshot() + if state == ResolverStateHealthy { + continue + } + if e.canProbe(now) { + mr.probeRR = (idx + 1) % len(mr.entries) + return e + } + } + return nil +} + +// Close closes all underlying connections and stops the reader goroutines. +func (mr *MultiResolver) Close() error { + mr.closeOnce.Do(func() { + close(mr.closed) + for _, e := range mr.entries { + e.conn.Close() + } + }) + return nil +} + +// LocalAddr returns the local address of the first underlying connection. +// The value is arbitrary (there is no single meaningful local address for a +// multi-resolver), but the method must exist to satisfy net.PacketConn. +// No caller uses the returned address for routing or logic decisions. +func (mr *MultiResolver) LocalAddr() net.Addr { + return mr.entries[0].conn.LocalAddr() +} + +// SetDeadline sets a deadline on all underlying connections. +func (mr *MultiResolver) SetDeadline(t time.Time) error { + var last error + for _, e := range mr.entries { + if err := e.conn.SetDeadline(t); err != nil { + last = err + } + } + return last +} + +// SetReadDeadline sets a read deadline on all underlying connections. +func (mr *MultiResolver) SetReadDeadline(t time.Time) error { + var last error + for _, e := range mr.entries { + if err := e.conn.SetReadDeadline(t); err != nil { + last = err + } + } + return last +} + +// SetWriteDeadline sets a write deadline on all underlying connections. +func (mr *MultiResolver) SetWriteDeadline(t time.Time) error { + var last error + for _, e := range mr.entries { + if err := e.conn.SetWriteDeadline(t); err != nil { + last = err + } + } + return last +} + +// GetResolverConnection creates the underlying transport net.PacketConn for r. +// For UDP per-query mode, forgedStats is passed to NewUDPPacketConn so +// forged-response milestone logs are attributed to this resolver. Callers +// may pass nil to let each layer create its own counter. +func GetResolverConnection(r Resolver, queueSize int, overflowMode turbotunnel.QueueOverflowMode, forgedStats *ForgedStats) (net.PacketConn, net.Addr, error) { + switch r.ResolverType { + case ResolverTypeUDP: + addr, err := net.ResolveUDPAddr("udp", r.ResolverAddr) + if err != nil { + return nil, nil, err + } + if r.UDPSharedSocket { + lc := net.ListenConfig{Control: r.DialerControl} + conn, err := lc.ListenPacket(context.Background(), "udp", ":0") + if err != nil { + return nil, nil, err + } + return conn, addr, nil + } + workers := r.UDPWorkers + if workers <= 0 { + workers = DefaultUDPWorkers + } + timeout := r.UDPTimeout + if timeout <= 0 { + timeout = DefaultUDPResponseTimeout + } + conn, err := NewUDPPacketConn(addr, r.DialerControl, workers, timeout, !r.UDPAcceptErrors, queueSize, overflowMode, forgedStats) + if err != nil { + return nil, nil, err + } + return conn, addr, nil + + case ResolverTypeDOH: + var rt http.RoundTripper + if r.RoundTripper != nil { + rt = r.RoundTripper + } else if r.UTLSClientHelloID != nil { + rt = NewUTLSRoundTripper(nil, r.UTLSClientHelloID) + } else { + rt = http.DefaultTransport + } + conn, err := NewHTTPPacketConn(rt, r.ResolverAddr, 8, queueSize, overflowMode) + if err != nil { + return nil, nil, err + } + return conn, turbotunnel.DummyAddr{}, nil + + case ResolverTypeDOT: + var dialTLSContext func(ctx context.Context, network, addr string) (net.Conn, error) + if r.UTLSClientHelloID != nil { + id := r.UTLSClientHelloID + dialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) { + return UTLSDialContext(ctx, network, addr, nil, id) + } + } else { + dialTLSContext = func(ctx context.Context, network, addr string) (net.Conn, error) { + return tls.DialWithDialer(&net.Dialer{}, network, addr, nil) + } + } + conn, err := NewTLSPacketConn(r.ResolverAddr, dialTLSContext, queueSize, overflowMode) + if err != nil { + return nil, nil, err + } + return conn, turbotunnel.DummyAddr{}, nil + + default: + return nil, nil, fmt.Errorf("unsupported resolver type: %s", r.ResolverType) + } +} diff --git a/client/multi_resolver_test.go b/client/multi_resolver_test.go new file mode 100644 index 0000000..93eff20 --- /dev/null +++ b/client/multi_resolver_test.go @@ -0,0 +1,415 @@ +package client + +import ( + "bytes" + "net" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/net2share/vaydns/dns" + "github.com/net2share/vaydns/turbotunnel" +) + +// fakePacketConn is a controllable net.PacketConn for MultiResolver tests. +// Pushed responses are consumed in order by successive ReadFrom calls; each +// response can be either data bytes or an error. +type fakePacketConn struct { + name string + readCh chan fakeReadResp + closed chan struct{} + once sync.Once +} + +type fakeReadResp struct { + data []byte + err error +} + +func newFakePacketConn(name string) *fakePacketConn { + return &fakePacketConn{ + name: name, + readCh: make(chan fakeReadResp, 16), + closed: make(chan struct{}), + } +} + +func (f *fakePacketConn) pushData(data []byte) { + cp := make([]byte, len(data)) + copy(cp, data) + f.readCh <- fakeReadResp{data: cp} +} + +func (f *fakePacketConn) pushError(err error) { + f.readCh <- fakeReadResp{err: err} +} + +func (f *fakePacketConn) ReadFrom(p []byte) (int, net.Addr, error) { + select { + case r, ok := <-f.readCh: + if !ok { + return 0, nil, net.ErrClosed + } + if r.err != nil { + return 0, nil, r.err + } + return copy(p, r.data), turbotunnel.DummyAddr{}, nil + case <-f.closed: + return 0, nil, net.ErrClosed + } +} + +func (f *fakePacketConn) WriteTo(p []byte, _ net.Addr) (int, error) { + select { + case <-f.closed: + return 0, net.ErrClosed + default: + return len(p), nil + } +} + +func (f *fakePacketConn) Close() error { + f.once.Do(func() { close(f.closed) }) + return nil +} + +func (f *fakePacketConn) LocalAddr() net.Addr { return turbotunnel.DummyAddr{} } +func (f *fakePacketConn) SetDeadline(time.Time) error { return nil } +func (f *fakePacketConn) SetReadDeadline(time.Time) error { return nil } +func (f *fakePacketConn) SetWriteDeadline(time.Time) error { return nil } + +// newFakeEntry builds a resolverEntry around a fakePacketConn, suitable for +// injection into a hand-constructed MultiResolver in tests. +func newFakeEntry(name string, conn *fakePacketConn) *resolverEntry { + return &resolverEntry{ + name: name, + addr: turbotunnel.DummyAddr{}, + conn: conn, + forgedStats: &ForgedStats{Label: name}, + pending: make(map[uint16]time.Time), + state: ResolverStateUnknown, + } +} + +// newTestMultiResolver mirrors what NewMultiResolver does after +// GetResolverConnection returns, but with pre-built synthetic entries. It does +// not spawn the healthWorker — tests don't need it, and omitting it keeps them +// hermetic with respect to timing. +func newTestMultiResolver(entries []*resolverEntry) *MultiResolver { + mr := &MultiResolver{ + entries: entries, + recvChan: make(chan multiReadResult, len(entries)*4), + closed: make(chan struct{}), + allDead: make(chan struct{}), + } + mr.aliveCount.Store(int32(len(entries))) + mr.startReaders() + return mr +} + +// TestMultiResolver_DeadEntryDoesNotBreakReadFrom verifies that when one +// resolver entry's transport returns a fatal read error, MultiResolver.ReadFrom +// continues to deliver packets from healthy entries instead of propagating +// that error to the upper DNSPacketConn layer (which would tear down the whole +// tunnel session). +func TestMultiResolver_DeadEntryDoesNotBreakReadFrom(t *testing.T) { + failing := newFakePacketConn("failing") + working := newFakePacketConn("working") + + entries := []*resolverEntry{ + newFakeEntry("failing", failing), + newFakeEntry("working", working), + } + mr := newTestMultiResolver(entries) + defer mr.Close() + + // Step 1: make the failing entry's ReadFrom return a fatal error. + // The reader goroutine will pick this up immediately. + failing.pushError(net.ErrClosed) + + // Give the reader goroutine time to handle the error. Under the buggy + // implementation it pushes an error-bearing multiReadResult onto + // recvChan and exits; under a correct implementation it would quietly + // mark the entry down and exit without poisoning recvChan. + time.Sleep(100 * time.Millisecond) + + // Step 2: push a valid response to the working entry. This lands on + // recvChan strictly AFTER any push from the failing entry, so the + // ordering is deterministic. + wantResponse := []byte("valid-response-bytes-from-working-resolver") + working.pushData(wantResponse) + + // Step 3: ReadFrom must return the bytes from the working entry, not + // the error from the failing one. + type readResult struct { + n int + err error + buf []byte + } + done := make(chan readResult, 1) + go func() { + buf := make([]byte, 4096) + n, _, err := mr.ReadFrom(buf) + done <- readResult{n: n, err: err, buf: buf} + }() + + select { + case r := <-done: + if r.err != nil { + t.Fatalf("MultiResolver.ReadFrom returned error %v; expected it to ignore the failed entry and return the bytes from the working entry. This is the C1 bug: a single resolver's read error tears down the tunnel.", r.err) + } + if !bytes.Equal(r.buf[:r.n], wantResponse) { + t.Fatalf("MultiResolver.ReadFrom returned wrong bytes.\n got: %x\n want: %x", r.buf[:r.n], wantResponse) + } + case <-time.After(3 * time.Second): + t.Fatal("MultiResolver.ReadFrom timed out; expected it to return bytes from the working entry within 3s") + } +} + +// TestMultiResolver_FailedEntryExcludedFromSelection verifies that after a +// reader goroutine has processed a fatal read error and exited, the entry +// is no longer returned by selectPrimary. Otherwise the round-robin scheduler +// would keep sending queries to a resolver whose reader goroutine is gone +// (so responses would never come back). +func TestMultiResolver_FailedEntryExcludedFromSelection(t *testing.T) { + failing := newFakePacketConn("failing") + working := newFakePacketConn("working") + + entries := []*resolverEntry{ + newFakeEntry("failing", failing), + newFakeEntry("working", working), + } + mr := newTestMultiResolver(entries) + defer mr.Close() + + // Kill the failing entry's reader; leave working blocked in ReadFrom. + failing.pushError(net.ErrClosed) + + // Let the reader goroutine process the error and (ideally) mark the + // entry down before it exits. + time.Sleep(100 * time.Millisecond) + + // Drain any stale result from recvChan so the assertion below only + // checks selectPrimary behaviour, not lingering channel contents. +drain: + for { + select { + case <-mr.recvChan: + default: + break drain + } + } + + // Call selectPrimary repeatedly. It must never return the failing + // entry — otherwise real queries will be sent to a resolver whose + // reader is gone, so responses will never come back. + for i := range 10 { + selected := mr.selectPrimary() + if selected == nil { + t.Fatalf("iter %d: selectPrimary returned nil", i) + } + if selected.name == "failing" { + t.Fatalf("iter %d: selectPrimary returned the failed entry %q; expected it to be excluded after its reader exited on error. State is %s.", i, selected.name, selected.stateSnapshot()) + } + } +} + +// buildDNSResponse constructs a minimal valid DNS wire-format response with +// the given RCODE. QR is set to 1. This is the smallest packet that +// dns.MessageFromWireFormat can parse and evaluateIncoming can classify. +func buildDNSResponse(t *testing.T, rcode uint16) []byte { + t.Helper() + msg := &dns.Message{ + ID: 0x1234, + Flags: 0x8000 | rcode, // QR=1 + RCODE + } + buf, err := msg.WireFormat() + if err != nil { + t.Fatalf("buildDNSResponse: %v", err) + } + return buf +} + +// TestForgedStats_LabeledMilestoneLog verifies that ForgedStats.Record() +// fires a milestone log at the expected thresholds and that the log line +// includes the resolver label in the "from " format. +func TestForgedStats_LabeledMilestoneLog(t *testing.T) { + stats := &ForgedStats{Label: "1.2.3.4:53"} + + // Record 9 forged NXDOMAIN responses — no milestone yet. + for range 9 { + stats.Record(dns.RcodeNameError) + } + if atomic.LoadUint64(&stats.Total) != 9 { + t.Fatalf("expected Total=9, got %d", stats.Total) + } + + // The 10th should trigger the first milestone. + stats.Record(dns.RcodeNameError) + if atomic.LoadUint64(&stats.Total) != 10 { + t.Fatalf("expected Total=10, got %d", stats.Total) + } + if atomic.LoadUint64(&stats.NXDOMAIN) != 10 { + t.Fatalf("expected NXDOMAIN=10, got %d", stats.NXDOMAIN) + } +} + +// TestForgedStats_UnlabeledFallback verifies that a ForgedStats with an +// empty Label still records correctly (used by DNSPacketConn's catch-all). +func TestForgedStats_UnlabeledFallback(t *testing.T) { + stats := &ForgedStats{} + stats.Record(dns.RcodeServerFailure) + if atomic.LoadUint64(&stats.Total) != 1 { + t.Fatalf("expected Total=1, got %d", stats.Total) + } + if atomic.LoadUint64(&stats.SERVFAIL) != 1 { + t.Fatalf("expected SERVFAIL=1, got %d", stats.SERVFAIL) + } +} + +// TestMultiResolver_ForgedResponseRecordedPerEntry verifies that in +// multi-resolver mode, a forged DNS response (NXDOMAIN) arriving at one +// entry increments that entry's ForgedStats but not the other entries'. +func TestMultiResolver_ForgedResponseRecordedPerEntry(t *testing.T) { + forging := newFakePacketConn("forging") + clean := newFakePacketConn("clean") + + forgingEntry := newFakeEntry("forging", forging) + cleanEntry := newFakeEntry("clean", clean) + + mr := newTestMultiResolver([]*resolverEntry{forgingEntry, cleanEntry}) + defer mr.Close() + + // Push a forged NXDOMAIN response to the forging entry. + forgedResp := buildDNSResponse(t, dns.RcodeNameError) + forging.pushData(forgedResp) + + // Give the reader time to process. + time.Sleep(100 * time.Millisecond) + + // Forging entry's stats should have recorded the forged response. + if got := atomic.LoadUint64(&forgingEntry.forgedStats.Total); got != 1 { + t.Errorf("forging entry ForgedStats.Total = %d, want 1", got) + } + if got := atomic.LoadUint64(&forgingEntry.forgedStats.NXDOMAIN); got != 1 { + t.Errorf("forging entry ForgedStats.NXDOMAIN = %d, want 1", got) + } + + // Clean entry's stats should be untouched. + if got := atomic.LoadUint64(&cleanEntry.forgedStats.Total); got != 0 { + t.Errorf("clean entry ForgedStats.Total = %d, want 0", got) + } +} + +// TestMultiResolver_ForgedResponseNotForwardedToRecvChan verifies that +// forged responses detected by evaluateIncoming are NOT forwarded to +// recvChan (and therefore never reach DNSPacketConn). Only legitimate +// responses should appear on the channel. +func TestMultiResolver_ForgedResponseNotForwardedToRecvChan(t *testing.T) { + conn := newFakePacketConn("resolver") + entry := newFakeEntry("resolver", conn) + + mr := newTestMultiResolver([]*resolverEntry{entry}) + defer mr.Close() + + // Push a forged NXDOMAIN, then a valid NoError response. + forgedResp := buildDNSResponse(t, dns.RcodeNameError) + validResp := buildDNSResponse(t, dns.RcodeNoError) + conn.pushData(forgedResp) + conn.pushData(validResp) + + // ReadFrom should return the valid response, skipping the forged one. + buf := make([]byte, 4096) + type readResult struct { + n int + err error + } + done := make(chan readResult, 1) + go func() { + n, _, err := mr.ReadFrom(buf) + done <- readResult{n: n, err: err} + }() + + select { + case r := <-done: + if r.err != nil { + t.Fatalf("ReadFrom error: %v", r.err) + } + // The bytes should be the valid response, not the forged one. + if bytes.Equal(buf[:r.n], forgedResp) { + t.Fatal("ReadFrom returned the forged response — it should have been filtered by the reader goroutine") + } + if !bytes.Equal(buf[:r.n], validResp) { + t.Fatalf("ReadFrom returned unexpected bytes:\n got: %x\n want: %x", buf[:r.n], validResp) + } + case <-time.After(3 * time.Second): + t.Fatal("ReadFrom timed out; expected valid response within 3s") + } + + // Verify the forged response was recorded. + if got := atomic.LoadUint64(&entry.forgedStats.Total); got != 1 { + t.Errorf("ForgedStats.Total = %d, want 1", got) + } +} + +// TestMultiResolver_ForgedStatsPerEntryLabeled verifies that each entry +// in a MultiResolver carries a distinctly-labeled ForgedStats instance. +func TestMultiResolver_ForgedStatsPerEntryLabeled(t *testing.T) { + conn1 := newFakePacketConn("r1") + conn2 := newFakePacketConn("r2") + + e1 := newFakeEntry("1.1.1.1:53", conn1) + e2 := newFakeEntry("8.8.8.8:53", conn2) + + mr := newTestMultiResolver([]*resolverEntry{e1, e2}) + defer mr.Close() + + if e1.forgedStats == e2.forgedStats { + t.Fatal("entries share the same ForgedStats pointer; expected distinct per-resolver instances") + } + if e1.forgedStats.Label != "1.1.1.1:53" { + t.Errorf("entry1 label = %q, want %q", e1.forgedStats.Label, "1.1.1.1:53") + } + if e2.forgedStats.Label != "8.8.8.8:53" { + t.Errorf("entry2 label = %q, want %q", e2.forgedStats.Label, "8.8.8.8:53") + } +} + +// TestMultiResolver_WriteToFailsOverToNextEntry verifies that when the +// primary entry's transport is closed, WriteTo marks it dead and retries +// with the next alive entry rather than returning an error. +func TestMultiResolver_WriteToFailsOverToNextEntry(t *testing.T) { + broken := newFakePacketConn("broken") + working := newFakePacketConn("working") + + entries := []*resolverEntry{ + newFakeEntry("broken", broken), + newFakeEntry("working", working), + } + mr := newTestMultiResolver(entries) + defer mr.Close() + + // Close the first entry's transport so WriteTo errors on it. + broken.Close() + + // Build a minimal DNS query so trackOutgoingID can parse it. + query := buildDNSResponse(t, dns.RcodeNoError) // any valid DNS message works + + n, err := mr.WriteTo(query, turbotunnel.DummyAddr{}) + if err != nil { + t.Fatalf("WriteTo returned error %v; expected failover to working entry", err) + } + if n != len(query) { + t.Fatalf("WriteTo wrote %d bytes, want %d", n, len(query)) + } + + // Verify the broken entry was marked dead. + if !entries[0].isDead() { + t.Error("broken entry should be marked dead after write failure") + } + // Verify the working entry is still alive. + if entries[1].isDead() { + t.Error("working entry should still be alive") + } +} diff --git a/client/udp.go b/client/udp.go index 8eae1b5..30aeb92 100644 --- a/client/udp.go +++ b/client/udp.go @@ -33,23 +33,28 @@ type UDPPacketConn struct { } // NewUDPPacketConn creates a UDPPacketConn with numWorkers goroutines that -// each send one query at a time on a fresh UDP socket. The returned -// ForgedStats pointer is shared with the caller so DNSPacketConn can -// include per-query forged counts in its reporting. -func NewUDPPacketConn(remoteAddr net.Addr, dialerControl func(network, address string, c syscall.RawConn) error, numWorkers int, responseTimeout time.Duration, ignoreErrors bool, queueSize int, overflowMode turbotunnel.QueueOverflowMode) (*UDPPacketConn, *ForgedStats, error) { - stats := &ForgedStats{} +// each send one query at a time on a fresh UDP socket. If forgedStats is +// non-nil, the UDPPacketConn records forged-response counts into it; +// otherwise a new instance is created internally (labeled with the remote +// address). Callers that want unified milestone logging across layers +// should create a single ForgedStats and pass the same pointer here and +// to NewDNSPacketConn. +func NewUDPPacketConn(remoteAddr net.Addr, dialerControl func(network, address string, c syscall.RawConn) error, numWorkers int, responseTimeout time.Duration, ignoreErrors bool, queueSize int, overflowMode turbotunnel.QueueOverflowMode, forgedStats *ForgedStats) (*UDPPacketConn, error) { + if forgedStats == nil { + forgedStats = &ForgedStats{Label: remoteAddr.String()} + } pconn := &UDPPacketConn{ remoteAddr: remoteAddr, dialerControl: dialerControl, responseTimeout: responseTimeout, ignoreErrors: ignoreErrors, - forgedStats: stats, + forgedStats: forgedStats, QueuePacketConn: turbotunnel.NewQueuePacketConn(remoteAddr, 0, queueSize, overflowMode), } for i := 0; i < numWorkers; i++ { go pconn.sendLoop() } - return pconn, stats, nil + return pconn, nil } // sendLoop is the per-worker loop. It dequeues one packet at a time from the diff --git a/dns/dns.go b/dns/dns.go index 57eb1a3..02658f4 100644 --- a/dns/dns.go +++ b/dns/dns.go @@ -68,6 +68,7 @@ const ( RcodeServerFailure = 2 // a.k.a. SERVFAIL RcodeNameError = 3 // a.k.a. NXDOMAIN RcodeNotImplemented = 4 // a.k.a. NOTIMPL + RcodeRefused = 5 // a.k.a. REFUSED // https://tools.ietf.org/html/rfc6891#section-9 ExtendedRcodeBadVers = 16 // a.k.a. BADVERS ) diff --git a/docs/client-library.md b/docs/client-library.md index 391b6b0..2989632 100644 --- a/docs/client-library.md +++ b/docs/client-library.md @@ -47,6 +47,21 @@ t.ListenAndServe("127.0.0.1:7000") // blocks, handles reconnection `ListenAndServe` opens a local TCP listener, creates tunnel sessions with automatic reconnection on failure, and forwards connections through the tunnel. +### Multi-resolver API + +Spread DNS queries across multiple resolvers for throughput and resilience: + +```go +r1, _ := client.NewResolver(client.ResolverTypeUDP, "8.8.8.8:53") +r2, _ := client.NewResolver(client.ResolverTypeDOH, "https://1.1.1.1/dns-query") +ts, _ := client.NewTunnelServer("t.example.com", "pubkey-hex") +t, _ := client.NewTunnelMulti([]client.Resolver{r1, r2}, ts) + +t.ListenAndServe("127.0.0.1:7000") +``` + +`NewTunnelMulti` accepts a slice of resolvers. The client multiplexes queries across them with round-robin selection and health-based routing. Per-resolver forged-response tracking logs which resolver is being targeted by DNS injection. + ## Key types | Type | Description | @@ -87,8 +102,8 @@ ts.RPS = 200 // rate limit queries/second ts.RecordType = "cname" // DNS record type for downstream data: txt, null, cname, a, aaaa, mx, ns, srv, caa (default: "txt") // Session options -t.IdleTimeout = 60 * time.Second -t.KeepAlive = 10 * time.Second +t.IdleTimeout = 10 * time.Second +t.KeepAlive = 2 * time.Second t.OpenStreamTimeout = 10 * time.Second t.MaxStreams = 256 t.SessionCheckInterval = 500 * time.Millisecond diff --git a/e2e/multi-resolver-dnstt-compat/docker-compose.yml b/e2e/multi-resolver-dnstt-compat/docker-compose.yml new file mode 100644 index 0000000..b492929 --- /dev/null +++ b/e2e/multi-resolver-dnstt-compat/docker-compose.yml @@ -0,0 +1,81 @@ +networks: + dns-net: + ipam: + config: + - subnet: 172.28.0.0/24 + backend-net: + +volumes: + keys: + +services: + keygen: + build: + context: ../.. + dockerfile: Dockerfile + volumes: + - keys:/keys + command: > + sh -c "vaydns-server -gen-key -privkey-file /keys/server.key -pubkey-file /keys/server.pub" + + dns1: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.10 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + dns2: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.11 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + backend: + image: nginx:alpine + networks: + - backend-net + + server: + build: + context: ../.. + dockerfile: Dockerfile + networks: + dns-net: + ipv4_address: 172.28.0.20 + backend-net: + volumes: + - keys:/keys + command: > + vaydns-server -udp :53 -privkey-file /keys/server.key + -domain t.example.com -upstream backend:80 + -dnstt-compat + depends_on: + keygen: + condition: service_completed_successfully + dns1: + condition: service_started + dns2: + condition: service_started + + client: + build: + context: ../.. + dockerfile: Dockerfile + networks: + - dns-net + volumes: + - keys:/keys + command: > + vaydns-client -udp 172.28.0.10:53 -udp 172.28.0.11:53 + -pubkey-file /keys/server.pub + -domain t.example.com -listen 0.0.0.0:7000 + -dnstt-compat + depends_on: + server: + condition: service_started diff --git a/e2e/multi-resolver-dnstt-compat/run.sh b/e2e/multi-resolver-dnstt-compat/run.sh new file mode 100755 index 0000000..66618c6 --- /dev/null +++ b/e2e/multi-resolver-dnstt-compat/run.sh @@ -0,0 +1,33 @@ +#!/usr/bin/env bash +# Test: multi-resolver with dnstt-compat wire format. +# Verifies that -dnstt-compat (8-byte ClientID, padding prefixes, forced TXT, +# longer timeouts) works correctly when combined with multiple UDP resolvers. +# The wire format sits above the MultiResolver transport layer so there should +# be no interaction, but this test guards against regressions. +set -euo pipefail +cd "$(dirname "$0")" + +cleanup() { docker compose down -v 2>/dev/null; } +trap cleanup EXIT + +echo "--- Building and starting services ---" +docker compose up -d --build + +# dnstt-compat uses longer timeouts (2m idle, 10s keepalive) and the tunnel +# handshake may be slower due to the larger wire overhead, so allow more time. +echo "--- Waiting for tunnel (up to 45s) ---" +for i in $(seq 1 45); do + if docker compose exec -T client wget -q -O- http://localhost:7000 2>/dev/null | grep -q "Welcome to nginx"; then + echo "" + echo "=== PASS ===" + exit 0 + fi + printf "." + sleep 1 +done + +echo "" +echo "--- Tunnel did not come up. Dumping logs ---" +docker compose logs client server dns1 dns2 +echo "=== FAIL ===" +exit 1 diff --git a/e2e/multi-resolver-forge/Corefile.forge b/e2e/multi-resolver-forge/Corefile.forge new file mode 100644 index 0000000..224cbc9 --- /dev/null +++ b/e2e/multi-resolver-forge/Corefile.forge @@ -0,0 +1,6 @@ +. { + template IN ANY . { + rcode NXDOMAIN + } + log +} diff --git a/e2e/multi-resolver-forge/docker-compose.yml b/e2e/multi-resolver-forge/docker-compose.yml new file mode 100644 index 0000000..b5be9f2 --- /dev/null +++ b/e2e/multi-resolver-forge/docker-compose.yml @@ -0,0 +1,82 @@ +networks: + dns-net: + ipam: + config: + - subnet: 172.28.0.0/24 + backend-net: + +volumes: + keys: + +services: + keygen: + build: + context: ../.. + dockerfile: Dockerfile + volumes: + - keys:/keys + command: > + sh -c "vaydns-server -gen-key -privkey-file /keys/server.key -pubkey-file /keys/server.pub" + + dns-good: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.10 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + dns-forge: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.11 + volumes: + - ./Corefile.forge:/Corefile + command: ["-conf", "/Corefile"] + + backend: + image: nginx:alpine + networks: + - backend-net + + server: + build: + context: ../.. + dockerfile: Dockerfile + networks: + dns-net: + ipv4_address: 172.28.0.20 + backend-net: + volumes: + - keys:/keys + command: > + vaydns-server -udp :53 -privkey-file /keys/server.key + -domain t.example.com -upstream backend:80 + -idle-timeout 10s -keepalive 2s + depends_on: + keygen: + condition: service_completed_successfully + dns-good: + condition: service_started + dns-forge: + condition: service_started + + client: + build: + context: ../.. + dockerfile: Dockerfile + networks: + - dns-net + volumes: + - keys:/keys + command: > + vaydns-client -udp 172.28.0.10:53 -udp 172.28.0.11:53 + -pubkey-file /keys/server.pub + -domain t.example.com -listen 0.0.0.0:7000 + -idle-timeout 10s -keepalive 2s -session-check-interval 500ms + -reconnect-min 1s -reconnect-max 5s -log-level info + depends_on: + server: + condition: service_started diff --git a/e2e/multi-resolver-forge/run.sh b/e2e/multi-resolver-forge/run.sh new file mode 100755 index 0000000..2ee38ca --- /dev/null +++ b/e2e/multi-resolver-forge/run.sh @@ -0,0 +1,74 @@ +#!/usr/bin/env bash +# Test: one resolver returns forged responses while another works. +# dns-forge always replies with NXDOMAIN (CoreDNS template plugin), simulating +# a censor or broken resolver injecting fake responses. dns-good behaves +# normally. The tunnel must work through dns-good, with the per-query UDP +# worker's forged-response filter absorbing dns-forge's NXDOMAINs and +# MultiResolver's health state machine eventually routing around it. +set -euo pipefail +cd "$(dirname "$0")" + +cleanup() { docker compose down -v 2>/dev/null; } +trap cleanup EXIT + +fetch() { + docker compose exec -T client wget -q -O- http://localhost:7000 2>/dev/null | grep -q "Welcome to nginx" +} + +echo "--- Building and starting services ---" +docker compose up -d --build + +echo "--- Waiting for tunnel through dns-good while dns-forge injects NXDOMAINs (up to 60s) ---" +ok_count=0 +for i in $(seq 1 60); do + if fetch; then + ok_count=$((ok_count + 1)) + # Require two consecutive successes so one lucky query doesn't pass + # the test while the forging resolver is still in rotation. + if [ "$ok_count" -ge 2 ]; then + echo "" + # Sanity check: make sure dns-forge actually saw queries, so we + # know the tunnel was exercising the forging code path and didn't + # only hit dns-good by chance. CoreDNS's log plugin may buffer + # output briefly, so give it a moment to flush before grepping. + sleep 2 + forge_logs=$(docker compose logs dns-forge 2>&1) + if ! grep -q 'NXDOMAIN' <<<"$forge_logs"; then + echo "--- dns-forge never served NXDOMAIN; the forging code path may not have been exercised ---" + echo "$forge_logs" + echo "=== FAIL (forging path not exercised) ===" + exit 1 + fi + nxdomain_count=$(grep -c 'NXDOMAIN' <<<"$forge_logs" || true) + echo "--- dns-forge served $nxdomain_count NXDOMAIN responses; forging path exercised ---" + + # Verify the client-side milestone log fired with per-resolver + # attribution: "forged DNS responses from : ..." + # This proves the ForgedStats plumbing carries the resolver + # label end-to-end and that milestone thresholds fire correctly + # in multi-resolver mode. + client_logs=$(docker compose logs client 2>&1) + if ! grep -q 'forged DNS responses from 172.28.0.11' <<<"$client_logs"; then + echo "--- Client did not log labeled forged milestone for dns-forge (172.28.0.11) ---" + echo "$client_logs" | tail -30 + echo "=== FAIL (missing per-resolver forged milestone log) ===" + exit 1 + fi + milestone_line=$(grep 'forged DNS responses from 172.28.0.11' <<<"$client_logs" | tail -1) + echo "--- Client milestone: $milestone_line ---" + echo "--- Tunnel delivers consistent responses despite forged NXDOMAINs ---" + echo "=== PASS ===" + exit 0 + fi + else + ok_count=0 + fi + printf "." + sleep 1 +done + +echo "" +echo "--- Tunnel did not come up through dns-good ---" +docker compose logs client server dns-good dns-forge +echo "=== FAIL ===" +exit 1 diff --git a/e2e/multi-resolver-health-detection/docker-compose.yml b/e2e/multi-resolver-health-detection/docker-compose.yml new file mode 100644 index 0000000..9a3f25c --- /dev/null +++ b/e2e/multi-resolver-health-detection/docker-compose.yml @@ -0,0 +1,83 @@ +networks: + dns-net: + ipam: + config: + - subnet: 172.28.0.0/24 + backend-net: + +volumes: + keys: + +services: + keygen: + build: + context: ../.. + dockerfile: Dockerfile + volumes: + - keys:/keys + command: > + sh -c "vaydns-server -gen-key -privkey-file /keys/server.key -pubkey-file /keys/server.pub" + + dns-good: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.10 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + dns-forge: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.11 + volumes: + - ../multi-resolver-forge/Corefile.forge:/Corefile + command: ["-conf", "/Corefile"] + + backend: + image: nginx:alpine + networks: + - backend-net + + server: + build: + context: ../.. + dockerfile: Dockerfile + networks: + dns-net: + ipv4_address: 172.28.0.20 + backend-net: + volumes: + - keys:/keys + command: > + vaydns-server -udp :53 -privkey-file /keys/server.key + -domain t.example.com -upstream backend:80 + -idle-timeout 10s -keepalive 2s + depends_on: + keygen: + condition: service_completed_successfully + dns-good: + condition: service_started + dns-forge: + condition: service_started + + client: + build: + context: ../.. + dockerfile: Dockerfile + networks: + - dns-net + volumes: + - keys:/keys + command: > + vaydns-client -udp 172.28.0.10:53 -udp 172.28.0.11:53 + -udp-shared-socket + -pubkey-file /keys/server.pub + -domain t.example.com -listen 0.0.0.0:7000 + -idle-timeout 10s -keepalive 2s -session-check-interval 500ms + -reconnect-min 1s -reconnect-max 5s -log-level info + depends_on: + server: + condition: service_started diff --git a/e2e/multi-resolver-health-detection/run.sh b/e2e/multi-resolver-health-detection/run.sh new file mode 100755 index 0000000..b2162b4 --- /dev/null +++ b/e2e/multi-resolver-health-detection/run.sh @@ -0,0 +1,99 @@ +#!/usr/bin/env bash +# Test: health state machine detects and excludes a forging resolver. +# +# dns-forge returns NXDOMAIN for every query. The client uses +# -udp-shared-socket so there is no per-query UDP worker filtering — +# forged NXDOMAIN responses reach evaluateIncoming directly, which is +# the code path affected by the per-response decay bug. +# +# With per-query UDP (the default), the UDP worker absorbs forged +# responses and detection happens through the pending-timeout path +# instead. That path works at high query rates even without the decay fix, +# so it cannot demonstrate the bug. Shared-socket mode is required to +# exercise the evaluateIncoming accumulation path. +# +# The health state machine decays penalty counters only on the 1/sec +# health tick, not per-response. This allows invalidCount to grow at +# the response rate (~several/sec) and cross the detection threshold +# within seconds. dns-forge should be marked Down and excluded from +# round-robin. During the measurement burst, dns-forge should receive +# only occasional probe queries, not 50% of all traffic. +set -euo pipefail +cd "$(dirname "$0")" + +cleanup() { docker compose down -v 2>/dev/null; } +trap cleanup EXIT + +fetch() { + docker compose exec -T client wget -q -O /dev/null http://localhost:7000 2>/dev/null +} + +query_count() { + # Count the number of DNS queries logged by a CoreDNS service. + docker compose logs "$1" 2>&1 | grep -c 't\.example\.com' || echo 0 +} + +echo "--- Building and starting services ---" +docker compose up -d --build + +echo "--- Waiting for tunnel (up to 30s) ---" +for i in $(seq 1 30); do + if fetch; then + echo "" + echo "--- Tunnel is up ---" + break + fi + if [ "$i" -eq 30 ]; then + echo "" + docker compose logs client server dns-good dns-forge + echo "=== FAIL (tunnel not ready) ===" + exit 1 + fi + printf "." + sleep 1 +done + +# Phase 1: let the health state machine observe the forging resolver. +# The pending-response timeout is 5s and the detection threshold is 8 +# timeouts / 5 invalids, so ~15s should be enough for the health machine +# to mark dns-forge as Down. +echo "--- Waiting 20s for health machine to observe dns-forge ---" +# Drive some traffic during the detection window so pending IDs accumulate. +for i in $(seq 1 20); do + fetch || true + sleep 1 +done + +# Phase 2: snapshot dns-forge query count, then drive a traffic burst +# and measure how many NEW queries dns-forge receives. +count_before=$(query_count dns-forge) +echo "--- dns-forge query count before burst: $count_before ---" + +echo "--- Driving 10 HTTP fetches through tunnel ---" +for i in $(seq 1 10); do + fetch || true + sleep 1 +done + +count_after=$(query_count dns-forge) +delta=$((count_after - count_before)) +echo "--- dns-forge query count after burst: $count_after (delta: $delta) ---" + +# Also check dns-good for comparison. +good_before_ignored=0 # not tracked, but log the total for debugging +good_total=$(query_count dns-good) +echo "--- dns-good total queries: $good_total ---" + +# dns-forge should receive very few queries during the burst (only probe +# queries at ~1/sec = ~10 max), not ~50% of all traffic. +max_allowed=15 +if [ "$delta" -gt "$max_allowed" ]; then + echo "--- dns-forge received $delta queries during the burst (max allowed: $max_allowed) ---" + echo "--- Health machine did not route away from the forging resolver (decay bug) ---" + docker compose logs client | tail -20 + echo "=== FAIL ===" + exit 1 +fi + +echo "--- dns-forge received only $delta queries during burst (≤$max_allowed); health machine routed away ---" +echo "=== PASS ===" diff --git a/e2e/multi-resolver-runtime-failure/docker-compose.yml b/e2e/multi-resolver-runtime-failure/docker-compose.yml new file mode 100644 index 0000000..d456182 --- /dev/null +++ b/e2e/multi-resolver-runtime-failure/docker-compose.yml @@ -0,0 +1,82 @@ +networks: + dns-net: + ipam: + config: + - subnet: 172.28.0.0/24 + backend-net: + +volumes: + keys: + +services: + keygen: + build: + context: ../.. + dockerfile: Dockerfile + volumes: + - keys:/keys + command: > + sh -c "vaydns-server -gen-key -privkey-file /keys/server.key -pubkey-file /keys/server.pub" + + dns1: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.10 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + dns2: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.11 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + backend: + image: nginx:alpine + networks: + - backend-net + + server: + build: + context: ../.. + dockerfile: Dockerfile + networks: + dns-net: + ipv4_address: 172.28.0.20 + backend-net: + volumes: + - keys:/keys + command: > + vaydns-server -udp :53 -privkey-file /keys/server.key + -domain t.example.com -upstream backend:80 + -idle-timeout 10s -keepalive 2s + depends_on: + keygen: + condition: service_completed_successfully + dns1: + condition: service_started + dns2: + condition: service_started + + client: + build: + context: ../.. + dockerfile: Dockerfile + networks: + - dns-net + volumes: + - keys:/keys + command: > + vaydns-client -udp 172.28.0.10:53 -udp 172.28.0.11:53 + -pubkey-file /keys/server.pub + -domain t.example.com -listen 0.0.0.0:7000 + -idle-timeout 10s -keepalive 2s -session-check-interval 500ms + -reconnect-min 1s -reconnect-max 5s -log-level info + depends_on: + server: + condition: service_started diff --git a/e2e/multi-resolver-runtime-failure/run.sh b/e2e/multi-resolver-runtime-failure/run.sh new file mode 100755 index 0000000..a863b79 --- /dev/null +++ b/e2e/multi-resolver-runtime-failure/run.sh @@ -0,0 +1,82 @@ +#!/usr/bin/env bash +# Test: one of several UDP resolvers dies mid-flight. +# Start the tunnel with two working DNS resolvers, verify HTTP traffic works, +# kill one resolver, and verify traffic continues through the remaining one. +# +# A failing UDP resolver does not surface an error to MultiResolver (the +# per-query worker retries silently), so this test validates that the health +# state machine + round-robin selection routes around it without tearing +# down the tunnel session. +set -euo pipefail +cd "$(dirname "$0")" + +cleanup() { docker compose down -v 2>/dev/null; } +trap cleanup EXIT + +fetch() { + docker compose exec -T client wget -q -O- http://localhost:7000 2>/dev/null | grep -q "Welcome to nginx" +} + +echo "--- Building and starting services ---" +docker compose up -d --build + +echo "--- Waiting for initial tunnel (up to 30s) ---" +for i in $(seq 1 30); do + if fetch; then + echo "" + echo "--- Initial tunnel is up ---" + break + fi + if [ "$i" -eq 30 ]; then + echo "" + docker compose logs client server dns1 dns2 + echo "=== FAIL (initial tunnel not ready) ===" + exit 1 + fi + printf "." + sleep 1 +done + +# Snapshot the session id before the kill, to detect any reconnect later. +pre_kill_sessions=$(docker compose logs client 2>&1 | grep -c 'session .* ready' || true) + +echo "--- Killing dns1 (half the queries will start dropping) ---" +docker compose kill dns1 + +# Give the client a moment to notice and start routing around dns1. +sleep 3 + +echo "--- Verifying tunnel still works through dns2 (up to 45s) ---" +ok_count=0 +for i in $(seq 1 45); do + if fetch; then + ok_count=$((ok_count + 1)) + # Require two consecutive successes so we don't declare victory on + # a lucky query that happened to go to dns2. + if [ "$ok_count" -ge 2 ]; then + post_kill_sessions=$(docker compose logs client 2>&1 | grep -c 'session .* ready' || true) + new_sessions=$((post_kill_sessions - pre_kill_sessions)) + if [ "$new_sessions" -gt 0 ]; then + echo "" + echo "--- Tunnel recovered but triggered $new_sessions new session(s) — resolver failure should be isolated without a full reconnect ---" + docker compose logs client | tail -30 + echo "=== FAIL (session was rebuilt instead of isolated) ===" + exit 1 + fi + echo "" + echo "--- Tunnel survived with 0 new sessions (single-entry failure isolated) ---" + echo "=== PASS ===" + exit 0 + fi + else + ok_count=0 + fi + printf "." + sleep 1 +done + +echo "" +echo "--- Tunnel did not survive dns1 kill ---" +docker compose logs client server dns1 dns2 +echo "=== FAIL ===" +exit 1 diff --git a/e2e/multi-resolver/docker-compose.yml b/e2e/multi-resolver/docker-compose.yml new file mode 100644 index 0000000..36f9222 --- /dev/null +++ b/e2e/multi-resolver/docker-compose.yml @@ -0,0 +1,82 @@ +networks: + dns-net: + ipam: + config: + - subnet: 172.28.0.0/24 + backend-net: + +volumes: + keys: + +services: + keygen: + build: + context: ../.. + dockerfile: Dockerfile + volumes: + - keys:/keys + command: > + sh -c "vaydns-server -gen-key -privkey-file /keys/server.key -pubkey-file /keys/server.pub" + + dns1: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.10 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + dns2: + image: coredns/coredns + networks: + dns-net: + ipv4_address: 172.28.0.11 + volumes: + - ../Corefile:/Corefile + command: ["-conf", "/Corefile"] + + backend: + image: nginx:alpine + networks: + - backend-net + + server: + build: + context: ../.. + dockerfile: Dockerfile + networks: + dns-net: + ipv4_address: 172.28.0.20 + backend-net: + volumes: + - keys:/keys + command: > + vaydns-server -udp :53 -privkey-file /keys/server.key + -domain t.example.com -upstream backend:80 + -idle-timeout 10s -keepalive 2s + depends_on: + keygen: + condition: service_completed_successfully + dns1: + condition: service_started + dns2: + condition: service_started + + client: + build: + context: ../.. + dockerfile: Dockerfile + networks: + - dns-net + volumes: + - keys:/keys + command: > + vaydns-client -udp 172.28.0.10:53 -udp 172.28.0.11:53 + -pubkey-file /keys/server.pub + -domain t.example.com -listen 0.0.0.0:7000 + -idle-timeout 10s -keepalive 2s -session-check-interval 500ms + -reconnect-min 1s -reconnect-max 5s + depends_on: + server: + condition: service_started diff --git a/e2e/multi-resolver/run.sh b/e2e/multi-resolver/run.sh new file mode 100755 index 0000000..c64921d --- /dev/null +++ b/e2e/multi-resolver/run.sh @@ -0,0 +1,29 @@ +#!/usr/bin/env bash +# Test: multi-resolver smoke test. +# Verifies that vaydns-client works when configured with two UDP resolvers, +# and that an HTTP request through the tunnel succeeds. +set -euo pipefail +cd "$(dirname "$0")" + +cleanup() { docker compose down -v 2>/dev/null; } +trap cleanup EXIT + +echo "--- Building and starting services ---" +docker compose up -d --build + +echo "--- Waiting for tunnel (up to 30s) ---" +for i in $(seq 1 30); do + if docker compose exec -T client wget -q -O- http://localhost:7000 2>/dev/null | grep -q "Welcome to nginx"; then + echo "" + echo "=== PASS ===" + exit 0 + fi + printf "." + sleep 1 +done + +echo "" +echo "--- Tunnel did not come up. Dumping logs ---" +docker compose logs client server dns1 dns2 +echo "=== FAIL ===" +exit 1 diff --git a/e2e/run-test.sh b/e2e/run-test.sh index b07c0b7..70eaf89 100755 --- a/e2e/run-test.sh +++ b/e2e/run-test.sh @@ -20,7 +20,7 @@ for rt in txt cname a aaaa mx ns srv; do fi done -for test_dir in socks-download recovery transport-recovery; do +for test_dir in socks-download recovery transport-recovery multi-resolver multi-resolver-runtime-failure multi-resolver-forge multi-resolver-health-detection multi-resolver-dnstt-compat; do total=$((total + 1)) echo "" echo "========================================" diff --git a/man/vaydns-client.1 b/man/vaydns-client.1 index 0b92691..d916fb4 100644 --- a/man/vaydns-client.1 +++ b/man/vaydns-client.1 @@ -13,7 +13,9 @@ .Sh SYNOPSIS .Nm -.Op Fl doh Ar URL | Fl dot Ar HOST : Ns Ar PORT | Fl udp Ar HOST : Ns Ar PORT +.Op Fl doh Ar URL ... +.Op Fl dot Ar HOST : Ns Ar PORT ... +.Op Fl udp Ar HOST : Ns Ar PORT ... .Op Fl pubkey Ar HEX | Fl pubkey-file Ar FILENAME .Fl domain Ar DOMAIN .Fl listen Ar LOCALADDR : Ns Ar LOCALPORT @@ -41,13 +43,19 @@ when a session fails. .Ss TRANSPORT -You must use exactly one of the +You must use at least one of the .Fl doh , .Fl dot , or .Fl udp -options, +options to specify what form of DNS to use: +each option may be repeated, and transport types may be mixed. +When multiple resolvers are configured, +.Nm +spreads queries across healthy resolvers, +routes around resolvers that stop returning usable responses, +and periodically probes unhealthy resolvers for recovery. .Bl -tag @@ -164,7 +172,7 @@ with Maximum concurrent streams per session. 0 means unlimited. Default: -.Cm 256 . +.Cm 0 . .It Fl open-stream-timeout Ar DURATION Timeout for opening an smux stream. @@ -226,7 +234,7 @@ Only applies when using without .Fl udp-shared-socket . Default: -.Cm 400ms . +.Cm 500ms . .It Fl udp-shared-socket Use a single shared UDP socket instead of per-query sockets. @@ -293,7 +301,7 @@ is set. .It Fl log-level Ar LEVEL Log level: debug, info, warning, error. Default: -.Cm warning . +.Cm info . .El @@ -314,6 +322,21 @@ vaydns-client -doh https://resolver.example/dns-query \e -listen 127.0.0.1:7000 .Ed +.Pp +Tunnel through multiple DNS over HTTPS and DNS over TLS resolvers. +Use +.Fl log-level +.Cm debug +to show the per-resolver health table. + +.Bd -literal -offset indent +vaydns-client -doh https://dns.google/dns-query \e + -doh https://cloudflare-dns.com/dns-query \e + -dot one.one.one.one:853 \e + -pubkey-file server.pub -domain t.example.com \e + -listen 127.0.0.1:7000 -log-level debug +.Ed + .Pp Tunnel through the DNS over TLS resolver at .Cm resolver.example:853 . diff --git a/vaydns-client/main.go b/vaydns-client/main.go index fc7e621..7ac3a94 100644 --- a/vaydns-client/main.go +++ b/vaydns-client/main.go @@ -2,7 +2,7 @@ // // Usage: // -// vaydns-client [-doh URL|-dot ADDR|-udp ADDR] [-pubkey HEX|-pubkey-file FILENAME] -domain DOMAIN -listen LOCALADDR +// vaydns-client [-doh URL]... [-dot ADDR]... [-udp ADDR]... [-pubkey HEX|-pubkey-file FILENAME] -domain DOMAIN -listen LOCALADDR package main import ( @@ -20,6 +20,17 @@ import ( log "github.com/sirupsen/logrus" ) +type StringSliceFlag []string + +func (s *StringSliceFlag) String() string { + return fmt.Sprint(*s) +} + +func (s *StringSliceFlag) Set(value string) error { + *s = append(*s, value) + return nil +} + var version = "dev" func readKeyFromFile(filename string) ([]byte, error) { @@ -33,13 +44,13 @@ func readKeyFromFile(filename string) ([]byte, error) { func main() { var showVersion bool - var dohURL string - var dotAddr string + var dohURLs StringSliceFlag + var dotAddrs StringSliceFlag var domainArg string var listenAddr string var pubkeyFilename string var pubkeyString string - var udpAddr string + var udpAddrs StringSliceFlag var utlsDistribution string var maxQnameLen int var maxNumLabels int @@ -64,11 +75,12 @@ func main() { flag.Usage = func() { fmt.Fprintf(flag.CommandLine.Output(), `Usage: - %[1]s [-doh URL|-dot ADDR|-udp ADDR] -pubkey-file PUBKEYFILE -domain DOMAIN -listen LOCALADDR + %[1]s [-doh URL]... [-dot ADDR]... [-udp ADDR]... -pubkey-file PUBKEYFILE -domain DOMAIN -listen LOCALADDR Examples: %[1]s -doh https://resolver.example/dns-query -pubkey-file server.pub -domain t.example.com -listen 127.0.0.1:7000 %[1]s -dot resolver.example:853 -pubkey-file server.pub -domain t.example.com -listen 127.0.0.1:7000 + %[1]s -doh https://dns.google/dns-query -doh https://cloudflare-dns.com/dns-query -dot one.one.one.one:853 -pubkey-file server.pub -domain t.example.com -listen 127.0.0.1:7000 `, os.Args[0]) flag.CommandLine.VisitAll(func(f *flag.Flag) { @@ -110,11 +122,11 @@ Known TLS fingerprints for -utls are: fmt.Fprintln(flag.CommandLine.Output(), line.String()) } } - flag.StringVar(&dohURL, "doh", "", "URL of DoH resolver") - flag.StringVar(&dotAddr, "dot", "", "address of DoT resolver") + flag.Var(&dohURLs, "doh", "URL of DoH resolver (repeatable)") + flag.Var(&dotAddrs, "dot", "address of DoT resolver (repeatable)") flag.StringVar(&pubkeyString, "pubkey", "", fmt.Sprintf("server public key (%d hex digits)", noise.KeyLen*2)) flag.StringVar(&pubkeyFilename, "pubkey-file", "", "read server public key from file") - flag.StringVar(&udpAddr, "udp", "", "address of UDP DNS resolver") + flag.Var(&udpAddrs, "udp", "address of UDP DNS resolver (repeatable)") flag.StringVar(&utlsDistribution, "utls", "4*random,3*Firefox_120,1*Firefox_105,3*Chrome_120,1*Chrome_102,1*iOS_14,1*iOS_13", "choose TLS fingerprint from weighted distribution") @@ -210,32 +222,44 @@ Known TLS fingerprints for -utls are: if utlsClientHelloID != nil { log.Infof("uTLS fingerprint %s %s", utlsClientHelloID.Client, utlsClientHelloID.Version) } - - // Select resolver transport. - var resolverType client.ResolverType - var resolverAddr string - transportCount := 0 - if dohURL != "" { - resolverType = client.ResolverTypeDOH - resolverAddr = dohURL - transportCount++ - } - if dotAddr != "" { - resolverType = client.ResolverTypeDOT - resolverAddr = dotAddr - transportCount++ - } - if udpAddr != "" { - resolverType = client.ResolverTypeUDP - resolverAddr = udpAddr - transportCount++ - } - if transportCount == 0 { - fmt.Fprintf(os.Stderr, "one of -doh, -dot, or -udp is required\n") + udpTimeout, err := time.ParseDuration(udpTimeoutStr) + if err != nil { + fmt.Fprintf(os.Stderr, "invalid -udp-timeout: %v\n", err) os.Exit(1) } - if transportCount > 1 { - fmt.Fprintf(os.Stderr, "only one of -doh, -dot, and -udp may be given\n") + // Select resolver transport. + resolvers := make([]client.Resolver, 0) + + for _, dohURL := range dohURLs { + resolver := client.Resolver{ + ResolverType: client.ResolverTypeDOH, + ResolverAddr: dohURL, + } + resolver.UTLSClientHelloID = utlsClientHelloID + resolvers = append(resolvers, resolver) + } + for _, dotAddr := range dotAddrs { + resolver := client.Resolver{ + ResolverType: client.ResolverTypeDOT, + ResolverAddr: dotAddr, + } + resolver.UTLSClientHelloID = utlsClientHelloID + resolvers = append(resolvers, resolver) + } + for _, udpAddr := range udpAddrs { + resolver := client.Resolver{ + ResolverType: client.ResolverTypeUDP, + ResolverAddr: udpAddr, + } + resolver.UDPWorkers = udpWorkers + resolver.UDPSharedSocket = udpSharedSocket + resolver.UDPTimeout = udpTimeout + resolver.UDPAcceptErrors = udpAcceptErrors + resolvers = append(resolvers, resolver) + } + + if len(resolvers) == 0 { + fmt.Fprintf(os.Stderr, "one of -doh, -dot, or -udp is required\n") os.Exit(1) } @@ -270,11 +294,6 @@ Known TLS fingerprints for -utls are: fmt.Fprintf(os.Stderr, "invalid -open-stream-timeout: %v\n", err) os.Exit(1) } - udpTimeout, err := time.ParseDuration(udpTimeoutStr) - if err != nil { - fmt.Fprintf(os.Stderr, "invalid -udp-timeout: %v\n", err) - os.Exit(1) - } // Validate. if keepAlive >= idleTimeout { @@ -348,16 +367,7 @@ Known TLS fingerprints for -utls are: } // Build resolver. - resolver, err := client.NewResolver(resolverType, resolverAddr) - if err != nil { - fmt.Fprintf(os.Stderr, "resolver: %v\n", err) - os.Exit(1) - } - resolver.UTLSClientHelloID = utlsClientHelloID - resolver.UDPWorkers = udpWorkers - resolver.UDPSharedSocket = udpSharedSocket - resolver.UDPTimeout = udpTimeout - resolver.UDPAcceptErrors = udpAcceptErrors + if udpAcceptErrors { if udpSharedSocket { log.Warnf("-udp-accept-errors has no effect when -udp-shared-socket is set") @@ -380,7 +390,7 @@ Known TLS fingerprints for -utls are: ts.RecordType = recordTypeStr // Build tunnel. - tunnel, err := client.NewTunnel(resolver, ts) + tunnel, err := client.NewTunnelMulti(resolvers, ts) if err != nil { fmt.Fprintf(os.Stderr, "%v\n", err) os.Exit(1)