package discovery import ( "log" "net" "strings" "testing" "time" "github.com/miekg/dns" ) func TestDNSDiscovery_Interception(t *testing.T) { serviceIP := "192.0.2.100" upstreamDNS := []string{"8.8.8.8"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") // Test intercepting Bose service m := new(dns.Msg) m.SetQuestion("api.bose.com.", dns.TypeA) rw := &mockResponseWriter{} d.ServeDNS(rw, m) if rw.msg == nil { t.Fatal("Expected a response message, got nil") } if len(rw.msg.Answer) == 0 { t.Fatal("Expected an answer in the response") } if a, ok := rw.msg.Answer[0].(*dns.A); ok { if a.A.String() != serviceIP { t.Errorf("Expected intercepted IP %s, got %s", serviceIP, a.A.String()) } } else { t.Errorf("Expected A record, got %T", rw.msg.Answer[0]) } // Test intercepting streamingoauth.bose.com if !d.shouldIntercept("streamingoauth.bose.com") { t.Error("Expected streamingoauth.bose.com to be intercepted") } // Test aftertouch.test m2 := new(dns.Msg) m2.SetQuestion("aftertouch.test.", dns.TypeA) rw2 := &mockResponseWriter{} d.ServeDNS(rw2, m2) if rw2.msg == nil || len(rw2.msg.Answer) == 0 { t.Fatal("Expected response for aftertouch.test") } if a, ok := rw2.msg.Answer[0].(*dns.A); ok { if a.A.String() != serviceIP { t.Errorf("Expected intercepted IP %s for aftertouch.test, got %s", serviceIP, a.A.String()) } } else { t.Errorf("Expected A record for aftertouch.test, got %T", rw2.msg.Answer[0]) } } func TestDNSDiscovery_Forwarding(t *testing.T) { // This test is harder because it needs a real upstream or a mock. // For now, let's just test that it calls forward and record. serviceIP := "192.0.2.100" upstreamDNS := []string{"127.0.0.1:5353"} // Use a port that is likely closed or we can mock d := NewDNSDiscovery(upstreamDNS, serviceIP, "") m := new(dns.Msg) m.SetQuestion("google.com.", dns.TypeA) rw := &mockResponseWriter{} // Start a mock upstream DNS server mux := dns.NewServeMux() mux.HandleFunc("google.com.", func(w dns.ResponseWriter, r *dns.Msg) { m := new(dns.Msg) m.SetReply(r) _ = w.WriteMsg(m) }) ts := &dns.Server{Addr: "127.0.0.1:5353", Net: "udp", Handler: mux, ReadTimeout: 100 * time.Millisecond, WriteTimeout: 100 * time.Millisecond} go func() { _ = ts.ListenAndServe() }() defer func() { _ = ts.Shutdown() }() // Give it a moment to start time.Sleep(100 * time.Millisecond) // We expect forward to succeed d.ServeDNS(rw, m) d.mu.RLock() host, exists := d.discovered["google.com"] d.mu.RUnlock() if !exists { t.Error("Expected google.com to be recorded in discovery") } if host.IsBoseService { t.Error("google.com should not be identified as a Bose service") } } func TestDNSDiscovery_StartTCP(t *testing.T) { serviceIP := "192.0.2.100" upstreamDNS := []string{"8.8.8.8"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") addr := "127.0.0.1:5354" go func() { _ = d.Start(addr) }() // Give it a moment to start time.Sleep(200 * time.Millisecond) // Test TCP resolution m := new(dns.Msg) m.SetQuestion("api.bose.com.", dns.TypeA) c := new(dns.Client) c.Net = "tcp" in, _, err := c.Exchange(m, addr) if err != nil { t.Fatalf("Failed to exchange via TCP: %v", err) } if len(in.Answer) == 0 { t.Fatal("Expected answer in TCP response") } if a, ok := in.Answer[0].(*dns.A); ok { if a.A.String() != serviceIP { t.Errorf("Expected intercepted IP %s via TCP, got %s", serviceIP, a.A.String()) } } else { t.Errorf("Expected A record via TCP, got %T", in.Answer[0]) } // Test Shutdown err = d.Shutdown() if err != nil { t.Errorf("Shutdown failed: %v", err) } // Verify it's really shut down by trying to connect _, _, err = c.Exchange(m, addr) if err == nil { t.Error("Expected error after shutdown, but could still exchange") } } func TestDNSDiscovery_SelfForwarding(t *testing.T) { serviceIP := "soundtouch.local" upstreamDNS := []string{"127.0.0.1:5357"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") // Mock upstream DNS server for soundtouch.local mux := dns.NewServeMux() mux.HandleFunc("soundtouch.local.", func(w dns.ResponseWriter, r *dns.Msg) { m := new(dns.Msg) m.SetReply(r) rr, _ := dns.NewRR("soundtouch.local. 60 IN A 192.0.2.10") m.Answer = append(m.Answer, rr) _ = w.WriteMsg(m) }) ts := &dns.Server{Addr: "127.0.0.1:5357", Net: "udp", Handler: mux, ReadTimeout: 100 * time.Millisecond, WriteTimeout: 100 * time.Millisecond} go func() { _ = ts.ListenAndServe() }() defer func() { _ = ts.Shutdown() }() time.Sleep(100 * time.Millisecond) m := new(dns.Msg) m.SetQuestion("soundtouch.local.", dns.TypeA) rw := &mockResponseWriter{} d.ServeDNS(rw, m) if rw.msg == nil { t.Fatal("Expected a response for soundtouch.local") } if rw.msg.Rcode != dns.RcodeSuccess { t.Errorf("Expected Success (0) for soundtouch.local being forwarded, got %d", rw.msg.Rcode) } if len(rw.msg.Answer) == 0 { t.Fatal("Expected an answer in the response") } if a, ok := rw.msg.Answer[0].(*dns.A); ok { if a.A.String() != "192.0.2.10" { t.Errorf("Expected IP 192.0.2.10, got %s", a.A.String()) } } // Check if d.recordQuery logged it correctly. d.mu.RLock() host, exists := d.discovered["soundtouch.local"] d.mu.RUnlock() if !exists { t.Error("Expected soundtouch.local to be recorded") } // It should NOT be intercepted anymore if host != nil && host.IsIntercepted { t.Error("Expected soundtouch.local NOT to be intercepted anymore, but forwarded") } } func TestDNSDiscovery_ForwardLocal(t *testing.T) { serviceIP := "192.0.2.100" upstreamDNS := []string{"127.0.0.1:5356"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") m := new(dns.Msg) m.SetQuestion("someone-else.local.", dns.TypeA) rw := &mockResponseWriter{} // Start a mock upstream DNS server that returns SUCCESS for .local mux := dns.NewServeMux() mux.HandleFunc("someone-else.local.", func(w dns.ResponseWriter, r *dns.Msg) { m := new(dns.Msg) m.SetReply(r) rr, _ := dns.NewRR("someone-else.local. 60 IN A 192.0.2.50") m.Answer = append(m.Answer, rr) _ = w.WriteMsg(m) }) ts := &dns.Server{Addr: "127.0.0.1:5356", Net: "udp", Handler: mux, ReadTimeout: 100 * time.Millisecond, WriteTimeout: 100 * time.Millisecond} go func() { _ = ts.ListenAndServe() }() defer func() { _ = ts.Shutdown() }() time.Sleep(100 * time.Millisecond) d.ServeDNS(rw, m) if rw.msg == nil { t.Fatal("Expected a response message") } if rw.msg.Rcode != dns.RcodeSuccess { t.Errorf("Expected Success (0) for .local being forwarded, got %d", rw.msg.Rcode) } if len(rw.msg.Answer) == 0 { t.Fatal("Expected an answer in the response") } } func TestDNSDiscovery_IsRunning(t *testing.T) { serviceIP := "192.0.2.100" upstreamDNS := []string{"8.8.8.8"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") addr := "127.0.0.1:5355" if d.IsRunning(addr) { t.Error("Expected IsRunning to be false before Start") } go func() { _ = d.Start(addr) }() // Give it a moment to start time.Sleep(200 * time.Millisecond) if !d.IsRunning(addr) { t.Error("Expected IsRunning to be true after Start") } if d.IsRunning("127.0.0.1:9999") { t.Error("Expected IsRunning to be false for wrong address") } _ = d.Shutdown() if d.IsRunning(addr) { t.Error("Expected IsRunning to be false after Shutdown") } } type mockResponseWriter struct { msg *dns.Msg } func (m *mockResponseWriter) LocalAddr() net.Addr { return nil } func (m *mockResponseWriter) RemoteAddr() net.Addr { return nil } func (m *mockResponseWriter) WriteMsg(msg *dns.Msg) error { m.msg = msg; return nil } func (m *mockResponseWriter) Write([]byte) (int, error) { return 0, nil } func (m *mockResponseWriter) Close() error { return nil } func (m *mockResponseWriter) TsigStatus() error { return nil } func (m *mockResponseWriter) TsigTimersOnly(bool) {} func (m *mockResponseWriter) Hijack() {} func TestDNSDiscovery_LogThrottling(t *testing.T) { d := NewDNSDiscovery([]string{"8.8.8.8"}, "192.0.2.100", "") // Capture log output var logBuf strings.Builder oldOutput := log.Writer() log.SetOutput(&logBuf) defer log.SetOutput(oldOutput) msg := "Test log message" d.throttledLog(msg) d.throttledLog(msg) d.throttledLog(msg) count := strings.Count(logBuf.String(), msg) if count != 1 { t.Errorf("Expected log message to appear once due to throttling, but appeared %d times", count) } // Advance time by 11 seconds to bypass throttling d.lastLogMu.Lock() d.lastLog[msg] = time.Now().Add(-11 * time.Second) d.lastLogMu.Unlock() d.throttledLog(msg) count = strings.Count(logBuf.String(), msg) if count != 2 { t.Errorf("Expected log message to appear twice after advancing time, but appeared %d times", count) } } func TestDNSDiscovery_LoopPrevention(t *testing.T) { serviceIP := "192.0.2.100" bindAddr := "127.0.0.1:53" upstreamDNS := []string{"127.0.0.1:53"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") d.bindAddr = bindAddr // Capture log output to avoid panic if it's being throttled/logged var logBuf strings.Builder oldOutput := log.Writer() log.SetOutput(&logBuf) defer log.SetOutput(oldOutput) m := new(dns.Msg) m.SetQuestion("google.com.", dns.TypeA) rw := &mockResponseWriter{} d.forward(rw, m) if rw.msg == nil { t.Fatal("Expected a response message") } if rw.msg.Rcode != dns.RcodeServerFailure { t.Errorf("Expected RcodeServerFailure (2), got %d", rw.msg.Rcode) } } func TestDNSDiscovery_EmptyUpstream(t *testing.T) { serviceIP := "192.0.2.100" var upstreamDNS []string // Empty upstream d := NewDNSDiscovery(upstreamDNS, serviceIP, "") d.bindAddr = ":53" m := new(dns.Msg) m.SetQuestion("google.com.", dns.TypeA) rw := &mockResponseWriter{} d.ServeDNS(rw, m) if rw.msg == nil { t.Fatal("Expected a response message, got nil") } if rw.msg.Rcode != dns.RcodeServerFailure { t.Errorf("Expected RcodeServerFailure (2) for empty upstream, got %d", rw.msg.Rcode) } // Verify log message (optional, but good to check it's the simplified one) } func TestDNSDiscovery_ForwardTimeout(t *testing.T) { serviceIP := "192.0.2.100" // Mock server that deliberately delays its response mux := dns.NewServeMux() mux.HandleFunc("google.com.", func(w dns.ResponseWriter, r *dns.Msg) { time.Sleep(200 * time.Millisecond) // Longer than the timeout m := new(dns.Msg) m.SetReply(r) _ = w.WriteMsg(m) }) ts := &dns.Server{Addr: "127.0.0.1:5358", Net: "udp", Handler: mux} go func() { _ = ts.ListenAndServe() }() defer func() { _ = ts.Shutdown() }() // Give the server a moment to start time.Sleep(50 * time.Millisecond) upstreamDNS := []string{"127.0.0.1:5358"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") d.timeout = 100 * time.Millisecond m := new(dns.Msg) m.SetQuestion("google.com.", dns.TypeA) rw := &mockResponseWriter{} start := time.Now() d.forward(rw, m) duration := time.Since(start) // Since we're forwarding to a local server that sleeps for 200ms, // and our timeout is 100ms, it should take at least 100ms. if duration < 100*time.Millisecond { t.Errorf("Expected forward to take at least 100ms (timeout), but took %v", duration) } if rw.msg == nil || rw.msg.Rcode != dns.RcodeServerFailure { t.Errorf("Expected RcodeServerFailure after timeout, got msg: %v", rw.msg) } } func TestDNSDiscovery_MultipleUpstreams(t *testing.T) { serviceIP := "192.0.2.100" // Mock server 1: returns NXDOMAIN mux1 := dns.NewServeMux() mux1.HandleFunc("test.com.", func(w dns.ResponseWriter, r *dns.Msg) { m := new(dns.Msg) m.SetReply(r) m.Rcode = dns.RcodeNameError _ = w.WriteMsg(m) }) ts1 := &dns.Server{Addr: "127.0.0.1:5356", Net: "udp", Handler: mux1} go func() { _ = ts1.ListenAndServe() }() defer func() { _ = ts1.Shutdown() }() // Mock server 2: succeeds mux2 := dns.NewServeMux() mux2.HandleFunc("test.com.", func(w dns.ResponseWriter, r *dns.Msg) { m := new(dns.Msg) m.SetReply(r) m.Answer = append(m.Answer, &dns.A{ Hdr: dns.RR_Header{Name: r.Question[0].Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("1.2.3.4"), }) _ = w.WriteMsg(m) }) ts2 := &dns.Server{Addr: "127.0.0.1:5357", Net: "udp", Handler: mux2} go func() { _ = ts2.ListenAndServe() }() defer func() { _ = ts2.Shutdown() }() time.Sleep(100 * time.Millisecond) upstreamDNS := []string{"127.0.0.1:5356", "127.0.0.1:5357"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") m := new(dns.Msg) m.SetQuestion("test.com.", dns.TypeA) rw := &mockResponseWriter{} d.forward(rw, m) if rw.msg == nil { t.Fatal("Expected a response message") } // It should succeed because it falls back to the second upstream if rw.msg.Rcode != dns.RcodeSuccess { t.Errorf("Expected RcodeSuccess (0), got %d. Fallback failed.", rw.msg.Rcode) } if len(rw.msg.Answer) == 0 { t.Fatal("Expected an answer from the second upstream") } } func TestDNSDiscovery_HostnameServiceIP(t *testing.T) { // Use localhost which should resolve to 127.0.0.1 serviceIP := "localhost" upstreamDNS := []string{"8.8.8.8"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") m := new(dns.Msg) m.SetQuestion("api.bose.com.", dns.TypeA) rw := &mockResponseWriter{} d.ServeDNS(rw, m) if rw.msg == nil { t.Fatal("Expected a response message, got nil") } if len(rw.msg.Answer) == 0 { t.Fatal("Expected an answer in the response") } if a, ok := rw.msg.Answer[0].(*dns.A); ok { // It should be resolved to 127.0.0.1 (or whatever localhost resolves to) if a.A.String() == "" { t.Error("Expected a non-empty IP address") } log.Printf("Resolved localhost to %s", a.A.String()) } else if cname, ok := rw.msg.Answer[0].(*dns.CNAME); ok { // Fallback to CNAME is also acceptable if resolution failed but it shouldn't for localhost if cname.Target != "localhost." { t.Errorf("Expected CNAME to localhost., got %s", cname.Target) } } else { t.Errorf("Expected A or CNAME record, got %T", rw.msg.Answer[0]) } } func TestDNSDiscovery_UnresolvableHostname(t *testing.T) { // Use a likely unresolvable hostname serviceIP := "this.hostname.does.not.exist.at.all.invalid" upstreamDNS := []string{"8.8.8.8"} d := NewDNSDiscovery(upstreamDNS, serviceIP, "") m := new(dns.Msg) m.SetQuestion("api.bose.com.", dns.TypeA) rw := &mockResponseWriter{} d.ServeDNS(rw, m) if rw.msg == nil { t.Fatal("Expected a response message, got nil") } if len(rw.msg.Answer) == 0 { t.Fatal("Expected an answer in the response (CNAME fallback)") } if cname, ok := rw.msg.Answer[0].(*dns.CNAME); ok { expected := serviceIP + "." if cname.Target != expected { t.Errorf("Expected CNAME to %s, got %s", expected, cname.Target) } } else { t.Errorf("Expected CNAME record for unresolvable hostname, got %T", rw.msg.Answer[0]) } } // mockResponseWriterWithAddr is like mockResponseWriter but returns a // configurable remote address; used to simulate queries from speaker IPs. type mockResponseWriterWithAddr struct { mockResponseWriter remote net.Addr } func (m *mockResponseWriterWithAddr) RemoteAddr() net.Addr { return m.remote } // mockUDPAddr implements net.Addr for test purposes. type mockUDPAddr struct{ addr string } func (a *mockUDPAddr) Network() string { return "udp" } func (a *mockUDPAddr) String() string { return a.addr } func TestDNSDiscovery_InterceptClientTracking_NonLoopback(t *testing.T) { d := NewDNSDiscovery([]string{"8.8.8.8"}, "192.0.2.100", "") // Query for an intercepted Bose hostname from a non-loopback address. msg := new(dns.Msg) msg.SetQuestion("content.api.bose.io.", dns.TypeA) rw := &mockResponseWriterWithAddr{ remote: &mockUDPAddr{addr: "10.1.103.209:54321"}, } d.ServeDNS(rw, msg) clients := d.InterceptClientIPs() if _, ok := clients["10.1.103.209"]; !ok { t.Errorf("expected 10.1.103.209 in interceptClients, got %v", clients) } } func TestDNSDiscovery_InterceptClientTracking_LoopbackExcluded(t *testing.T) { d := NewDNSDiscovery([]string{"8.8.8.8"}, "192.0.2.100", "") // Query for an intercepted hostname from loopback (health-probe scenario). msg := new(dns.Msg) msg.SetQuestion("api.bose.com.", dns.TypeA) rw := &mockResponseWriterWithAddr{ remote: &mockUDPAddr{addr: "127.0.0.1:55069"}, } d.ServeDNS(rw, msg) clients := d.InterceptClientIPs() if len(clients) != 0 { t.Errorf("expected no clients (loopback should be excluded), got %v", clients) } } func TestDNSDiscovery_InterceptClientTracking_NonInterceptedNotRecorded(t *testing.T) { // Forwarded (non-intercepted) queries should not populate interceptClients. upstreamMux := dns.NewServeMux() upstreamMux.HandleFunc("google.com.", func(w dns.ResponseWriter, r *dns.Msg) { m := new(dns.Msg) m.SetReply(r) _ = w.WriteMsg(m) }) ts := &dns.Server{Addr: "127.0.0.1:5360", Net: "udp", Handler: upstreamMux, ReadTimeout: 100 * time.Millisecond, WriteTimeout: 100 * time.Millisecond} go func() { _ = ts.ListenAndServe() }() defer func() { _ = ts.Shutdown() }() time.Sleep(100 * time.Millisecond) d := NewDNSDiscovery([]string{"127.0.0.1:5360"}, "192.0.2.100", "") msg := new(dns.Msg) msg.SetQuestion("google.com.", dns.TypeA) rw := &mockResponseWriterWithAddr{ remote: &mockUDPAddr{addr: "10.1.103.209:54321"}, } d.ServeDNS(rw, msg) clients := d.InterceptClientIPs() if len(clients) != 0 { t.Errorf("expected no intercept clients for non-intercepted query, got %v", clients) } }