refactor(netx/resolver): add CloseIdleConnections to RoundTripper (#501)

While there, also change to pointer receiver and use internal
testing for what are clearly unit tests.

Part of https://github.com/ooni/probe/issues/1591.
This commit is contained in:
Simone Basso 2021-09-09 20:49:12 +02:00 committed by GitHub
commit 1eb9e8c9b0
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
13 changed files with 203 additions and 115 deletions

View file

@ -1,17 +1,15 @@
package resolver_test
package resolver
import (
"context"
"errors"
"net"
"testing"
"github.com/ooni/probe-cli/v3/internal/engine/netx/resolver"
)
func TestDNSOverTCPTransportQueryTooLarge(t *testing.T) {
const address = "9.9.9.9:53"
txp := resolver.NewDNSOverTCP(new(net.Dialer).DialContext, address)
txp := NewDNSOverTCP(new(net.Dialer).DialContext, address)
reply, err := txp.RoundTrip(context.Background(), make([]byte, 1<<18))
if err == nil {
t.Fatal("expected an error here")
@ -24,8 +22,8 @@ func TestDNSOverTCPTransportQueryTooLarge(t *testing.T) {
func TestDNSOverTCPTransportDialFailure(t *testing.T) {
const address = "9.9.9.9:53"
mocked := errors.New("mocked error")
fakedialer := resolver.FakeDialer{Err: mocked}
txp := resolver.NewDNSOverTCP(fakedialer.DialContext, address)
fakedialer := FakeDialer{Err: mocked}
txp := NewDNSOverTCP(fakedialer.DialContext, address)
reply, err := txp.RoundTrip(context.Background(), make([]byte, 1<<11))
if !errors.Is(err, mocked) {
t.Fatal("not the error we expected")
@ -38,10 +36,10 @@ func TestDNSOverTCPTransportDialFailure(t *testing.T) {
func TestDNSOverTCPTransportSetDealineFailure(t *testing.T) {
const address = "9.9.9.9:53"
mocked := errors.New("mocked error")
fakedialer := resolver.FakeDialer{Conn: &resolver.FakeConn{
fakedialer := FakeDialer{Conn: &FakeConn{
SetDeadlineError: mocked,
}}
txp := resolver.NewDNSOverTCP(fakedialer.DialContext, address)
txp := NewDNSOverTCP(fakedialer.DialContext, address)
reply, err := txp.RoundTrip(context.Background(), make([]byte, 1<<11))
if !errors.Is(err, mocked) {
t.Fatal("not the error we expected")
@ -54,10 +52,10 @@ func TestDNSOverTCPTransportSetDealineFailure(t *testing.T) {
func TestDNSOverTCPTransportWriteFailure(t *testing.T) {
const address = "9.9.9.9:53"
mocked := errors.New("mocked error")
fakedialer := resolver.FakeDialer{Conn: &resolver.FakeConn{
fakedialer := FakeDialer{Conn: &FakeConn{
WriteError: mocked,
}}
txp := resolver.NewDNSOverTCP(fakedialer.DialContext, address)
txp := NewDNSOverTCP(fakedialer.DialContext, address)
reply, err := txp.RoundTrip(context.Background(), make([]byte, 1<<11))
if !errors.Is(err, mocked) {
t.Fatal("not the error we expected")
@ -70,10 +68,10 @@ func TestDNSOverTCPTransportWriteFailure(t *testing.T) {
func TestDNSOverTCPTransportReadFailure(t *testing.T) {
const address = "9.9.9.9:53"
mocked := errors.New("mocked error")
fakedialer := resolver.FakeDialer{Conn: &resolver.FakeConn{
fakedialer := FakeDialer{Conn: &FakeConn{
ReadError: mocked,
}}
txp := resolver.NewDNSOverTCP(fakedialer.DialContext, address)
txp := NewDNSOverTCP(fakedialer.DialContext, address)
reply, err := txp.RoundTrip(context.Background(), make([]byte, 1<<11))
if !errors.Is(err, mocked) {
t.Fatal("not the error we expected")
@ -86,11 +84,11 @@ func TestDNSOverTCPTransportReadFailure(t *testing.T) {
func TestDNSOverTCPTransportSecondReadFailure(t *testing.T) {
const address = "9.9.9.9:53"
mocked := errors.New("mocked error")
fakedialer := resolver.FakeDialer{Conn: &resolver.FakeConn{
fakedialer := FakeDialer{Conn: &FakeConn{
ReadError: mocked,
ReadData: []byte{byte(0), byte(2)},
}}
txp := resolver.NewDNSOverTCP(fakedialer.DialContext, address)
txp := NewDNSOverTCP(fakedialer.DialContext, address)
reply, err := txp.RoundTrip(context.Background(), make([]byte, 1<<11))
if !errors.Is(err, mocked) {
t.Fatal("not the error we expected")
@ -103,11 +101,11 @@ func TestDNSOverTCPTransportSecondReadFailure(t *testing.T) {
func TestDNSOverTCPTransportAllGood(t *testing.T) {
const address = "9.9.9.9:53"
mocked := errors.New("mocked error")
fakedialer := resolver.FakeDialer{Conn: &resolver.FakeConn{
fakedialer := FakeDialer{Conn: &FakeConn{
ReadError: mocked,
ReadData: []byte{byte(0), byte(1), byte(1)},
}}
txp := resolver.NewDNSOverTCP(fakedialer.DialContext, address)
txp := NewDNSOverTCP(fakedialer.DialContext, address)
reply, err := txp.RoundTrip(context.Background(), make([]byte, 1<<11))
if err != nil {
t.Fatal(err)
@ -119,7 +117,7 @@ func TestDNSOverTCPTransportAllGood(t *testing.T) {
func TestDNSOverTCPTransportOK(t *testing.T) {
const address = "9.9.9.9:53"
txp := resolver.NewDNSOverTCP(new(net.Dialer).DialContext, address)
txp := NewDNSOverTCP(new(net.Dialer).DialContext, address)
if txp.RequiresPadding() != false {
t.Fatal("invalid RequiresPadding")
}
@ -133,7 +131,7 @@ func TestDNSOverTCPTransportOK(t *testing.T) {
func TestDNSOverTLSTransportOK(t *testing.T) {
const address = "9.9.9.9:853"
txp := resolver.NewDNSOverTLS(resolver.DialTLSContext, address)
txp := NewDNSOverTLS(DialTLSContext, address)
if txp.RequiresPadding() != true {
t.Fatal("invalid RequiresPadding")
}