This commit was merged in pull request #12.
This commit is contained in:
@@ -328,3 +328,156 @@ func TestQueryNoTCPFallbackWhenDisabled(t *testing.T) {
|
||||
t.Error("expected truncated response to be returned as-is")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeSuccess(t *testing.T) {
|
||||
answerResp := new(dns.Msg)
|
||||
answerResp.SetReply(new(dns.Msg))
|
||||
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
if msg.RecursionDesired {
|
||||
t.Error("IterativeQuery should send RD=false")
|
||||
}
|
||||
return answerResp.Copy(), nil
|
||||
}
|
||||
|
||||
server := net.ParseIP("198.41.0.4")
|
||||
resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, nil, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(resp.Answer) == 0 {
|
||||
t.Fatal("expected answer records")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeRetry(t *testing.T) {
|
||||
callCount := 0
|
||||
answerResp := new(dns.Msg)
|
||||
answerResp.SetReply(new(dns.Msg))
|
||||
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
callCount++
|
||||
if callCount < 2 {
|
||||
return nil, errors.New("transient error")
|
||||
}
|
||||
return answerResp.Copy(), nil
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{UDPSize: 2048, Retries: 3, AllowTCP: true}
|
||||
server := net.ParseIP("198.41.0.4")
|
||||
resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if resp == nil {
|
||||
t.Fatal("expected non-nil response after retry")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeAllFail(t *testing.T) {
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, errors.New("server unreachable")
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false}
|
||||
server := net.ParseIP("198.41.0.4")
|
||||
_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn)
|
||||
if err == nil {
|
||||
t.Fatal("expected error when all attempts fail")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeTCPFallback(t *testing.T) {
|
||||
truncatedResp := new(dns.Msg)
|
||||
truncatedResp.SetReply(new(dns.Msg))
|
||||
truncatedResp.Truncated = true
|
||||
|
||||
fullResp := new(dns.Msg)
|
||||
fullResp.SetReply(new(dns.Msg))
|
||||
fullResp.Answer = append(fullResp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
if !useTCP {
|
||||
return truncatedResp.Copy(), nil
|
||||
}
|
||||
return fullResp.Copy(), nil
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{UDPSize: 2048, Retries: 1, AllowTCP: true}
|
||||
server := net.ParseIP("198.41.0.4")
|
||||
resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(resp.Answer) == 0 {
|
||||
t.Fatal("expected answer after TCP fallback")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeContextCancelled(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// Cancel the context immediately so the retry loop aborts during backoff
|
||||
cancel()
|
||||
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
return nil, errors.New("error")
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{UDPSize: 2048, Retries: 5, AllowTCP: false}
|
||||
server := net.ParseIP("198.41.0.4")
|
||||
_, err := IterativeQueryWithExchange(ctx, server, "example.com", dns.TypeA, cfg, exchangeFn)
|
||||
if err == nil {
|
||||
t.Fatal("expected error when context cancelled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeNilResponse(t *testing.T) {
|
||||
callCount := 0
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
callCount++
|
||||
return nil, nil // nil response, no error
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false}
|
||||
server := net.ParseIP("198.41.0.4")
|
||||
_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for nil responses")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIterativeQueryWithExchangeUseTCP(t *testing.T) {
|
||||
var wasTCP bool
|
||||
answerResp := new(dns.Msg)
|
||||
answerResp.SetReply(new(dns.Msg))
|
||||
answerResp.Answer = append(answerResp.Answer, &dns.A{
|
||||
Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300},
|
||||
A: net.ParseIP("1.2.3.4"),
|
||||
})
|
||||
|
||||
exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
|
||||
wasTCP = useTCP
|
||||
return answerResp.Copy(), nil
|
||||
}
|
||||
|
||||
cfg := &QueryConfig{UDPSize: 2048, Retries: 1, UseTCP: true}
|
||||
server := net.ParseIP("198.41.0.4")
|
||||
_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !wasTCP {
|
||||
t.Error("expected TCP exchange when UseTCP=true")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user