diff --git a/internal/cmd/jafar/README.md b/internal/cmd/jafar/README.md index a769753..8b4722a 100644 --- a/internal/cmd/jafar/README.md +++ b/internal/cmd/jafar/README.md @@ -156,16 +156,14 @@ response for every request whose `Host` contains the specified string. ### tls-proxy -TLS proxy is a TCP proxy that routes traffic to specific servers depending +TLS proxy is a proxy that routes traffic to specific servers depending on their SNI value. It is controlled by the following flags: ```bash -tls-proxy-address string - Address where the TCP+TLS proxy should listen (default "127.0.0.1:443") + Address where the HTTP proxy should listen (default "127.0.0.1:443") -tls-proxy-block value - Register SNI header keyword triggering TLS censorship - -tls-proxy-outbound-port - Define the outbound port requests are proxied to (default "443 for HTTPS) + Register keyword triggering TLS censorship ``` The `-tls-proxy-address` flags has the same semantics it has for the DNS diff --git a/internal/cmd/jafar/main.go b/internal/cmd/jafar/main.go index 7d376d8..61f5285 100644 --- a/internal/cmd/jafar/main.go +++ b/internal/cmd/jafar/main.go @@ -58,9 +58,8 @@ var ( tag *string - tlsProxyAddress *string - tlsProxyBlock flagx.StringArray - tlsProxyOutboundPort *string + tlsProxyAddress *string + tlsProxyBlock flagx.StringArray uncensoredResolverDoH *string ) @@ -160,16 +159,12 @@ func init() { // tlsProxy tlsProxyAddress = flag.String( "tls-proxy-address", "127.0.0.1:443", - "Address where the TCP+TLS proxy should listen", + "Address where the HTTP proxy should listen", ) flag.Var( &tlsProxyBlock, "tls-proxy-block", "Register keyword triggering TLS censorship", ) - tlsProxyOutboundPort = flag.String( - "tls-proxy-outbound-port", "443", - "The outbound port where requests should be proxied", - ) // uncensored uncensoredResolverDoH = flag.String( @@ -232,7 +227,7 @@ func iptablesStart() *iptables.CensoringPolicy { } func tlsProxyStart(uncensored *uncensored.Client) net.Listener { - proxy := tlsproxy.NewCensoringProxy(tlsProxyBlock, uncensored, tlsProxyOutboundPort) + proxy := tlsproxy.NewCensoringProxy(tlsProxyBlock, uncensored) listener, err := proxy.Start(*tlsProxyAddress) runtimex.PanicOnError(err, "proxy.Start failed") return listener diff --git a/internal/cmd/jafar/tlsproxy/tlsproxy.go b/internal/cmd/jafar/tlsproxy/tlsproxy.go index d98751f..91cede5 100644 --- a/internal/cmd/jafar/tlsproxy/tlsproxy.go +++ b/internal/cmd/jafar/tlsproxy/tlsproxy.go @@ -21,9 +21,8 @@ type Dialer interface { // CensoringProxy is a censoring TLS proxy type CensoringProxy struct { - keywords []string - dial func(network, address string) (net.Conn, error) - outboundPort string + keywords []string + dial func(network, address string) (net.Conn, error) } // NewCensoringProxy creates a new CensoringProxy instance using @@ -32,18 +31,13 @@ type CensoringProxy struct { // the SNII record of a ClientHello. dnsNetwork and dnsAddress are // settings to configure the upstream, non censored DNS. func NewCensoringProxy( - keywords []string, uncensored Dialer, outboundPort *string, + keywords []string, uncensored Dialer, ) *CensoringProxy { - defaultPort := "443" - if outboundPort == nil { - outboundPort = &defaultPort - } return &CensoringProxy{ keywords: keywords, dial: func(network, address string) (net.Conn, error) { return uncensored.DialContext(context.Background(), network, address) }, - outboundPort: *outboundPort, } } @@ -152,7 +146,7 @@ func (p *CensoringProxy) handle(clientconn net.Conn) { return } } - serverconn, err := p.dial("tcp", net.JoinHostPort(sni, p.outboundPort)) + serverconn, err := p.dial("tcp", net.JoinHostPort(sni, "443")) if err != nil { log.WithError(err).Warn("tlsproxy: p.dial failed") alertclose(clientconn) diff --git a/internal/cmd/jafar/tlsproxy/tlsproxy_test.go b/internal/cmd/jafar/tlsproxy/tlsproxy_test.go index 1394ba6..86a21d8 100644 --- a/internal/cmd/jafar/tlsproxy/tlsproxy_test.go +++ b/internal/cmd/jafar/tlsproxy/tlsproxy_test.go @@ -94,7 +94,7 @@ func TestFailWriteAfterConnect(t *testing.T) { func TestListenError(t *testing.T) { proxy := NewCensoringProxy( - []string{""}, uncensored.NewClient("https://1.1.1.1/dns-query"), nil, + []string{""}, uncensored.NewClient("https://1.1.1.1/dns-query"), ) listener, err := proxy.Start("8.8.8.8:80") if err == nil { @@ -107,7 +107,7 @@ func TestListenError(t *testing.T) { func newproxy(t *testing.T, blocked string) net.Listener { proxy := NewCensoringProxy( - []string{blocked}, uncensored.NewClient("https://1.1.1.1/dns-query"), nil, + []string{blocked}, uncensored.NewClient("https://1.1.1.1/dns-query"), ) listener, err := proxy.Start("127.0.0.1:0") if err != nil { diff --git a/internal/engine/experiment.go b/internal/engine/experiment.go index 325bb6f..f40f4c3 100644 --- a/internal/engine/experiment.go +++ b/internal/engine/experiment.go @@ -92,12 +92,7 @@ func (eaw *experimentAsyncWrapper) RunAsync( out := make(chan *model.ExperimentAsyncTestKeys) measurement := eaw.experiment.newMeasurement(input) start := time.Now() - args := &model.ExperimentArgs{ - Callbacks: eaw.callbacks, - Measurement: measurement, - Session: eaw.session, - } - err := eaw.experiment.measurer.Run(ctx, args) + err := eaw.experiment.measurer.Run(ctx, eaw.session, measurement, eaw.callbacks) stop := time.Now() if err != nil { return nil, err diff --git a/internal/engine/experiment/dash/dash.go b/internal/engine/experiment/dash/dash.go index 8489286..a44021d 100644 --- a/internal/engine/experiment/dash/dash.go +++ b/internal/engine/experiment/dash/dash.go @@ -249,10 +249,10 @@ func (m Measurer) ExperimentVersion() string { } // Run implements model.ExperimentMeasurer.Run. -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { tk := new(TestKeys) measurement.TestKeys = tk saver := &tracex.Saver{} diff --git a/internal/engine/experiment/dash/dash_test.go b/internal/engine/experiment/dash/dash_test.go index b2a910c..7b5addc 100644 --- a/internal/engine/experiment/dash/dash_test.go +++ b/internal/engine/experiment/dash/dash_test.go @@ -270,15 +270,15 @@ func TestMeasureWithCancelledContext(t *testing.T) { cancel() // cause failure measurement := new(model.Measurement) m := &Measurer{} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := m.Run( + ctx, + &mockable.Session{ MockableHTTPClient: http.DefaultClient, MockableLogger: log.Log, }, - } - err := m.Run(ctx, args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) // See corresponding comment in Measurer.Run implementation to // understand why here it's correct to return nil. if !errors.Is(err, nil) { diff --git a/internal/engine/experiment/dnscheck/dnscheck.go b/internal/engine/experiment/dnscheck/dnscheck.go index d9912cb..66b805d 100644 --- a/internal/engine/experiment/dnscheck/dnscheck.go +++ b/internal/engine/experiment/dnscheck/dnscheck.go @@ -120,11 +120,10 @@ var ( ) // Run implements model.ExperimentSession.Run -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { // 1. fill the measurement with test keys tk := new(TestKeys) tk.Lookups = make(map[string]urlgetter.TestKeys) diff --git a/internal/engine/experiment/dnscheck/dnscheck_test.go b/internal/engine/experiment/dnscheck/dnscheck_test.go index c217eb6..2bc1966 100644 --- a/internal/engine/experiment/dnscheck/dnscheck_test.go +++ b/internal/engine/experiment/dnscheck/dnscheck_test.go @@ -56,12 +56,12 @@ func TestExperimentNameAndVersion(t *testing.T) { func TestDNSCheckFailsWithoutInput(t *testing.T) { measurer := NewExperimentMeasurer(Config{Domain: "example.com"}) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: new(model.Measurement), - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + new(model.Measurement), + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, ErrInputRequired) { t.Fatal("expected no input error") } @@ -69,12 +69,12 @@ func TestDNSCheckFailsWithoutInput(t *testing.T) { func TestDNSCheckFailsWithInvalidURL(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &model.Measurement{Input: "Not a valid URL \x7f"}, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + &model.Measurement{Input: "Not a valid URL \x7f"}, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, ErrInvalidURL) { t.Fatal("expected invalid input error") } @@ -82,12 +82,12 @@ func TestDNSCheckFailsWithInvalidURL(t *testing.T) { func TestDNSCheckFailsWithUnsupportedProtocol(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &model.Measurement{Input: "file://1.1.1.1"}, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + &model.Measurement{Input: "file://1.1.1.1"}, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, ErrUnsupportedURLScheme) { t.Fatal("expected unsupported scheme error") } @@ -100,12 +100,12 @@ func TestWithCancelledContext(t *testing.T) { DefaultAddrs: "1.1.1.1 1.0.0.1", }) measurement := &model.Measurement{Input: "dot://one.one.one.one"} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: newsession(), - } - err := measurer.Run(ctx, args) + err := measurer.Run( + ctx, + newsession(), + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -147,12 +147,12 @@ func TestDNSCheckValid(t *testing.T) { DefaultAddrs: "1.1.1.1 1.0.0.1", }) measurement := model.Measurement{Input: "dot://one.one.one.one:853"} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &measurement, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + &measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatalf("unexpected error: %s", err.Error()) } @@ -195,12 +195,12 @@ func TestDNSCheckWait(t *testing.T) { measurer := &Measurer{Endpoints: endpoints} run := func(input string) { measurement := model.Measurement{Input: model.MeasurementTarget(input)} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &measurement, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + &measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatalf("unexpected error: %s", err.Error()) } diff --git a/internal/engine/experiment/dnsping/dnsping.go b/internal/engine/experiment/dnsping/dnsping.go index cef7a93..648b323 100644 --- a/internal/engine/experiment/dnsping/dnsping.go +++ b/internal/engine/experiment/dnsping/dnsping.go @@ -85,10 +85,12 @@ var ( ) // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { if measurement.Input == "" { return errNoInputProvided } diff --git a/internal/engine/experiment/dnsping/dnsping_test.go b/internal/engine/experiment/dnsping/dnsping_test.go index 0a9ef50..3255a31 100644 --- a/internal/engine/experiment/dnsping/dnsping_test.go +++ b/internal/engine/experiment/dnsping/dnsping_test.go @@ -61,12 +61,7 @@ func TestMeasurer_run(t *testing.T) { MockableLogger: model.DiscardLogger, } callbacks := model.NewPrinterCallbacks(model.DiscardLogger) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: meas, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, meas, callbacks) return meas, m, err } diff --git a/internal/engine/experiment/example/example.go b/internal/engine/experiment/example/example.go index c2c1c01..1a7f329 100644 --- a/internal/engine/experiment/example/example.go +++ b/internal/engine/experiment/example/example.go @@ -57,10 +57,10 @@ func (m Measurer) ExperimentVersion() string { var ErrFailure = errors.New("mocked error") // Run implements model.ExperimentMeasurer.Run. -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { var err error if m.config.ReturnError { err = ErrFailure diff --git a/internal/engine/experiment/example/example_test.go b/internal/engine/experiment/example/example_test.go index 29c28fa..dc7e218 100644 --- a/internal/engine/experiment/example/example_test.go +++ b/internal/engine/experiment/example/example_test.go @@ -26,12 +26,7 @@ func TestSuccess(t *testing.T) { sess := &mockable.Session{MockableLogger: log.Log} callbacks := model.NewPrinterCallbacks(sess.Logger()) measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -52,12 +47,7 @@ func TestFailure(t *testing.T) { ctx := context.Background() sess := &mockable.Session{MockableLogger: log.Log} callbacks := model.NewPrinterCallbacks(sess.Logger()) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: new(model.Measurement), - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, new(model.Measurement), callbacks) if !errors.Is(err, example.ErrFailure) { t.Fatal("expected an error here") } diff --git a/internal/engine/experiment/fbmessenger/fbmessenger.go b/internal/engine/experiment/fbmessenger/fbmessenger.go index 3901845..49bcbeb 100644 --- a/internal/engine/experiment/fbmessenger/fbmessenger.go +++ b/internal/engine/experiment/fbmessenger/fbmessenger.go @@ -157,10 +157,10 @@ func (m Measurer) ExperimentVersion() string { } // Run implements ExperimentMeasurer.Run -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { ctx, cancel := context.WithTimeout(ctx, 60*time.Second) defer cancel() urlgetter.RegisterExtensions(measurement) diff --git a/internal/engine/experiment/fbmessenger/fbmessenger_test.go b/internal/engine/experiment/fbmessenger/fbmessenger_test.go index fa631c3..0545cd4 100644 --- a/internal/engine/experiment/fbmessenger/fbmessenger_test.go +++ b/internal/engine/experiment/fbmessenger/fbmessenger_test.go @@ -35,12 +35,7 @@ func TestSuccess(t *testing.T) { sess := newsession(t) measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -102,12 +97,7 @@ func TestWithCancelledContext(t *testing.T) { sess := &mockable.Session{MockableLogger: log.Log} measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } diff --git a/internal/engine/experiment/hhfm/hhfm.go b/internal/engine/experiment/hhfm/hhfm.go index 8518432..574cfbd 100644 --- a/internal/engine/experiment/hhfm/hhfm.go +++ b/internal/engine/experiment/hhfm/hhfm.go @@ -90,10 +90,10 @@ var ( ) // Run implements ExperimentMeasurer.Run. -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { ctx, cancel := context.WithTimeout(ctx, 30*time.Second) defer cancel() urlgetter.RegisterExtensions(measurement) diff --git a/internal/engine/experiment/hhfm/hhfm_test.go b/internal/engine/experiment/hhfm/hhfm_test.go index 157ab18..4d95f70 100644 --- a/internal/engine/experiment/hhfm/hhfm_test.go +++ b/internal/engine/experiment/hhfm/hhfm_test.go @@ -45,12 +45,7 @@ func TestSuccess(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -158,12 +153,7 @@ func TestCancelledContext(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -269,12 +259,7 @@ func TestNoHelpers(t *testing.T) { sess := &mockable.Session{} measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, hhfm.ErrNoAvailableTestHelpers) { t.Fatal("not the error we expected") } @@ -324,12 +309,7 @@ func TestNoActualHelpersInList(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, hhfm.ErrNoAvailableTestHelpers) { t.Fatal("not the error we expected") } @@ -382,12 +362,7 @@ func TestWrongTestHelperType(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, hhfm.ErrInvalidHelperType) { t.Fatal("not the error we expected") } @@ -440,12 +415,7 @@ func TestNewRequestFailure(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err == nil || !strings.HasSuffix(err.Error(), "invalid control character in URL") { t.Fatal("not the error we expected") } @@ -502,12 +472,7 @@ func TestInvalidJSONBody(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } diff --git a/internal/engine/experiment/hirl/hirl.go b/internal/engine/experiment/hirl/hirl.go index e4931d9..6cc3517 100644 --- a/internal/engine/experiment/hirl/hirl.go +++ b/internal/engine/experiment/hirl/hirl.go @@ -78,10 +78,10 @@ var ( ) // Run implements ExperimentMeasurer.Run. -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { tk := new(TestKeys) measurement.TestKeys = tk if len(m.Methods) < 1 { diff --git a/internal/engine/experiment/hirl/hirl_test.go b/internal/engine/experiment/hirl/hirl_test.go index 40fec23..1299307 100644 --- a/internal/engine/experiment/hirl/hirl_test.go +++ b/internal/engine/experiment/hirl/hirl_test.go @@ -42,12 +42,7 @@ func TestSuccess(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -96,12 +91,7 @@ func TestCancelledContext(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -200,12 +190,7 @@ func TestWithFakeMethods(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -266,12 +251,7 @@ func TestWithNoMethods(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, hirl.ErrNoMeasurementMethod) { t.Fatal("not the error we expected") } @@ -299,12 +279,7 @@ func TestNoHelpers(t *testing.T) { sess := &mockable.Session{} measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, hirl.ErrNoAvailableTestHelpers) { t.Fatal("not the error we expected") } @@ -336,12 +311,7 @@ func TestNoActualHelperInList(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, hirl.ErrNoAvailableTestHelpers) { t.Fatal("not the error we expected") } @@ -376,12 +346,7 @@ func TestWrongTestHelperType(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, hirl.ErrInvalidHelperType) { t.Fatal("not the error we expected") } diff --git a/internal/engine/experiment/httphostheader/httphostheader.go b/internal/engine/experiment/httphostheader/httphostheader.go index 0a6a5b4..70e4087 100644 --- a/internal/engine/experiment/httphostheader/httphostheader.go +++ b/internal/engine/experiment/httphostheader/httphostheader.go @@ -46,10 +46,12 @@ func (m *Measurer) ExperimentVersion() string { } // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { if measurement.Input == "" { return errors.New("experiment requires input") } diff --git a/internal/engine/experiment/httphostheader/httphostheader_test.go b/internal/engine/experiment/httphostheader/httphostheader_test.go index 2c9c556..efd88ce 100644 --- a/internal/engine/experiment/httphostheader/httphostheader_test.go +++ b/internal/engine/experiment/httphostheader/httphostheader_test.go @@ -30,12 +30,12 @@ func TestMeasurerMeasureNoMeasurementInput(t *testing.T) { measurer := NewExperimentMeasurer(Config{ TestHelperURL: "http://www.google.com", }) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &model.Measurement{}, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + new(model.Measurement), + model.NewPrinterCallbacks(log.Log), + ) if err == nil || err.Error() != "experiment requires input" { t.Fatal("not the error we expected") } @@ -44,12 +44,12 @@ func TestMeasurerMeasureNoMeasurementInput(t *testing.T) { func TestMeasurerMeasureNoTestHelper(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) measurement := &model.Measurement{Input: "x.org"} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -75,12 +75,12 @@ func TestRunnerHTTPSetHostHeader(t *testing.T) { measurement := &model.Measurement{ Input: "x.org", } - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + measurement, + model.NewPrinterCallbacks(log.Log), + ) if host != "x.org" { t.Fatal("not the host we expected") } diff --git a/internal/engine/experiment/imap/.smtp.go.swp b/internal/engine/experiment/imap/.smtp.go.swp new file mode 100644 index 0000000..c4f8bb6 Binary files /dev/null and b/internal/engine/experiment/imap/.smtp.go.swp differ diff --git a/internal/engine/experiment/imap/.smtp_test.go.swp b/internal/engine/experiment/imap/.smtp_test.go.swp new file mode 100644 index 0000000..5a742be Binary files /dev/null and b/internal/engine/experiment/imap/.smtp_test.go.swp differ diff --git a/internal/engine/experiment/imap/imap.go b/internal/engine/experiment/imap/imap.go new file mode 100644 index 0000000..93e9386 --- /dev/null +++ b/internal/engine/experiment/imap/imap.go @@ -0,0 +1,416 @@ +package imap + +import ( + "bufio" + "context" + "crypto/tls" + "fmt" + "github.com/pkg/errors" + "net" + //"net/smtp" + "net/url" + "strings" + "time" + + "github.com/ooni/probe-cli/v3/internal/engine/experiment/urlgetter" + "github.com/ooni/probe-cli/v3/internal/measurexlite" + "github.com/ooni/probe-cli/v3/internal/model" + "github.com/ooni/probe-cli/v3/internal/tracex" +) + +var ( + // errNoInputProvided indicates you didn't provide any input + errNoInputProvided = errors.New("not input provided") + + // errInputIsNotAnURL indicates that input is not an URL + errInputIsNotAnURL = errors.New("input is not an URL") + + // errInvalidScheme indicates that the scheme is invalid + errInvalidScheme = errors.New("scheme must be smtp(s)") +) + +const ( + testName = "imap" + testVersion = "0.0.1" +) + +// Config contains the experiment config. +type Config struct{} + +type RuntimeConfig struct { + host string + port string + forced_tls bool + noop_count uint8 +} + +func config(input model.MeasurementTarget) (*RuntimeConfig, error) { + if input == "" { + // TODO: static input data (eg. gmail/riseup..) + return nil, errNoInputProvided + } + + parsed, err := url.Parse(string(input)) + if err != nil { + return nil, fmt.Errorf("%w: %s", errInputIsNotAnURL, err.Error()) + } + if parsed.Scheme != "imap" && parsed.Scheme != "imaps" { + return nil, errInvalidScheme + } + + port := "" + + if parsed.Port() == "" { + // Default ports for StartTLS and forced TLS respectively + if parsed.Scheme == "imap" { + port = "143" + } else { + port = "993" + } + } else { + // Valid port is checked by URL parsing + port = parsed.Port() + } + + valid_config := RuntimeConfig{ + host: parsed.Hostname(), + forced_tls: parsed.Scheme == "imaps", + port: port, + noop_count: 10, + } + + return &valid_config, nil +} + +// TestKeys contains the experiment results + +type TestKeys struct { + Queries []*model.ArchivalDNSLookupResult `json:"queries"` + Runs map[string]*IndividualTestKeys `json:"runs"` + // Used for global failure (DNS resolution) + Failure string `json:"failure"` + // Indicates global failure or individual test failure + Failed bool `json:"failed"` +} + +// IndividualTestKeys contains results for TCP/IP level stuff for each address found +// in the DNS lookup +type IndividualTestKeys struct { + NoOpCounter uint8 + TCPConnect []*model.ArchivalTCPConnectResult `json:"tcp_connect"` + TLSHandshakes []*model.ArchivalTLSOrQUICHandshakeResult `json:"tls_handshakes"` + // Individual failure aborting the test run for this address/port combo + Failure *string `json:"failure"` +} + +type Measurer struct { + // Config contains the experiment settings. If empty we + // will be using default settings. + Config Config + + // Getter is an optional getter to be used for testing. + Getter urlgetter.MultiGetter +} + +// ExperimentName implements ExperimentMeasurer.ExperimentName +func (m Measurer) ExperimentName() string { + return testName +} + +// ExperimentVersion implements ExperimentMeasurer.ExperimentVersion +func (m Measurer) ExperimentVersion() string { + return testVersion +} + +// Manages sequential TCP sessions to the same hostname (over different IPs) +// don't use in parallel! +type TCPRunner struct { + trace *measurexlite.Trace + logger model.Logger + ctx context.Context + tk *TestKeys + tlsconfig *tls.Config + host string + port string + // addr is changed everytime TCPRunner.conn(addr) is called + addr string +} + +type TCPSession struct { + addr string + port string + runner *TCPRunner + tk *IndividualTestKeys + tls bool + raw_conn *net.Conn + tls_conn *net.Conn +} + +func (s *TCPSession) Close() { + if s.tls { + var conn = *s.tls_conn + conn.Close() + } else { + var conn = *s.raw_conn + conn.Close() + } +} + +func (s *TCPSession) current_conn() net.Conn { + if s.tls { + return *s.tls_conn + } else { + return *s.raw_conn + } +} + +func (r *TCPRunner) run_key() string { + return net.JoinHostPort(r.addr, r.port) +} + +func (r *TCPRunner) get_run() *IndividualTestKeys { + if r.tk.Runs == nil { + r.tk.Runs = make(map[string]*IndividualTestKeys) + } + key := r.run_key() + val, exists := r.tk.Runs[key] + if exists { + return val + } else { + r.tk.Runs[key] = &IndividualTestKeys{} + return r.tk.Runs[key] + } +} + +func (r *TCPRunner) conn(addr string, port string) (*TCPSession, bool) { + r.addr = addr + run := r.get_run() + + s := new(TCPSession) + if !s.conn(addr, port, r, run) { + return nil, false + } + return s, true +} + +func (r *TCPRunner) dial(addr string, port string) (net.Conn, error) { + dialer := r.trace.NewDialerWithoutResolver(r.logger) + conn, err := dialer.DialContext(r.ctx, "tcp", net.JoinHostPort(addr, port)) + run := r.get_run() + run.TCPConnect = append(run.TCPConnect, r.trace.TCPConnects()...) + return conn, err + +} + +func (s *TCPSession) conn(addr string, port string, runner *TCPRunner, tk *IndividualTestKeys) bool { + // Initialize addr field and corresponding errors in TestKeys + s.addr = addr + s.port = port + s.tls = false + s.runner = runner + s.tk = tk + + conn, err := runner.dial(addr, port) + if err != nil { + s.error(err) + return false + } + s.raw_conn = &conn + + return true +} + +func (s *TCPSession) error(err error) { + s.runner.tk.Failed = true + s.tk.Failure = tracex.NewFailure(err) + //s. = append(s.errors, tracex.NewFailure(err)) +} + +func (r *TCPRunner) resolve(host string) ([]string, bool) { + r.logger.Infof("Resolving DNS for %s", host) + resolver := r.trace.NewStdlibResolver(r.logger) + addrs, err := resolver.LookupHost(r.ctx, host) + r.tk.Queries = append(r.tk.Queries, r.trace.DNSLookupsFromRoundTrip()...) + if err != nil { + r.tk.Failure = *tracex.NewFailure(err) + return []string{}, false + } + r.logger.Infof("Finished DNS for %s: %v", host, addrs) + + return addrs, true +} + +func (s *TCPSession) handshake() bool { + if s.tls { + // TLS already initialized... + return true + } + s.runner.logger.Infof("Starting TLS handshake with %s:%s", s.addr, s.port) + thx := s.runner.trace.NewTLSHandshakerStdlib(s.runner.logger) + tconn, _, err := thx.Handshake(s.runner.ctx, *s.raw_conn, s.runner.tlsconfig) + s.tk.TLSHandshakes = append(s.tk.TLSHandshakes, s.runner.trace.FirstTLSHandshakeOrNil()) + if err != nil { + s.error(err) + return false + } + + s.tls = true + s.tls_conn = &tconn + s.runner.logger.Infof("Handshake succeeded") + return true +} + +func (s *TCPSession) starttls(message string) bool { + if s.tls { + // TLS already initialized... + return true + } + if message != "" { + s.runner.logger.Infof("Asking for StartTLS upgrade") + s.current_conn().Write([]byte(message)) + } + return s.handshake() +} + +func (s *TCPSession) imap(noop uint8) bool { + conn := s.current_conn() + + command, err := bufio.NewReader(conn).ReadString('\n') + if err != nil { + s.error(err) + return false + } + if !strings.Contains(command, "CAPABILITY") { + s.error(errors.New("Unexpected IMAP reply: " + command)) + return false + } + + + if noop > 0 { + s.runner.logger.Infof("Trying to generate no-op traffic") + s.tk.NoOpCounter = 0 + for s.tk.NoOpCounter < noop { + s.tk.NoOpCounter += 1 + s.runner.logger.Infof("NoOp Iteration %d", s.tk.NoOpCounter) + + conn.Write([]byte("A1 NOOP\n")) + command, err := bufio.NewReader(conn).ReadString('\n') + if err != nil { + s.error(err) + break + } + if !strings.Contains(command, "OK NOOP") { + s.error(errors.New("Unexpected IMAP reply: " + command)) + break + } + } + + if s.tk.NoOpCounter == noop { + s.runner.logger.Infof("Successfully generated no-op traffic") + return true + } else { + s.runner.logger.Infof("Failed no-op traffic at iteration %d", s.tk.NoOpCounter) + return false + } + } + + return true +} + +// Run implements ExperimentMeasurer.Run +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { + log := sess.Logger() + trace := measurexlite.NewTrace(0, measurement.MeasurementStartTimeSaved) + + config, err := config(measurement.Input) + if err != nil { + // Invalid input data, we don't even generate report + return err + } + + tk := new(TestKeys) + measurement.TestKeys = tk + + ctx, cancel := context.WithTimeout(ctx, 60*time.Second) + defer cancel() + + tlsconfig := tls.Config{ + InsecureSkipVerify: false, + ServerName: config.host, + } + + runner := &TCPRunner{ + trace: trace, + logger: log, + ctx: ctx, + tk: tk, + tlsconfig: &tlsconfig, + host: config.host, + port: config.port, + } + + // First resolve DNS + addrs, success := runner.resolve(config.host) + if !success { + return nil + } + + for _, addr := range addrs { + tcp_session, success := runner.conn(addr, config.port) + if !success { + continue + } + defer tcp_session.Close() + + if config.forced_tls { + // Direct TLS connection + if !tcp_session.handshake() { + continue + } + + // Try EHLO + NoOps + if !tcp_session.imap(config.noop_count) { + continue + } + } else { + // StartTLS... + if !tcp_session.starttls("A1 STARTTLS\n") { + continue + } + + if !tcp_session.imap(config.noop_count) { + continue + } + } + } + + return nil +} + +// NewExperimentMeasurer creates a new ExperimentMeasurer. +func NewExperimentMeasurer(config Config) model.ExperimentMeasurer { + return Measurer{Config: config} +} + +// SummaryKeys contains summary keys for this experiment. +// +// Note that this structure is part of the ABI contract with ooniprobe +// therefore we should be careful when changing it. +type SummaryKeys struct { + //DNSBlocking bool `json:"facebook_dns_blocking"` + //TCPBlocking bool `json:"facebook_tcp_blocking"` + IsAnomaly bool `json:"-"` +} + +// GetSummaryKeys implements model.ExperimentMeasurer.GetSummaryKeys. +func (m Measurer) GetSummaryKeys(measurement *model.Measurement) (interface{}, error) { + sk := SummaryKeys{IsAnomaly: false} + _, ok := measurement.TestKeys.(*TestKeys) + if !ok { + return sk, errors.New("invalid test keys type") + } + return sk, nil +} diff --git a/internal/engine/experiment/imap/imap_test.go b/internal/engine/experiment/imap/imap_test.go new file mode 100644 index 0000000..1fbfc15 --- /dev/null +++ b/internal/engine/experiment/imap/imap_test.go @@ -0,0 +1,186 @@ +package imap + +import ( + "bufio" + "context" + "crypto/tls" + //"encoding/json" + "errors" + "fmt" + "net" + "strings" + "testing" + + "github.com/ooni/probe-cli/v3/internal/engine/mockable" + "github.com/ooni/probe-cli/v3/internal/model" +) + +func plaintextListener() net.Listener { + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + if l, err = net.Listen("tcp6", "[::1]:0"); err != nil { + panic(fmt.Sprintf("httptest: failed to listen on a port: %v", err)) + } + } + return l +} + +func tlsListener(l net.Listener) net.Listener { + return tls.NewListener(l, &tls.Config{}) +} + +func listener_addr(l net.Listener) string { + return l.Addr().String() +} + +func ValidIMAPServer(conn net.Conn) { + for { + command, err := bufio.NewReader(conn).ReadString('\n') + if err != nil { + return + } + + if strings.Contains(command, "NOOP") { + conn.Write([]byte("A1 OK NOOP completed.\n")) + } else if command == "STARTTLS" { + conn.Write([]byte("A1 OK Begin TLS negotiation now.\n")) + // TODO: conn.Close does not actually close connection? or does client not detect it? + conn.Close() + return + } + conn.Write([]byte("\n")) + } +} + +func TCPServer(l net.Listener) { + for { + conn, err := l.Accept() + if err != nil { + continue + } + defer conn.Close() + conn.Write([]byte("* OK [CAPABILITY IMAP4rev1 SASL-IR LOGIN-REFERRALS ID ENABLE IDLE LITERAL+ STARTTLS LOGINDISABLED] howdy, ready.\n")) + ValidIMAPServer(conn) + } +} + +func TestMeasurer_run(t *testing.T) { + // runHelper is an helper function to run this set of tests. + runHelper := func(input string) (*model.Measurement, model.ExperimentMeasurer, error) { + m := NewExperimentMeasurer(Config{}) + if m.ExperimentName() != "imap" { + t.Fatal("invalid experiment name") + } + if m.ExperimentVersion() != "0.0.1" { + t.Fatal("invalid experiment version") + } + ctx := context.Background() + meas := &model.Measurement{ + Input: model.MeasurementTarget(input), + } + sess := &mockable.Session{ + MockableLogger: model.DiscardLogger, + } + callbacks := model.NewPrinterCallbacks(model.DiscardLogger) + err := m.Run(ctx, sess, meas, callbacks) + return meas, m, err + } + + t.Run("with empty input", func(t *testing.T) { + _, _, err := runHelper("") + if !errors.Is(err, errNoInputProvided) { + t.Fatal("unexpected error", err) + } + }) + + t.Run("with invalid URL", func(t *testing.T) { + _, _, err := runHelper("\t") + if !errors.Is(err, errInputIsNotAnURL) { + t.Fatal("unexpected error", err) + } + }) + + t.Run("with invalid scheme", func(t *testing.T) { + _, _, err := runHelper("https://8.8.8.8:443/") + if !errors.Is(err, errInvalidScheme) { + t.Fatal("unexpected error", err) + } + }) + + t.Run("with broken TLS", func(t *testing.T) { + p := plaintextListener() + defer p.Close() + + l := tlsListener(p) + defer l.Close() + addr := listener_addr(l) + go TCPServer(l) + + meas, m, err := runHelper("imaps://" + addr) + if err != nil { + t.Fatal(err) + } + + tk := meas.TestKeys.(*TestKeys) + + for _, run := range tk.Runs { + for _, handshake := range run.TLSHandshakes { + if *handshake.Failure != "unknown_failure: remote error: tls: unrecognized name" { + t.Fatal("expected unrecognized_name in TLS handshake") + } + } + + if run.NoOpCounter != 0 { + t.Fatalf("expected to not have any noops, not %d noops", run.NoOpCounter) + } + } + + ask, err := m.GetSummaryKeys(meas) + if err != nil { + t.Fatal("cannot obtain summary") + } + summary := ask.(SummaryKeys) + if summary.IsAnomaly { + t.Fatal("expected no anomaly") + } + }) + + t.Run("with broken starttls", func(t *testing.T) { + l := plaintextListener() + defer l.Close() + addr := listener_addr(l) + + go TCPServer(l) + + meas, m, err := runHelper("imap://" + addr) + if err != nil { + t.Fatal(err) + } + + tk := meas.TestKeys.(*TestKeys) + //bs, _ := json.Marshal(tk) + //fmt.Println(string(bs)) + + for _, run := range tk.Runs { + for _, handshake := range run.TLSHandshakes { + if *handshake.Failure != "unknown_failure: tls: first record does not look like a TLS handshake" { + + t.Fatal("expected broken handshake") + } + } + + if run.NoOpCounter != 0 { + t.Fatalf("expected to not have any noops, not %d noops", run.NoOpCounter) + } + } + + ask, err := m.GetSummaryKeys(meas) + if err != nil { + t.Fatal("cannot obtain summary") + } + summary := ask.(SummaryKeys) + if summary.IsAnomaly { + t.Fatal("expected no anomaly") + } + }) +} diff --git a/internal/engine/experiment/ndt7/ndt7.go b/internal/engine/experiment/ndt7/ndt7.go index 7e4bb23..8b403c6 100644 --- a/internal/engine/experiment/ndt7/ndt7.go +++ b/internal/engine/experiment/ndt7/ndt7.go @@ -210,10 +210,10 @@ func (m *Measurer) doUpload( } // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { tk := new(TestKeys) tk.Protocol = 7 measurement.TestKeys = tk diff --git a/internal/engine/experiment/ndt7/ndt7_test.go b/internal/engine/experiment/ndt7/ndt7_test.go index 7da663e..5ad4f1c 100644 --- a/internal/engine/experiment/ndt7/ndt7_test.go +++ b/internal/engine/experiment/ndt7/ndt7_test.go @@ -84,12 +84,7 @@ func TestRunWithCancelledContext(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() // immediately cancel meas := &model.Measurement{} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: meas, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, meas, model.NewPrinterCallbacks(log.Log)) // Here we get nil because we still want to submit this measurement if !errors.Is(err, nil) { t.Fatal("not the error we expected") @@ -109,15 +104,15 @@ func TestGood(t *testing.T) { } measurement := new(model.Measurement) measurer := NewExperimentMeasurer(Config{}) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableHTTPClient: http.DefaultClient, MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -138,15 +133,15 @@ func TestFailDownload(t *testing.T) { cancel() } meas := &model.Measurement{} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: meas, - Session: &mockable.Session{ + err := measurer.Run( + ctx, + &mockable.Session{ MockableHTTPClient: http.DefaultClient, MockableLogger: log.Log, }, - } - err := measurer.Run(ctx, args) + meas, + model.NewPrinterCallbacks(log.Log), + ) // We expect a nil failure here because we want to submit anyway // a measurement that failed to connect to m-lab. if err != nil { @@ -169,15 +164,15 @@ func TestFailUpload(t *testing.T) { cancel() } meas := &model.Measurement{} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: meas, - Session: &mockable.Session{ + err := measurer.Run( + ctx, + &mockable.Session{ MockableHTTPClient: http.DefaultClient, MockableLogger: log.Log, }, - } - err := measurer.Run(ctx, args) + meas, + model.NewPrinterCallbacks(log.Log), + ) // Here we expect a nil error because we want to submit this measurement if err != nil { t.Fatal(err) @@ -202,15 +197,15 @@ func TestDownloadJSONUnmarshalFail(t *testing.T) { seenError = true return expected } - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &model.Measurement{}, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableHTTPClient: http.DefaultClient, MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + new(model.Measurement), + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } diff --git a/internal/engine/experiment/portfiltering/measurer.go b/internal/engine/experiment/portfiltering/measurer.go index 487f036..76c9a2d 100644 --- a/internal/engine/experiment/portfiltering/measurer.go +++ b/internal/engine/experiment/portfiltering/measurer.go @@ -38,10 +38,12 @@ var ( ) // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { // TODO(DecFox): Replace the localhost deployment with an OONI testhelper // Ensure that we only do this once we have a deployed testhelper testhelper := "http://127.0.0.1" diff --git a/internal/engine/experiment/portfiltering/measurer_test.go b/internal/engine/experiment/portfiltering/measurer_test.go index c5ba589..e3a3876 100644 --- a/internal/engine/experiment/portfiltering/measurer_test.go +++ b/internal/engine/experiment/portfiltering/measurer_test.go @@ -29,12 +29,7 @@ func TestMeasurer_run(t *testing.T) { } callbacks := model.NewPrinterCallbacks(model.DiscardLogger) ctx := context.Background() - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: meas, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, meas, callbacks) if err != nil { t.Fatal(err) } diff --git a/internal/engine/experiment/psiphon/psiphon.go b/internal/engine/experiment/psiphon/psiphon.go index 45ed762..5f55aea 100644 --- a/internal/engine/experiment/psiphon/psiphon.go +++ b/internal/engine/experiment/psiphon/psiphon.go @@ -66,10 +66,10 @@ func (m *Measurer) printprogress( } // Run runs the measurement -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { const maxruntime = 300 ctx, cancel := context.WithTimeout(ctx, maxruntime*time.Second) var ( diff --git a/internal/engine/experiment/psiphon/psiphon_test.go b/internal/engine/experiment/psiphon/psiphon_test.go index 4186448..112ccdf 100644 --- a/internal/engine/experiment/psiphon/psiphon_test.go +++ b/internal/engine/experiment/psiphon/psiphon_test.go @@ -33,12 +33,8 @@ func TestRunWithCancelledContext(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() // fail immediately measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: newfakesession(), - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, newfakesession(), measurement, + model.NewPrinterCallbacks(log.Log)) if !errors.Is(err, nil) { // nil because we want to submit the measurement t.Fatal("expected another error here") } @@ -68,12 +64,8 @@ func TestRunWithCustomInputAndCancelledContext(t *testing.T) { } ctx, cancel := context.WithCancel(context.Background()) cancel() // fail immediately - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: newfakesession(), - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, newfakesession(), measurement, + model.NewPrinterCallbacks(log.Log)) if !errors.Is(err, nil) { // nil because we want to submit the measurement t.Fatal("expected another error here") } @@ -92,12 +84,7 @@ func TestRunWillPrintSomethingWithCancelledContext(t *testing.T) { cancel() // fail after we've given the printer a chance to run } observer := observerCallbacks{progress: &atomicx.Int64{}} - args := &model.ExperimentArgs{ - Callbacks: observer, - Measurement: measurement, - Session: newfakesession(), - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, newfakesession(), measurement, observer) if !errors.Is(err, nil) { // nil because we want to submit the measurement t.Fatal("expected another error here") } diff --git a/internal/engine/experiment/quicping/quicping.go b/internal/engine/experiment/quicping/quicping.go index cedf7c4..972f68e 100644 --- a/internal/engine/experiment/quicping/quicping.go +++ b/internal/engine/experiment/quicping/quicping.go @@ -221,11 +221,12 @@ func (m *Measurer) receiver( } // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { host := string(measurement.Input) // allow URL input if u, err := url.ParseRequestURI(host); err == nil { diff --git a/internal/engine/experiment/quicping/quicping_test.go b/internal/engine/experiment/quicping/quicping_test.go index 34a8023..79a214f 100644 --- a/internal/engine/experiment/quicping/quicping_test.go +++ b/internal/engine/experiment/quicping/quicping_test.go @@ -33,12 +33,8 @@ func TestInvalidHost(t *testing.T) { measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("a.a.a.a") sess := &mockable.Session{MockableLogger: log.Log} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run(context.Background(), sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err == nil { t.Fatal("expected an error here") } @@ -57,12 +53,8 @@ func TestURLInput(t *testing.T) { measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("https://google.com/") sess := &mockable.Session{MockableLogger: log.Log} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run(context.Background(), sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err != nil { t.Fatal("unexpected error") } @@ -81,12 +73,8 @@ func TestSuccess(t *testing.T) { measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("google.com") sess := &mockable.Session{MockableLogger: log.Log} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run(context.Background(), sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err != nil { t.Fatal("did not expect an error here") } @@ -129,12 +117,8 @@ func TestWithCancelledContext(t *testing.T) { sess := &mockable.Session{MockableLogger: log.Log} ctx, cancel := context.WithCancel(context.Background()) cancel() - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err != nil { t.Fatal("did not expect an error here") } @@ -154,12 +138,8 @@ func TestListenFails(t *testing.T) { measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("google.com") sess := &mockable.Session{MockableLogger: log.Log} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run(context.Background(), sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err == nil { t.Fatal("expected an error here") } @@ -202,12 +182,8 @@ func TestWriteFails(t *testing.T) { measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("google.com") sess := &mockable.Session{MockableLogger: log.Log} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run(context.Background(), sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err != nil { t.Fatal("unexpected error") } @@ -263,12 +239,8 @@ func TestReadFails(t *testing.T) { measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("google.com") sess := &mockable.Session{MockableLogger: log.Log} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run(context.Background(), sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err != nil { t.Fatal("unexpected error") } @@ -299,12 +271,8 @@ func TestNoResponse(t *testing.T) { measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("ooni.org") sess := &mockable.Session{MockableLogger: log.Log} - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run(context.Background(), sess, measurement, + model.NewPrinterCallbacks(log.Log)) if err != nil { t.Fatal("did not expect an error here") } diff --git a/internal/engine/experiment/riseupvpn/riseupvpn.go b/internal/engine/experiment/riseupvpn/riseupvpn.go index 10aa8de..62dd7e3 100644 --- a/internal/engine/experiment/riseupvpn/riseupvpn.go +++ b/internal/engine/experiment/riseupvpn/riseupvpn.go @@ -175,11 +175,8 @@ func (m Measurer) ExperimentVersion() string { } // Run implements ExperimentMeasurer.Run. -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m Measurer) Run(ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks) error { ctx, cancel := context.WithTimeout(ctx, 90*time.Second) defer cancel() testkeys := NewTestKeys() diff --git a/internal/engine/experiment/riseupvpn/riseupvpn_test.go b/internal/engine/experiment/riseupvpn/riseupvpn_test.go index 5ae4637..0e67541 100644 --- a/internal/engine/experiment/riseupvpn/riseupvpn_test.go +++ b/internal/engine/experiment/riseupvpn/riseupvpn_test.go @@ -100,7 +100,7 @@ const ( "cert": "XXXXXXXXXXXXXXXXXXXXXXXXX", "iatMode": "0" } - }, + }, { "type":"openvpn", "protocols":[ @@ -328,12 +328,7 @@ func TestInvalidCaCert(t *testing.T) { sess := &mockable.Session{MockableLogger: log.Log} measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -604,12 +599,7 @@ func TestMissingTransport(t *testing.T) { sess := &mockable.Session{MockableLogger: log.Log} measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err = measurer.Run(ctx, args) + err = measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -800,14 +790,14 @@ func runDefaultMockTest(t *testing.T, multiGetter urlgetter.MultiGetter) *model. } measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) diff --git a/internal/engine/experiment/run/dnscheck.go b/internal/engine/experiment/run/dnscheck.go index 538e9a0..da3fd84 100644 --- a/internal/engine/experiment/run/dnscheck.go +++ b/internal/engine/experiment/run/dnscheck.go @@ -21,10 +21,5 @@ func (m *dnsCheckMain) do(ctx context.Context, input StructuredInput, measurement.TestName = exp.ExperimentName() measurement.TestVersion = exp.ExperimentVersion() measurement.Input = model.MeasurementTarget(input.Input) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - return exp.Run(ctx, args) + return exp.Run(ctx, sess, measurement, callbacks) } diff --git a/internal/engine/experiment/run/run.go b/internal/engine/experiment/run/run.go index a25cf89..6c38057 100644 --- a/internal/engine/experiment/run/run.go +++ b/internal/engine/experiment/run/run.go @@ -46,10 +46,10 @@ type StructuredInput struct { } // Run implements ExperimentMeasurer.ExperimentVersion. -func (Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { var input StructuredInput if err := json.Unmarshal([]byte(measurement.Input), &input); err != nil { return err diff --git a/internal/engine/experiment/run/run_test.go b/internal/engine/experiment/run/run_test.go index 6fb3fde..06ae406 100644 --- a/internal/engine/experiment/run/run_test.go +++ b/internal/engine/experiment/run/run_test.go @@ -31,12 +31,7 @@ func TestRunDNSCheckWithCancelledContext(t *testing.T) { cancel() // fail immediately sess := &mockable.Session{MockableLogger: log.Log} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) // TODO(bassosimone): here we could improve the tests by checking // whether the result makes sense for a cancelled context. if err != nil { @@ -67,12 +62,7 @@ func TestRunURLGetterWithCancelledContext(t *testing.T) { cancel() // fail immediately sess := &mockable.Session{MockableLogger: log.Log} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { // here we expected nil b/c we want to submit the measurement t.Fatal(err) } @@ -96,12 +86,7 @@ func TestRunWithInvalidJSON(t *testing.T) { ctx := context.Background() sess := &mockable.Session{MockableLogger: log.Log} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err == nil || err.Error() != "invalid character '}' looking for beginning of value" { t.Fatalf("not the error we expected: %+v", err) } @@ -115,12 +100,7 @@ func TestRunWithUnknownExperiment(t *testing.T) { ctx := context.Background() sess := &mockable.Session{MockableLogger: log.Log} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err == nil || err.Error() != "no such experiment: antani" { t.Fatalf("not the error we expected: %+v", err) } diff --git a/internal/engine/experiment/run/urlgetter.go b/internal/engine/experiment/run/urlgetter.go index 753c051..085f78f 100644 --- a/internal/engine/experiment/run/urlgetter.go +++ b/internal/engine/experiment/run/urlgetter.go @@ -18,10 +18,5 @@ func (m *urlGetterMain) do(ctx context.Context, input StructuredInput, measurement.TestName = exp.ExperimentName() measurement.TestVersion = exp.ExperimentVersion() measurement.Input = model.MeasurementTarget(input.Input) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - return exp.Run(ctx, args) + return exp.Run(ctx, sess, measurement, callbacks) } diff --git a/internal/engine/experiment/signal/signal.go b/internal/engine/experiment/signal/signal.go index df89fac..07e01bb 100644 --- a/internal/engine/experiment/signal/signal.go +++ b/internal/engine/experiment/signal/signal.go @@ -141,10 +141,8 @@ func (m Measurer) ExperimentVersion() string { } // Run implements ExperimentMeasurer.Run -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m Measurer) Run(ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks) error { ctx, cancel := context.WithTimeout(ctx, 60*time.Second) defer cancel() urlgetter.RegisterExtensions(measurement) diff --git a/internal/engine/experiment/signal/signal_test.go b/internal/engine/experiment/signal/signal_test.go index f78e859..ecadebe 100644 --- a/internal/engine/experiment/signal/signal_test.go +++ b/internal/engine/experiment/signal/signal_test.go @@ -25,14 +25,14 @@ func TestNewExperimentMeasurer(t *testing.T) { func TestGood(t *testing.T) { measurer := signal.NewExperimentMeasurer(signal.Config{}) measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -103,14 +103,14 @@ func TestBadSignalCA(t *testing.T) { SignalCA: "INVALIDCA", }) measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err.Error() != "AppendCertsFromPEM failed" { t.Fatal("not the error we expected") } diff --git a/internal/engine/experiment/simplequicping/simplequicping.go b/internal/engine/experiment/simplequicping/simplequicping.go index 44d2f7e..eb1ebca 100644 --- a/internal/engine/experiment/simplequicping/simplequicping.go +++ b/internal/engine/experiment/simplequicping/simplequicping.go @@ -112,11 +112,12 @@ var ( ) // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { if measurement.Input == "" { return errNoInputProvided } diff --git a/internal/engine/experiment/simplequicping/simplequicping_test.go b/internal/engine/experiment/simplequicping/simplequicping_test.go index 12bd5ce..cdd4138 100644 --- a/internal/engine/experiment/simplequicping/simplequicping_test.go +++ b/internal/engine/experiment/simplequicping/simplequicping_test.go @@ -65,12 +65,7 @@ func TestMeasurer_run(t *testing.T) { MockableLogger: model.DiscardLogger, } callbacks := model.NewPrinterCallbacks(model.DiscardLogger) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: meas, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, meas, callbacks) return meas, m, err } diff --git a/internal/engine/experiment/smtp/smtp.go b/internal/engine/experiment/smtp/smtp.go new file mode 100644 index 0000000..568db04 --- /dev/null +++ b/internal/engine/experiment/smtp/smtp.go @@ -0,0 +1,413 @@ +package smtp + +import ( + "context" + "crypto/tls" + "fmt" + "github.com/pkg/errors" + "net" + "net/smtp" + "net/url" + "time" + + "github.com/ooni/probe-cli/v3/internal/engine/experiment/urlgetter" + "github.com/ooni/probe-cli/v3/internal/measurexlite" + "github.com/ooni/probe-cli/v3/internal/model" + "github.com/ooni/probe-cli/v3/internal/tracex" +) + +var ( + // errNoInputProvided indicates you didn't provide any input + errNoInputProvided = errors.New("not input provided") + + // errInputIsNotAnURL indicates that input is not an URL + errInputIsNotAnURL = errors.New("input is not an URL") + + // errInvalidScheme indicates that the scheme is invalid + errInvalidScheme = errors.New("scheme must be smtp(s)") +) + +const ( + testName = "smtp" + testVersion = "0.0.1" +) + +// Config contains the experiment config. +type Config struct{} + +type RuntimeConfig struct { + host string + port string + forced_tls bool + noop_count uint8 +} + +func config(input model.MeasurementTarget) (*RuntimeConfig, error) { + if input == "" { + // TODO: static input data (eg. gmail/riseup..) + return nil, errNoInputProvided + } + + parsed, err := url.Parse(string(input)) + if err != nil { + return nil, fmt.Errorf("%w: %s", errInputIsNotAnURL, err.Error()) + } + if parsed.Scheme != "smtp" && parsed.Scheme != "smtps" { + return nil, errInvalidScheme + } + + port := "" + + if parsed.Port() == "" { + // Default ports for StartTLS and forced TLS respectively + if parsed.Scheme == "smtp" { + port = "587" + } else { + port = "465" + } + } else { + // Valid port is checked by URL parsing + port = parsed.Port() + } + + valid_config := RuntimeConfig{ + host: parsed.Hostname(), + forced_tls: parsed.Scheme == "smtps", + port: port, + noop_count: 10, + } + + return &valid_config, nil +} + +// TestKeys contains the experiment results + +type TestKeys struct { + Queries []*model.ArchivalDNSLookupResult `json:"queries"` + Runs map[string]*IndividualTestKeys `json:"runs"` + // Used for global failure (DNS resolution) + Failure string `json:"failure"` + // Indicates global failure or individual test failure + Failed bool `json:"failed"` +} + +// IndividualTestKeys contains results for TCP/IP level stuff for each address found +// in the DNS lookup +type IndividualTestKeys struct { + NoOpCounter uint8 + TCPConnect []*model.ArchivalTCPConnectResult `json:"tcp_connect"` + TLSHandshakes []*model.ArchivalTLSOrQUICHandshakeResult `json:"tls_handshakes"` + // Individual failure aborting the test run for this address/port combo + Failure *string `json:"failure"` +} + +type Measurer struct { + // Config contains the experiment settings. If empty we + // will be using default settings. + Config Config + + // Getter is an optional getter to be used for testing. + Getter urlgetter.MultiGetter +} + +// ExperimentName implements ExperimentMeasurer.ExperimentName +func (m Measurer) ExperimentName() string { + return testName +} + +// ExperimentVersion implements ExperimentMeasurer.ExperimentVersion +func (m Measurer) ExperimentVersion() string { + return testVersion +} + +// Manages sequential TCP sessions to the same hostname (over different IPs) +// don't use in parallel! +type TCPRunner struct { + trace *measurexlite.Trace + logger model.Logger + ctx context.Context + tk *TestKeys + tlsconfig *tls.Config + host string + port string + // addr is changed everytime TCPRunner.conn(addr) is called + addr string +} + +type TCPSession struct { + addr string + port string + runner *TCPRunner + tk *IndividualTestKeys + tls bool + raw_conn *net.Conn + tls_conn *net.Conn +} + +func (s *TCPSession) Close() { + if s.tls { + var conn = *s.tls_conn + conn.Close() + } else { + var conn = *s.raw_conn + conn.Close() + } +} + +func (s *TCPSession) current_conn() net.Conn { + if s.tls { + return *s.tls_conn + } else { + return *s.raw_conn + } +} + +func (r *TCPRunner) run_key() string { + return net.JoinHostPort(r.addr, r.port) +} + +func (r *TCPRunner) get_run() *IndividualTestKeys { + if r.tk.Runs == nil { + r.tk.Runs = make(map[string]*IndividualTestKeys) + } + key := r.run_key() + val, exists := r.tk.Runs[key] + if exists { + return val + } else { + r.tk.Runs[key] = &IndividualTestKeys{} + return r.tk.Runs[key] + } +} + +func (r *TCPRunner) conn(addr string, port string) (*TCPSession, bool) { + r.addr = addr + run := r.get_run() + + s := new(TCPSession) + if !s.conn(addr, port, r, run) { + return nil, false + } + return s, true +} + +func (r *TCPRunner) dial(addr string, port string) (net.Conn, error) { + dialer := r.trace.NewDialerWithoutResolver(r.logger) + conn, err := dialer.DialContext(r.ctx, "tcp", net.JoinHostPort(addr, port)) + run := r.get_run() + run.TCPConnect = append(run.TCPConnect, r.trace.TCPConnects()...) + return conn, err + +} + +func (s *TCPSession) conn(addr string, port string, runner *TCPRunner, tk *IndividualTestKeys) bool { + // Initialize addr field and corresponding errors in TestKeys + s.addr = addr + s.port = port + s.tls = false + s.runner = runner + s.tk = tk + + conn, err := runner.dial(addr, port) + if err != nil { + s.error(err) + return false + } + s.raw_conn = &conn + + return true +} + +func (s *TCPSession) error(err error) { + s.runner.tk.Failed = true + s.tk.Failure = tracex.NewFailure(err) + //s. = append(s.errors, tracex.NewFailure(err)) +} + +func (r *TCPRunner) resolve(host string) ([]string, bool) { + r.logger.Infof("Resolving DNS for %s", host) + resolver := r.trace.NewStdlibResolver(r.logger) + addrs, err := resolver.LookupHost(r.ctx, host) + r.tk.Queries = append(r.tk.Queries, r.trace.DNSLookupsFromRoundTrip()...) + if err != nil { + r.tk.Failure = *tracex.NewFailure(err) + return []string{}, false + } + r.logger.Infof("Finished DNS for %s: %v", host, addrs) + + return addrs, true +} + +func (s *TCPSession) handshake() bool { + if s.tls { + // TLS already initialized... + return true + } + s.runner.logger.Infof("Starting TLS handshake with %s:%s", s.addr, s.port) + thx := s.runner.trace.NewTLSHandshakerStdlib(s.runner.logger) + tconn, _, err := thx.Handshake(s.runner.ctx, *s.raw_conn, s.runner.tlsconfig) + s.tk.TLSHandshakes = append(s.tk.TLSHandshakes, s.runner.trace.FirstTLSHandshakeOrNil()) + if err != nil { + s.error(err) + return false + } + + s.tls = true + s.tls_conn = &tconn + s.runner.logger.Infof("Handshake succeeded") + return true +} + +func (s *TCPSession) starttls(message string) bool { + if s.tls { + // TLS already initialized... + return true + } + if message != "" { + s.runner.logger.Infof("Asking for StartTLS upgrade") + s.current_conn().Write([]byte(message)) + } + return s.handshake() +} + +func (s *TCPSession) smtp(ehlo string, noop uint8) bool { + // Auto-choose plaintext/TCP session + client, err := smtp.NewClient(s.current_conn(), ehlo) + if err != nil { + s.error(err) + return false + } + err = client.Hello(ehlo) + if err != nil { + s.error(err) + return false + } + + if noop > 0 { + s.runner.logger.Infof("Trying to generate more no-op traffic") + // TODO: noop counter per IP address + s.tk.NoOpCounter = 0 + for s.tk.NoOpCounter < noop { + s.tk.NoOpCounter += 1 + s.runner.logger.Infof("NoOp Iteration %d", s.tk.NoOpCounter) + err = client.Noop() + if err != nil { + s.error(err) + break + } + } + + if s.tk.NoOpCounter == noop { + s.runner.logger.Infof("Successfully generated no-op traffic") + return true + } else { + s.runner.logger.Infof("Failed no-op traffic at iteration %d", s.tk.NoOpCounter) + return false + } + } + + return true +} + +// Run implements ExperimentMeasurer.Run +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { + log := sess.Logger() + trace := measurexlite.NewTrace(0, measurement.MeasurementStartTimeSaved) + + config, err := config(measurement.Input) + if err != nil { + // Invalid input data, we don't even generate report + return err + } + + tk := new(TestKeys) + measurement.TestKeys = tk + + ctx, cancel := context.WithTimeout(ctx, 60*time.Second) + defer cancel() + + tlsconfig := tls.Config{ + InsecureSkipVerify: false, + ServerName: config.host, + } + + runner := &TCPRunner{ + trace: trace, + logger: log, + ctx: ctx, + tk: tk, + tlsconfig: &tlsconfig, + host: config.host, + port: config.port, + } + + // First resolve DNS + addrs, success := runner.resolve(config.host) + if !success { + return nil + } + + for _, addr := range addrs { + tcp_session, success := runner.conn(addr, config.port) + if !success { + continue + } + defer tcp_session.Close() + + if config.forced_tls { + // Direct TLS connection + if !tcp_session.handshake() { + continue + } + + // Try EHLO + NoOps + if !tcp_session.smtp("localhost", config.noop_count) { + continue + } + } else { + // StartTLS... first try plaintext EHLO + if !tcp_session.smtp("localhost", 0) { + continue + } + + // Upgrade via StartTLS and try EHLO + NoOps + if !tcp_session.starttls("STARTTLS\n") { + continue + } + + if !tcp_session.smtp("localhost", config.noop_count) { + continue + } + } + } + + return nil +} + +// NewExperimentMeasurer creates a new ExperimentMeasurer. +func NewExperimentMeasurer(config Config) model.ExperimentMeasurer { + return Measurer{Config: config} +} + +// SummaryKeys contains summary keys for this experiment. +// +// Note that this structure is part of the ABI contract with ooniprobe +// therefore we should be careful when changing it. +type SummaryKeys struct { + //DNSBlocking bool `json:"facebook_dns_blocking"` + //TCPBlocking bool `json:"facebook_tcp_blocking"` + IsAnomaly bool `json:"-"` +} + +// GetSummaryKeys implements model.ExperimentMeasurer.GetSummaryKeys. +func (m Measurer) GetSummaryKeys(measurement *model.Measurement) (interface{}, error) { + sk := SummaryKeys{IsAnomaly: false} + _, ok := measurement.TestKeys.(*TestKeys) + if !ok { + return sk, errors.New("invalid test keys type") + } + return sk, nil +} diff --git a/internal/engine/experiment/smtp/smtp_test.go b/internal/engine/experiment/smtp/smtp_test.go new file mode 100644 index 0000000..e49af52 --- /dev/null +++ b/internal/engine/experiment/smtp/smtp_test.go @@ -0,0 +1,185 @@ +package smtp + +import ( + "bufio" + "context" + "crypto/tls" + "errors" + "fmt" + "net" + "strings" + "testing" + + "github.com/ooni/probe-cli/v3/internal/engine/mockable" + "github.com/ooni/probe-cli/v3/internal/model" +) + +func plaintextListener() net.Listener { + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + if l, err = net.Listen("tcp6", "[::1]:0"); err != nil { + panic(fmt.Sprintf("httptest: failed to listen on a port: %v", err)) + } + } + return l +} + +func tlsListener(l net.Listener) net.Listener { + return tls.NewListener(l, &tls.Config{}) +} + +func listener_addr(l net.Listener) string { + return l.Addr().String() +} + +func ValidSMTPServer(conn net.Conn) { + for { + command, err := bufio.NewReader(conn).ReadString('\n') + if err != nil { + return + } + + if command == "" { + } else if command == "NOOP" { + conn.Write([]byte("250 2.0.0 Ok\n")) + } else if command == "STARTTLS" { + conn.Write([]byte("220 2.0.0 Ready to start TLS\n")) + // TODO: conn.Close does not actually close connection? or does client not detect it? + conn.Close() + return + } else if strings.HasPrefix(command, "EHLO") { + conn.Write([]byte("250 mock.example.com\n")) + } + conn.Write([]byte("\n")) + } +} + +func TCPServer(l net.Listener) { + for { + conn, err := l.Accept() + if err != nil { + continue + } + defer conn.Close() + conn.Write([]byte("220 mock.example.com ESMTP (spam is not appreciated)\n")) + ValidSMTPServer(conn) + } +} + +func TestMeasurer_run(t *testing.T) { + // runHelper is an helper function to run this set of tests. + runHelper := func(input string) (*model.Measurement, model.ExperimentMeasurer, error) { + m := NewExperimentMeasurer(Config{}) + if m.ExperimentName() != "smtp" { + t.Fatal("invalid experiment name") + } + if m.ExperimentVersion() != "0.0.1" { + t.Fatal("invalid experiment version") + } + ctx := context.Background() + meas := &model.Measurement{ + Input: model.MeasurementTarget(input), + } + sess := &mockable.Session{ + MockableLogger: model.DiscardLogger, + } + callbacks := model.NewPrinterCallbacks(model.DiscardLogger) + err := m.Run(ctx, sess, meas, callbacks) + return meas, m, err + } + + t.Run("with empty input", func(t *testing.T) { + _, _, err := runHelper("") + if !errors.Is(err, errNoInputProvided) { + t.Fatal("unexpected error", err) + } + }) + + t.Run("with invalid URL", func(t *testing.T) { + _, _, err := runHelper("\t") + if !errors.Is(err, errInputIsNotAnURL) { + t.Fatal("unexpected error", err) + } + }) + + t.Run("with invalid scheme", func(t *testing.T) { + _, _, err := runHelper("https://8.8.8.8:443/") + if !errors.Is(err, errInvalidScheme) { + t.Fatal("unexpected error", err) + } + }) + + t.Run("with broken TLS", func(t *testing.T) { + p := plaintextListener() + defer p.Close() + + l := tlsListener(p) + defer l.Close() + addr := listener_addr(l) + go TCPServer(l) + + meas, m, err := runHelper("smtps://" + addr) + if err != nil { + t.Fatal(err) + } + + tk := meas.TestKeys.(*TestKeys) + + for _, run := range tk.Runs { + for _, handshake := range run.TLSHandshakes { + if *handshake.Failure != "unknown_failure: remote error: tls: unrecognized name" { + t.Fatal("expected unrecognized_name in TLS handshake") + } + } + + if run.NoOpCounter != 0 { + t.Fatalf("expected to not have any noops, not %d noops", run.NoOpCounter) + } + } + + ask, err := m.GetSummaryKeys(meas) + if err != nil { + t.Fatal("cannot obtain summary") + } + summary := ask.(SummaryKeys) + if summary.IsAnomaly { + t.Fatal("expected no anomaly") + } + }) + + t.Run("with broken starttls", func(t *testing.T) { + l := plaintextListener() + defer l.Close() + addr := listener_addr(l) + + go TCPServer(l) + + meas, m, err := runHelper("smtp://" + addr) + if err != nil { + t.Fatal(err) + } + + tk := meas.TestKeys.(*TestKeys) + + for _, run := range tk.Runs { + for _, handshake := range run.TLSHandshakes { + if *handshake.Failure != "generic_timeout_error" { + t.Fatal("expected timeout in TLS handshake") + } + } + + if run.NoOpCounter != 0 { + t.Fatalf("expected to not have any noops, not %d noops", run.NoOpCounter) + } + } + + ask, err := m.GetSummaryKeys(meas) + if err != nil { + t.Fatal("cannot obtain summary") + } + summary := ask.(SummaryKeys) + if summary.IsAnomaly { + t.Fatal("expected no anomaly") + } + }) +} diff --git a/internal/engine/experiment/sniblocking/sniblocking.go b/internal/engine/experiment/sniblocking/sniblocking.go index 9bb3a40..4e4df18 100644 --- a/internal/engine/experiment/sniblocking/sniblocking.go +++ b/internal/engine/experiment/sniblocking/sniblocking.go @@ -233,10 +233,12 @@ func maybeURLToSNI(input model.MeasurementTarget) (model.MeasurementTarget, erro } // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { m.mu.Lock() if m.cache == nil { m.cache = make(map[string]Subresult) diff --git a/internal/engine/experiment/sniblocking/sniblocking_test.go b/internal/engine/experiment/sniblocking/sniblocking_test.go index 7424fbf..63ec9c7 100644 --- a/internal/engine/experiment/sniblocking/sniblocking_test.go +++ b/internal/engine/experiment/sniblocking/sniblocking_test.go @@ -116,12 +116,12 @@ func TestMeasurerMeasureNoMeasurementInput(t *testing.T) { measurer := NewExperimentMeasurer(Config{ ControlSNI: "example.com", }) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &model.Measurement{}, - Session: newsession(), - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + newsession(), + new(model.Measurement), + model.NewPrinterCallbacks(log.Log), + ) if err.Error() != "Experiment requires measurement.Input" { t.Fatal("not the error we expected") } @@ -136,12 +136,12 @@ func TestMeasurerMeasureWithInvalidInput(t *testing.T) { measurement := &model.Measurement{ Input: "\t", } - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: newsession(), - } - err := measurer.Run(ctx, args) + err := measurer.Run( + ctx, + newsession(), + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err == nil { t.Fatal("expected an error here") } @@ -156,12 +156,12 @@ func TestMeasurerMeasureWithCancelledContext(t *testing.T) { measurement := &model.Measurement{ Input: "kernel.org", } - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: newsession(), - } - err := measurer.Run(ctx, args) + err := measurer.Run( + ctx, + newsession(), + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } diff --git a/internal/engine/experiment/stunreachability/stunreachability.go b/internal/engine/experiment/stunreachability/stunreachability.go index 870500f..ba278c0 100644 --- a/internal/engine/experiment/stunreachability/stunreachability.go +++ b/internal/engine/experiment/stunreachability/stunreachability.go @@ -73,10 +73,10 @@ var errStunMissingPortInURL = errors.New("stun: missing port in URL") var errUnsupportedURLScheme = errors.New("stun: unsupported URL scheme") // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { tk := new(TestKeys) measurement.TestKeys = tk registerExtensions(measurement) diff --git a/internal/engine/experiment/stunreachability/stunreachability_test.go b/internal/engine/experiment/stunreachability/stunreachability_test.go index 209fd18..0b6c2be 100644 --- a/internal/engine/experiment/stunreachability/stunreachability_test.go +++ b/internal/engine/experiment/stunreachability/stunreachability_test.go @@ -32,12 +32,12 @@ func TestMeasurerExperimentNameVersion(t *testing.T) { func TestRunWithoutInput(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + &mockable.Session{}, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, errStunMissingInput) { t.Fatal("not the error we expected", err) } @@ -47,12 +47,12 @@ func TestRunWithInvalidURL(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("\t") // <- invalid URL - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + &mockable.Session{}, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err == nil || !strings.HasSuffix(err.Error(), "invalid control character in URL") { t.Fatal("not the error we expected", err) } @@ -62,12 +62,12 @@ func TestRunWithNoPort(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("stun://stun.ekiga.net") - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + &mockable.Session{}, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, errStunMissingPortInURL) { t.Fatal("not the error we expected", err) } @@ -77,12 +77,12 @@ func TestRunWithUnsupportedURLScheme(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget("https://stun.ekiga.net:3478") - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + &mockable.Session{}, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, errUnsupportedURLScheme) { t.Fatal("not the error we expected", err) } @@ -92,14 +92,14 @@ func TestRunWithInput(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget(defaultInput) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: model.DiscardLogger, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -124,14 +124,14 @@ func TestCancelledContext(t *testing.T) { measurer := NewExperimentMeasurer(Config{}) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget(defaultInput) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + ctx, + &mockable.Session{ MockableLogger: model.DiscardLogger, }, - } - err := measurer.Run(ctx, args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, nil) { // nil because we want to submit t.Fatal("not the error we expected", err) } @@ -166,14 +166,14 @@ func TestNewClientFailure(t *testing.T) { measurer := NewExperimentMeasurer(*config) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget(defaultInput) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: model.DiscardLogger, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, nil) { // nil because we want to submit t.Fatal("not the error we expected") } @@ -202,14 +202,14 @@ func TestStartFailure(t *testing.T) { measurer := NewExperimentMeasurer(*config) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget(defaultInput) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: model.DiscardLogger, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, nil) { // nil because we want to submit t.Fatal("not the error we expected") } @@ -242,14 +242,14 @@ func TestReadFailure(t *testing.T) { measurer := NewExperimentMeasurer(*config) measurement := new(model.Measurement) measurement.Input = model.MeasurementTarget(defaultInput) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: model.DiscardLogger, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, nil) { // nil because we want to submit t.Fatal("not the error we expected") } diff --git a/internal/engine/experiment/tcpping/tcpping.go b/internal/engine/experiment/tcpping/tcpping.go index 0354233..c8cfbae 100644 --- a/internal/engine/experiment/tcpping/tcpping.go +++ b/internal/engine/experiment/tcpping/tcpping.go @@ -82,10 +82,12 @@ var ( ) // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { if measurement.Input == "" { return errNoInputProvided } diff --git a/internal/engine/experiment/tcpping/tcpping_test.go b/internal/engine/experiment/tcpping/tcpping_test.go index 330b9cb..faec432 100644 --- a/internal/engine/experiment/tcpping/tcpping_test.go +++ b/internal/engine/experiment/tcpping/tcpping_test.go @@ -51,12 +51,7 @@ func TestMeasurer_run(t *testing.T) { MockableLogger: model.DiscardLogger, } callbacks := model.NewPrinterCallbacks(model.DiscardLogger) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: meas, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, meas, callbacks) return meas, m, err } diff --git a/internal/engine/experiment/telegram/telegram.go b/internal/engine/experiment/telegram/telegram.go index d9f606e..c9a0902 100644 --- a/internal/engine/experiment/telegram/telegram.go +++ b/internal/engine/experiment/telegram/telegram.go @@ -101,11 +101,8 @@ func (m Measurer) ExperimentVersion() string { } // Run implements ExperimentMeasurer.Run -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m Measurer) Run(ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks) error { ctx, cancel := context.WithTimeout(ctx, 60*time.Second) defer cancel() urlgetter.RegisterExtensions(measurement) diff --git a/internal/engine/experiment/telegram/telegram_test.go b/internal/engine/experiment/telegram/telegram_test.go index 8a1417c..4d206bf 100644 --- a/internal/engine/experiment/telegram/telegram_test.go +++ b/internal/engine/experiment/telegram/telegram_test.go @@ -28,14 +28,14 @@ func TestNewExperimentMeasurer(t *testing.T) { func TestGood(t *testing.T) { measurer := telegram.NewExperimentMeasurer(telegram.Config{}) measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -297,12 +297,7 @@ func TestWeConfigureWebChecksToFailOnHTTPError(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := measurer.Run(ctx, args); err != nil { + if err := measurer.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } if called.Load() < 1 { diff --git a/internal/engine/experiment/tlsmiddlebox/measurer.go b/internal/engine/experiment/tlsmiddlebox/measurer.go index 363e7a6..d157ab0 100644 --- a/internal/engine/experiment/tlsmiddlebox/measurer.go +++ b/internal/engine/experiment/tlsmiddlebox/measurer.go @@ -52,10 +52,12 @@ var ( ) // // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { if measurement.Input == "" { return errNoInputProvided } diff --git a/internal/engine/experiment/tlsmiddlebox/measurer_test.go b/internal/engine/experiment/tlsmiddlebox/measurer_test.go index 318fd5d..f5bf815 100644 --- a/internal/engine/experiment/tlsmiddlebox/measurer_test.go +++ b/internal/engine/experiment/tlsmiddlebox/measurer_test.go @@ -38,12 +38,7 @@ func TestMeasurer_input_failure(t *testing.T) { }, } callbacks := model.NewPrinterCallbacks(model.DiscardLogger) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: meas, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, meas, callbacks) return meas, m, err } diff --git a/internal/engine/experiment/tlsping/tlsping.go b/internal/engine/experiment/tlsping/tlsping.go index de51864..eda8023 100644 --- a/internal/engine/experiment/tlsping/tlsping.go +++ b/internal/engine/experiment/tlsping/tlsping.go @@ -112,10 +112,12 @@ var ( ) // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { if measurement.Input == "" { return errNoInputProvided } diff --git a/internal/engine/experiment/tlsping/tlsping_test.go b/internal/engine/experiment/tlsping/tlsping_test.go index fad4388..716549f 100644 --- a/internal/engine/experiment/tlsping/tlsping_test.go +++ b/internal/engine/experiment/tlsping/tlsping_test.go @@ -58,12 +58,7 @@ func TestMeasurer_run(t *testing.T) { MockableLogger: model.DiscardLogger, } callbacks := model.NewPrinterCallbacks(model.DiscardLogger) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: meas, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, meas, callbacks) return meas, m, err } diff --git a/internal/engine/experiment/tlstool/tlstool.go b/internal/engine/experiment/tlstool/tlstool.go index 9791b2a..812d081 100644 --- a/internal/engine/experiment/tlstool/tlstool.go +++ b/internal/engine/experiment/tlstool/tlstool.go @@ -78,11 +78,12 @@ var allMethods = []method{{ }} // Run implements ExperimentMeasurer.Run. -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { // TODO(bassosimone): wondering whether this experiment should // actually be merged with sniblocking instead? tk := new(TestKeys) diff --git a/internal/engine/experiment/tlstool/tlstool_test.go b/internal/engine/experiment/tlstool/tlstool_test.go index 55a33bc..6f95dde 100644 --- a/internal/engine/experiment/tlstool/tlstool_test.go +++ b/internal/engine/experiment/tlstool/tlstool_test.go @@ -27,12 +27,12 @@ func TestRunWithExplicitSNI(t *testing.T) { }) measurement := new(model.Measurement) measurement.Input = "8.8.8.8:853" - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := measurer.Run(ctx, args) + err := measurer.Run( + ctx, + &mockable.Session{}, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -43,12 +43,12 @@ func TestRunWithImplicitSNI(t *testing.T) { measurer := tlstool.NewExperimentMeasurer(tlstool.Config{}) measurement := new(model.Measurement) measurement.Input = "dns.google:853" - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := measurer.Run(ctx, args) + err := measurer.Run( + ctx, + &mockable.Session{}, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -60,12 +60,12 @@ func TestRunWithCancelledContext(t *testing.T) { measurer := tlstool.NewExperimentMeasurer(tlstool.Config{}) measurement := new(model.Measurement) measurement.Input = "dns.google:853" - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := measurer.Run(ctx, args) + err := measurer.Run( + ctx, + &mockable.Session{}, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } diff --git a/internal/engine/experiment/tor/tor.go b/internal/engine/experiment/tor/tor.go index ffd52fa..19c8f50 100644 --- a/internal/engine/experiment/tor/tor.go +++ b/internal/engine/experiment/tor/tor.go @@ -166,10 +166,12 @@ func (m *Measurer) ExperimentVersion() string { } // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { targets, err := m.gimmeTargets(ctx, sess) if err != nil { return err // fail the measurement if we cannot get any target diff --git a/internal/engine/experiment/tor/tor_test.go b/internal/engine/experiment/tor/tor_test.go index 4e02549..f2eabeb 100644 --- a/internal/engine/experiment/tor/tor_test.go +++ b/internal/engine/experiment/tor/tor_test.go @@ -36,14 +36,14 @@ func TestMeasurerMeasureFetchTorTargetsError(t *testing.T) { measurer.fetchTorTargets = func(ctx context.Context, sess model.ExperimentSession, cc string) (map[string]model.OOAPITorTarget, error) { return nil, expected } - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &model.Measurement{}, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + new(model.Measurement), + model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, expected) { t.Fatal("not the error we expected") } @@ -55,14 +55,14 @@ func TestMeasurerMeasureFetchTorTargetsEmptyList(t *testing.T) { return nil, nil } measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -79,14 +79,14 @@ func TestMeasurerMeasureGoodWithMockedOrchestra(t *testing.T) { measurer.fetchTorTargets = func(ctx context.Context, sess model.ExperimentSession, cc string) (map[string]model.OOAPITorTarget, error) { return nil, nil } - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: &model.Measurement{}, - Session: &mockable.Session{ + err := measurer.Run( + context.Background(), + &mockable.Session{ MockableLogger: log.Log, }, - } - err := measurer.Run(context.Background(), args) + new(model.Measurement), + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -99,12 +99,12 @@ func TestMeasurerMeasureGood(t *testing.T) { measurer := NewMeasurer(Config{}) sess := newsession() measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + sess, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } @@ -142,12 +142,12 @@ func TestMeasurerMeasureSanitiseOutput(t *testing.T) { key: staticPrivateTestingTarget, } measurement := new(model.Measurement) - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: sess, - } - err := measurer.Run(context.Background(), args) + err := measurer.Run( + context.Background(), + sess, + measurement, + model.NewPrinterCallbacks(log.Log), + ) if err != nil { t.Fatal(err) } diff --git a/internal/engine/experiment/torsf/integration_test.go b/internal/engine/experiment/torsf/integration_test.go index 11e846d..80feb82 100644 --- a/internal/engine/experiment/torsf/integration_test.go +++ b/internal/engine/experiment/torsf/integration_test.go @@ -34,12 +34,7 @@ func TestRunWithExistingTor(t *testing.T) { MockableLogger: log.Log, MockableTempDir: tempdir, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err = m.Run(ctx, args); err != nil { + if err = m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } } diff --git a/internal/engine/experiment/torsf/torsf.go b/internal/engine/experiment/torsf/torsf.go index 99ad54c..aef0f82 100644 --- a/internal/engine/experiment/torsf/torsf.go +++ b/internal/engine/experiment/torsf/torsf.go @@ -124,10 +124,10 @@ const maxRuntime = 600 * time.Second // set the relevant OONI error inside of the measurement and // return nil. This is important because the caller may not submit // the measurement if this method returns an error. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { ptl, sfdialer, err := m.setup(ctx, sess.Logger()) if err != nil { // we cannot setup the experiment diff --git a/internal/engine/experiment/torsf/torsf_test.go b/internal/engine/experiment/torsf/torsf_test.go index afdd81a..e71c666 100644 --- a/internal/engine/experiment/torsf/torsf_test.go +++ b/internal/engine/experiment/torsf/torsf_test.go @@ -47,12 +47,7 @@ func TestFailureWithInvalidRendezvousMethod(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := m.Run(ctx, args) + err := m.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, ptx.ErrSnowflakeNoSuchRendezvousMethod) { t.Fatal("unexpected error", err) } @@ -75,12 +70,7 @@ func TestFailureToStartPTXListener(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := m.Run(ctx, args); !errors.Is(err, expected) { + if err := m.Run(ctx, sess, measurement, callbacks); !errors.Is(err, expected) { t.Fatal("not the error we expected", err) } if tk := measurement.TestKeys; tk != nil { @@ -118,12 +108,7 @@ func TestSuccessWithMockedTunnelStart(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := m.Run(ctx, args); err != nil { + if err := m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } if called.Load() != 1 { @@ -183,12 +168,7 @@ func TestWithCancelledContext(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := m.Run(ctx, args); err != nil { + if err := m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } tk := measurement.TestKeys.(*TestKeys) @@ -251,12 +231,7 @@ func TestFailureToStartTunnel(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := m.Run(ctx, args); err != nil { + if err := m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } tk := measurement.TestKeys.(*TestKeys) diff --git a/internal/engine/experiment/urlgetter/urlgetter.go b/internal/engine/experiment/urlgetter/urlgetter.go index 3414607..8a6593c 100644 --- a/internal/engine/experiment/urlgetter/urlgetter.go +++ b/internal/engine/experiment/urlgetter/urlgetter.go @@ -97,10 +97,10 @@ func (m Measurer) ExperimentVersion() string { } // Run implements model.ExperimentSession.Run -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { // When using the urlgetter experiment directly, there is a nonconfigurable // default timeout that applies. When urlgetter is used as a library, it's // instead the responsibility of the user of urlgetter to set timeouts. Note diff --git a/internal/engine/experiment/urlgetter/urlgetter_test.go b/internal/engine/experiment/urlgetter/urlgetter_test.go index f9e62da..eaa8f7b 100644 --- a/internal/engine/experiment/urlgetter/urlgetter_test.go +++ b/internal/engine/experiment/urlgetter/urlgetter_test.go @@ -23,12 +23,10 @@ func TestMeasurer(t *testing.T) { } measurement := new(model.Measurement) measurement.Input = "https://www.google.com" - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := m.Run(ctx, args) + err := m.Run( + ctx, &mockable.Session{}, + measurement, model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, nil) { // nil because we want to submit the measurement t.Fatal("not the error we expected") } @@ -62,12 +60,10 @@ func TestMeasurerDNSCache(t *testing.T) { } measurement := new(model.Measurement) measurement.Input = "https://www.google.com" - args := &model.ExperimentArgs{ - Callbacks: model.NewPrinterCallbacks(log.Log), - Measurement: measurement, - Session: &mockable.Session{}, - } - err := m.Run(ctx, args) + err := m.Run( + ctx, &mockable.Session{}, + measurement, model.NewPrinterCallbacks(log.Log), + ) if !errors.Is(err, nil) { // nil because we want to submit the measurement t.Fatal("not the error we expected") } diff --git a/internal/engine/experiment/vanillator/integration_test.go b/internal/engine/experiment/vanillator/integration_test.go index c86ae9e..c569303 100644 --- a/internal/engine/experiment/vanillator/integration_test.go +++ b/internal/engine/experiment/vanillator/integration_test.go @@ -34,12 +34,7 @@ func TestRunWithExistingTor(t *testing.T) { MockableLogger: log.Log, MockableTempDir: tempdir, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err = m.Run(ctx, args); err != nil { + if err = m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } } diff --git a/internal/engine/experiment/vanillator/vanillator.go b/internal/engine/experiment/vanillator/vanillator.go index 1b6d687..593f07c 100644 --- a/internal/engine/experiment/vanillator/vanillator.go +++ b/internal/engine/experiment/vanillator/vanillator.go @@ -106,10 +106,10 @@ const maxRuntime = 200 * time.Second // set the relevant OONI error inside of the measurement and // return nil. This is important because the caller may not submit // the measurement if this method returns an error. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { m.registerExtensions(measurement) start := time.Now() ctx, cancel := context.WithTimeout(ctx, maxRuntime) diff --git a/internal/engine/experiment/vanillator/vanillator_test.go b/internal/engine/experiment/vanillator/vanillator_test.go index fda2f5e..e8cbdd4 100644 --- a/internal/engine/experiment/vanillator/vanillator_test.go +++ b/internal/engine/experiment/vanillator/vanillator_test.go @@ -59,12 +59,7 @@ func TestSuccessWithMockedTunnelStart(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := m.Run(ctx, args); err != nil { + if err := m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } if called.Load() != 1 { @@ -118,12 +113,7 @@ func TestWithCancelledContext(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := m.Run(ctx, args); err != nil { + if err := m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } tk := measurement.TestKeys.(*TestKeys) @@ -180,12 +170,7 @@ func TestFailureToStartTunnel(t *testing.T) { callbacks := &model.PrinterCallbacks{ Logger: model.DiscardLogger, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := m.Run(ctx, args); err != nil { + if err := m.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } tk := measurement.TestKeys.(*TestKeys) diff --git a/internal/engine/experiment/webconnectivity/control.go b/internal/engine/experiment/webconnectivity/control.go index bd953a1..ccaf870 100644 --- a/internal/engine/experiment/webconnectivity/control.go +++ b/internal/engine/experiment/webconnectivity/control.go @@ -4,10 +4,9 @@ import ( "context" "github.com/ooni/probe-cli/v3/internal/geoipx" - "github.com/ooni/probe-cli/v3/internal/httpapi" + "github.com/ooni/probe-cli/v3/internal/httpx" "github.com/ooni/probe-cli/v3/internal/model" "github.com/ooni/probe-cli/v3/internal/netxlite" - "github.com/ooni/probe-cli/v3/internal/runtimex" ) // Redirect to types defined inside the model package @@ -22,23 +21,22 @@ type ( // Control performs the control request and returns the response. func Control( ctx context.Context, sess model.ExperimentSession, - testhelpers []model.OOAPIService, creq ControlRequest) (ControlResponse, *model.OOAPIService, error) { - seqCaller := httpapi.NewSequenceCaller( - httpapi.MustNewPOSTJSONWithJSONResponseDescriptor(sess.Logger(), "/", creq).WithBodyLogging(true), - httpapi.NewEndpointList(sess.DefaultHTTPClient(), sess.UserAgent(), testhelpers...)..., - ) - sess.Logger().Infof("control for %s...", creq.HTTPRequest) - var out ControlResponse - idx, err := seqCaller.CallWithJSONResponse(ctx, &out) - sess.Logger().Infof("control for %s... %+v", creq.HTTPRequest, model.ErrorToStringOrOK(err)) - if err != nil { - // make sure error is wrapped - err = netxlite.NewTopLevelGenericErrWrapper(err) - return ControlResponse{}, nil, err + thAddr string, creq ControlRequest) (out ControlResponse, err error) { + clnt := &httpx.APIClientTemplate{ + BaseURL: thAddr, + HTTPClient: sess.DefaultHTTPClient(), + Logger: sess.Logger(), + UserAgent: sess.UserAgent(), } + sess.Logger().Infof("control for %s...", creq.HTTPRequest) + // make sure error is wrapped + err = clnt.WithBodyLogging().Build().PostJSON(ctx, "/", creq, &out) + if err != nil { + err = netxlite.NewTopLevelGenericErrWrapper(err) + } + sess.Logger().Infof("control for %s... %+v", creq.HTTPRequest, model.ErrorToStringOrOK(err)) fillASNs(&out.DNS) - runtimex.Assert(idx >= 0 && idx < len(testhelpers), "idx out of bounds") - return out, &testhelpers[idx], nil + return } // fillASNs fills the ASNs array of ControlDNSResult. For each Addr inside diff --git a/internal/engine/experiment/webconnectivity/webconnectivity.go b/internal/engine/experiment/webconnectivity/webconnectivity.go index 6792311..526b9f7 100644 --- a/internal/engine/experiment/webconnectivity/webconnectivity.go +++ b/internal/engine/experiment/webconnectivity/webconnectivity.go @@ -15,7 +15,7 @@ import ( const ( testName = "web_connectivity" - testVersion = "0.4.2" + testVersion = "0.4.1" ) // Config contains the experiment config. @@ -121,11 +121,12 @@ const ( ) // Run implements ExperimentMeasurer.Run. -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m Measurer) Run( + ctx context.Context, + sess model.ExperimentSession, + measurement *model.Measurement, + callbacks model.ExperimentCallbacks, +) error { ctx, cancel := context.WithTimeout(ctx, 60*time.Second) defer cancel() tk := new(TestKeys) @@ -144,9 +145,19 @@ func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { } // 1. find test helper testhelpers, _ := sess.GetTestHelpersByName("web-connectivity") - if len(testhelpers) < 1 { + var testhelper *model.OOAPIService + for _, th := range testhelpers { + if th.Type == "https" { + testhelper = &th + break + } + } + if testhelper == nil { return ErrNoAvailableTestHelpers } + measurement.TestHelpers = map[string]interface{}{ + "backend": testhelper, + } // 2. perform the DNS lookup step dnsBegin := time.Now() dnsResult := DNSLookup(ctx, DNSLookupConfig{ @@ -156,11 +167,10 @@ func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { tk.Queries = append(tk.Queries, dnsResult.TestKeys.Queries...) tk.DNSExperimentFailure = dnsResult.Failure epnts := NewEndpoints(URL, dnsResult.Addresses()) - sess.Logger().Infof("using control: %+v", testhelpers) + sess.Logger().Infof("using control: %s", testhelper.Address) // 3. perform the control measurement thBegin := time.Now() - var usedTH *model.OOAPIService - tk.Control, usedTH, err = Control(ctx, sess, testhelpers, ControlRequest{ + tk.Control, err = Control(ctx, sess, testhelper.Address, ControlRequest{ HTTPRequest: URL.String(), HTTPRequestHeaders: map[string][]string{ "Accept": {model.HTTPHeaderAccept}, @@ -169,11 +179,6 @@ func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { }, TCPConnect: epnts.Endpoints(), }) - if usedTH != nil { - measurement.TestHelpers = map[string]interface{}{ - "backend": usedTH, - } - } tk.THRuntime = time.Since(thBegin) tk.ControlFailure = tracex.NewFailure(err) // 4. analyze DNS results diff --git a/internal/engine/experiment/webconnectivity/webconnectivity_test.go b/internal/engine/experiment/webconnectivity/webconnectivity_test.go index 6b53693..c1fcb37 100644 --- a/internal/engine/experiment/webconnectivity/webconnectivity_test.go +++ b/internal/engine/experiment/webconnectivity/webconnectivity_test.go @@ -21,7 +21,7 @@ func TestNewExperimentMeasurer(t *testing.T) { if measurer.ExperimentName() != "web_connectivity" { t.Fatal("unexpected name") } - if measurer.ExperimentVersion() != "0.4.2" { + if measurer.ExperimentVersion() != "0.4.1" { t.Fatal("unexpected version") } } @@ -37,12 +37,7 @@ func TestSuccess(t *testing.T) { sess := newsession(t, true) measurement := &model.Measurement{Input: "http://www.example.com"} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -70,12 +65,7 @@ func TestMeasureWithCancelledContext(t *testing.T) { sess := newsession(t, true) measurement := &model.Measurement{Input: "http://www.example.com"} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := measurer.Run(ctx, args); err != nil { + if err := measurer.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } tk := measurement.TestKeys.(*webconnectivity.TestKeys) @@ -109,12 +99,7 @@ func TestMeasureWithNoInput(t *testing.T) { sess := newsession(t, true) measurement := &model.Measurement{Input: ""} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, webconnectivity.ErrNoInput) { t.Fatal(err) } @@ -142,12 +127,7 @@ func TestMeasureWithInputNotBeingAnURL(t *testing.T) { sess := newsession(t, true) measurement := &model.Measurement{Input: "\t\t\t\t\t\t"} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, webconnectivity.ErrInputIsNotAnURL) { t.Fatal(err) } @@ -175,12 +155,7 @@ func TestMeasureWithUnsupportedInput(t *testing.T) { sess := newsession(t, true) measurement := &model.Measurement{Input: "dnslookup://example.com"} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, webconnectivity.ErrUnsupportedInput) { t.Fatal(err) } @@ -208,12 +183,7 @@ func TestMeasureWithNoAvailableTestHelpers(t *testing.T) { sess := newsession(t, false) measurement := &model.Measurement{Input: "https://www.example.com"} callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if !errors.Is(err, webconnectivity.ErrNoAvailableTestHelpers) { t.Fatal(err) } diff --git a/internal/engine/experiment/whatsapp/whatsapp.go b/internal/engine/experiment/whatsapp/whatsapp.go index 1a1e3bf..24874e2 100644 --- a/internal/engine/experiment/whatsapp/whatsapp.go +++ b/internal/engine/experiment/whatsapp/whatsapp.go @@ -154,11 +154,10 @@ func (m Measurer) ExperimentVersion() string { } // Run implements ExperimentMeasurer.Run -func (m Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session - +func (m Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { ctx, cancel := context.WithTimeout(ctx, 60*time.Second) defer cancel() urlgetter.RegisterExtensions(measurement) diff --git a/internal/engine/experiment/whatsapp/whatsapp_test.go b/internal/engine/experiment/whatsapp/whatsapp_test.go index c6bb573..2162086 100644 --- a/internal/engine/experiment/whatsapp/whatsapp_test.go +++ b/internal/engine/experiment/whatsapp/whatsapp_test.go @@ -35,12 +35,7 @@ func TestSuccess(t *testing.T) { sess := &mockable.Session{MockableLogger: log.Log} measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -75,12 +70,7 @@ func TestFailureAllEndpoints(t *testing.T) { sess := &mockable.Session{MockableLogger: log.Log} measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - err := measurer.Run(ctx, args) + err := measurer.Run(ctx, sess, measurement, callbacks) if err != nil { t.Fatal(err) } @@ -608,12 +598,7 @@ func TestWeConfigureWebChecksCorrectly(t *testing.T) { } measurement := new(model.Measurement) callbacks := model.NewPrinterCallbacks(log.Log) - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err := measurer.Run(ctx, args); err != nil { + if err := measurer.Run(ctx, sess, measurement, callbacks); err != nil { t.Fatal(err) } if called.Load() != 263 { diff --git a/internal/engine/experiment_integration_test.go b/internal/engine/experiment_integration_test.go index 865a452..449edc6 100644 --- a/internal/engine/experiment_integration_test.go +++ b/internal/engine/experiment_integration_test.go @@ -475,7 +475,10 @@ func (am *antaniMeasurer) ExperimentVersion() string { return "0.1.1" } -func (am *antaniMeasurer) Run(ctx context.Context, args *model.ExperimentArgs) error { +func (am *antaniMeasurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { return nil } diff --git a/internal/experiment/webconnectivity/cleartextflow.go b/internal/experiment/webconnectivity/cleartextflow.go index d515d4e..351a27c 100644 --- a/internal/experiment/webconnectivity/cleartextflow.go +++ b/internal/experiment/webconnectivity/cleartextflow.go @@ -285,7 +285,7 @@ func (t *CleartextFlow) maybeFollowRedirects(ctx context.Context, resp *http.Res WaitGroup: t.WaitGroup, Referer: resp.Request.URL.String(), Session: nil, // no need to issue another control request - TestHelpers: nil, // ditto + THAddr: "", // ditto UDPAddress: t.UDPAddress, } resolvers.Start(ctx) diff --git a/internal/experiment/webconnectivity/control.go b/internal/experiment/webconnectivity/control.go index bf7003a..3e58f42 100644 --- a/internal/experiment/webconnectivity/control.go +++ b/internal/experiment/webconnectivity/control.go @@ -8,11 +8,10 @@ import ( "time" "github.com/ooni/probe-cli/v3/internal/engine/experiment/webconnectivity" - "github.com/ooni/probe-cli/v3/internal/httpapi" + "github.com/ooni/probe-cli/v3/internal/httpx" "github.com/ooni/probe-cli/v3/internal/measurexlite" "github.com/ooni/probe-cli/v3/internal/model" "github.com/ooni/probe-cli/v3/internal/netxlite" - "github.com/ooni/probe-cli/v3/internal/runtimex" ) // EndpointMeasurementsStarter is used by Control to start extra @@ -52,8 +51,8 @@ type Control struct { // Session is the MANDATORY session to use. Session model.ExperimentSession - // TestHelpers is the MANDATORY list of test helpers. - TestHelpers []model.OOAPIService + // THAddr is the MANDATORY TH's URL. + THAddr string // URL is the MANDATORY URL we are measuring. URL *url.URL @@ -103,20 +102,26 @@ func (c *Control) Run(parentCtx context.Context) { // create logger for this operation ol := measurexlite.NewOperationLogger( c.Logger, - "control for %s using %+v", + "control for %s using %s", creq.HTTPRequest, - c.TestHelpers, + c.THAddr, ) - // create an httpapi sequence caller - seqCaller := httpapi.NewSequenceCaller( - httpapi.MustNewPOSTJSONWithJSONResponseDescriptor(c.Logger, "/", creq).WithBodyLogging(true), - httpapi.NewEndpointList(c.Session.DefaultHTTPClient(), c.Session.UserAgent(), c.TestHelpers...)..., - ) + // create an API client + clnt := (&httpx.APIClientTemplate{ + Accept: "", + Authorization: "", + BaseURL: c.THAddr, + HTTPClient: c.Session.DefaultHTTPClient(), + Host: "", // use the one inside the URL + LogBody: true, + Logger: c.Logger, + UserAgent: c.Session.UserAgent(), + }).Build() // issue the control request and wait for the response var cresp webconnectivity.ControlResponse - idx, err := seqCaller.CallWithJSONResponse(opCtx, &cresp) + err := clnt.PostJSON(opCtx, "/", creq, &cresp) if err != nil { // make sure error is wrapped err = netxlite.NewTopLevelGenericErrWrapper(err) @@ -129,10 +134,6 @@ func (c *Control) Run(parentCtx context.Context) { c.TestKeys.SetControl(&cresp) ol.Stop(nil) - // record the specific TH that worked - runtimex.Assert(idx >= 0 && idx < len(c.TestHelpers), "idx out of bounds") - c.TestKeys.setTestHelper(&c.TestHelpers[idx]) - // if the TH returned us addresses we did not previously were // aware of, make sure we also measure them c.maybeStartExtraMeasurements(parentCtx, cresp.DNS.Addrs) diff --git a/internal/experiment/webconnectivity/dnsresolvers.go b/internal/experiment/webconnectivity/dnsresolvers.go index 8709635..0c9df8d 100644 --- a/internal/experiment/webconnectivity/dnsresolvers.go +++ b/internal/experiment/webconnectivity/dnsresolvers.go @@ -67,9 +67,8 @@ type DNSResolvers struct { // always follow the redirect chain caused by the provided URL. Session model.ExperimentSession - // TestHelpers is the OPTIONAL list of test helpers. If the list is - // empty, we are not going to try to contact any test helper. - TestHelpers []model.OOAPIService + // THAddr is the OPTIONAL test helper address. + THAddr string // UDPAddress is the OPTIONAL address of the UDP resolver to use. If this // field is not set we use a default one (e.g., `8.8.8.8:53`). @@ -499,15 +498,15 @@ func (t *DNSResolvers) startSecureFlows( } } -// maybeStartControlFlow starts the control flow iff .Session and .TestHelpers are set. +// maybeStartControlFlow starts the control flow iff .Session and .THAddr are set. func (t *DNSResolvers) maybeStartControlFlow( ctx context.Context, ps *prioritySelector, addresses []DNSEntry, ) { - // note: for subsequent requests we don't set .Session and .TestHelpers hence + // note: for subsequent requests we don't set .Session and .THAddr hence // we are not going to query the test helper more than once - if t.Session != nil && len(t.TestHelpers) > 0 { + if t.Session != nil && t.THAddr != "" { var addrs []string for _, addr := range addresses { addrs = append(addrs, addr.Addr) @@ -519,7 +518,7 @@ func (t *DNSResolvers) maybeStartControlFlow( PrioSelector: ps, TestKeys: t.TestKeys, Session: t.Session, - TestHelpers: t.TestHelpers, + THAddr: t.THAddr, URL: t.URL, WaitGroup: t.WaitGroup, } diff --git a/internal/experiment/webconnectivity/measurer.go b/internal/experiment/webconnectivity/measurer.go index dc520cf..c6731e4 100644 --- a/internal/experiment/webconnectivity/measurer.go +++ b/internal/experiment/webconnectivity/measurer.go @@ -36,17 +36,15 @@ func (m *Measurer) ExperimentName() string { // ExperimentVersion implements model.ExperimentMeasurer. func (m *Measurer) ExperimentVersion() string { - return "0.5.19" + return "0.5.18" } // Run implements model.ExperimentMeasurer. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { +func (m *Measurer) Run(ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks) error { // Reminder: when this function returns an error, the measurement result // WILL NOT be submitted to the OONI backend. You SHOULD only return an error // for fundamental errors (e.g., the input is invalid or missing). - _ = args.Callbacks - measurement := args.Measurement - sess := args.Session // make sure we have a cancellable context such that we can stop any // goroutine running in the background (e.g., priority.go's ones) @@ -91,7 +89,17 @@ func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { // obtain the test helper's address testhelpers, _ := sess.GetTestHelpersByName("web-connectivity") - if len(testhelpers) < 1 { + var thAddr string + for _, th := range testhelpers { + if th.Type == "https" { + thAddr = th.Address + measurement.TestHelpers = map[string]any{ + "backend": &th, + } + break + } + } + if thAddr == "" { sess.Logger().Warnf("continuing without a valid TH address") tk.SetControlFailure(webconnectivity.ErrNoAvailableTestHelpers) } @@ -112,7 +120,7 @@ func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { CookieJar: jar, Referer: "", Session: sess, - TestHelpers: testhelpers, + THAddr: thAddr, UDPAddress: "", } resos.Start(ctx) @@ -129,16 +137,6 @@ func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { // perform any deferred computation on the test keys tk.Finalize(sess.Logger()) - // set the test helper we used - // TODO(bassosimone): it may be more informative to know about all the - // test helpers we _tried_ to use, however the data format does not have - // support for that as far as I can tell... - if th := tk.getTestHelper(); th != nil { - measurement.TestHelpers = map[string]interface{}{ - "backend": th, - } - } - // return whether there was a fundamental failure, which would prevent // the measurement from being submitted to the OONI collector. return tk.fundamentalFailure diff --git a/internal/experiment/webconnectivity/secureflow.go b/internal/experiment/webconnectivity/secureflow.go index 6a47b73..68f41e1 100644 --- a/internal/experiment/webconnectivity/secureflow.go +++ b/internal/experiment/webconnectivity/secureflow.go @@ -337,7 +337,7 @@ func (t *SecureFlow) maybeFollowRedirects(ctx context.Context, resp *http.Respon WaitGroup: t.WaitGroup, Referer: resp.Request.URL.String(), Session: nil, // no need to issue another control request - TestHelpers: nil, // ditto + THAddr: "", // ditto UDPAddress: t.UDPAddress, } resolvers.Start(ctx) diff --git a/internal/experiment/webconnectivity/testkeys.go b/internal/experiment/webconnectivity/testkeys.go index 5cd917f..37dbc94 100644 --- a/internal/experiment/webconnectivity/testkeys.go +++ b/internal/experiment/webconnectivity/testkeys.go @@ -134,10 +134,6 @@ type TestKeys struct { // mu provides mutual exclusion for accessing the test keys. mu *sync.Mutex - - // testHelper is used to communicate the TH that worked to the main - // goroutine such that we can fill measurement.TestHelpers. - testHelper *model.OOAPIService } // ConnPriorityLogEntry is an entry in the TestKeys.ConnPriorityLog slice. @@ -306,21 +302,6 @@ func (tk *TestKeys) AppendConnPriorityLogEntry(entry *ConnPriorityLogEntry) { tk.mu.Unlock() } -// setTestHelper sets .testHelper in a thread safe way -func (tk *TestKeys) setTestHelper(th *model.OOAPIService) { - tk.mu.Lock() - tk.testHelper = th - tk.mu.Unlock() -} - -// getTestHelper gets .testHelper in a thread safe way -func (tk *TestKeys) getTestHelper() (th *model.OOAPIService) { - tk.mu.Lock() - th = tk.testHelper - tk.mu.Unlock() - return -} - // NewTestKeys creates a new instance of TestKeys. func NewTestKeys() *TestKeys { return &TestKeys{ @@ -367,7 +348,6 @@ func NewTestKeys() *TestKeys { ControlRequest: nil, fundamentalFailure: nil, mu: &sync.Mutex{}, - testHelper: nil, } } diff --git a/internal/httpapi/call.go b/internal/httpapi/call.go deleted file mode 100644 index 1228080..0000000 --- a/internal/httpapi/call.go +++ /dev/null @@ -1,181 +0,0 @@ -package httpapi - -// -// Calling HTTP APIs. -// - -import ( - "bytes" - "context" - "encoding/json" - "fmt" - "io" - "net/http" - "net/url" - "strings" - - "github.com/ooni/probe-cli/v3/internal/netxlite" -) - -// joinURLPath appends |resourcePath| to |urlPath|. -func joinURLPath(urlPath, resourcePath string) string { - if resourcePath == "" { - if urlPath == "" { - return "/" - } - return urlPath - } - if !strings.HasSuffix(urlPath, "/") { - urlPath += "/" - } - resourcePath = strings.TrimPrefix(resourcePath, "/") - return urlPath + resourcePath -} - -// newRequest creates a new http.Request from the given |ctx|, |endpoint|, and |desc|. -func newRequest(ctx context.Context, endpoint *Endpoint, desc *Descriptor) (*http.Request, error) { - URL, err := url.Parse(endpoint.BaseURL) - if err != nil { - return nil, err - } - // BaseURL and resource URL are joined if they have a path - URL.Path = joinURLPath(URL.Path, desc.URLPath) - if len(desc.URLQuery) > 0 { - URL.RawQuery = desc.URLQuery.Encode() - } else { - URL.RawQuery = "" // as documented we only honour desc.URLQuery - } - var reqBody io.Reader - if len(desc.RequestBody) > 0 { - reqBody = bytes.NewReader(desc.RequestBody) - desc.Logger.Debugf("httpapi: request body length: %d", len(desc.RequestBody)) - if desc.LogBody { - desc.Logger.Debugf("httpapi: request body: %s", string(desc.RequestBody)) - } - } - request, err := http.NewRequestWithContext(ctx, desc.Method, URL.String(), reqBody) - if err != nil { - return nil, err - } - request.Host = endpoint.Host // allow cloudfronting - if desc.Authorization != "" { - request.Header.Set("Authorization", desc.Authorization) - } - if desc.ContentType != "" { - request.Header.Set("Content-Type", desc.ContentType) - } - if desc.Accept != "" { - request.Header.Set("Accept", desc.Accept) - } - if endpoint.UserAgent != "" { - request.Header.Set("User-Agent", endpoint.UserAgent) - } - return request, nil -} - -// ErrHTTPRequestFailed indicates that the server returned >= 400. -type ErrHTTPRequestFailed struct { - // StatusCode is the status code that failed. - StatusCode int -} - -// Error implements error. -func (err *ErrHTTPRequestFailed) Error() string { - return fmt.Sprintf("httpapi: http request failed: %d", err.StatusCode) -} - -// errMaybeCensorship indicates that there was an error at the networking layer -// including, e.g., DNS, TCP connect, TLS. When we see this kind of error, we -// will consider retrying with another endpoint under the assumption that it -// may be that the current endpoint is censored. -type errMaybeCensorship struct { - // Err is the underlying error - Err error -} - -// Error implements error -func (err *errMaybeCensorship) Error() string { - return err.Err.Error() -} - -// Unwrap allows to get the underlying error -func (err *errMaybeCensorship) Unwrap() error { - return err.Err -} - -// docall calls the API represented by the given request |req| on the given |endpoint| -// and returns the response and its body or an error. -func docall(endpoint *Endpoint, desc *Descriptor, request *http.Request) (*http.Response, []byte, error) { - // Implementation note: remember to mark errors for which you want - // to retry with another endpoint using errMaybeCensorship. - response, err := endpoint.HTTPClient.Do(request) - if err != nil { - return nil, nil, &errMaybeCensorship{err} - } - defer response.Body.Close() - // Implementation note: always read and log the response body since - // it's quite useful to see the response JSON on API error. - r := io.LimitReader(response.Body, DefaultMaxBodySize) - data, err := netxlite.ReadAllContext(request.Context(), r) - if err != nil { - return response, nil, &errMaybeCensorship{err} - } - desc.Logger.Debugf("httpapi: response body length: %d bytes", len(data)) - if desc.LogBody { - desc.Logger.Debugf("httpapi: response body: %s", string(data)) - } - if response.StatusCode >= 400 { - return response, nil, &ErrHTTPRequestFailed{response.StatusCode} - } - return response, data, nil -} - -// call is like Call but also returns the response. -func call(ctx context.Context, desc *Descriptor, endpoint *Endpoint) (*http.Response, []byte, error) { - timeout := desc.Timeout - if timeout <= 0 { - timeout = DefaultCallTimeout // as documented - } - ctx, cancel := context.WithTimeout(ctx, timeout) - defer cancel() - request, err := newRequest(ctx, endpoint, desc) - if err != nil { - return nil, nil, err - } - return docall(endpoint, desc, request) -} - -// Call invokes the API described by |desc| on the given HTTP |endpoint| and -// returns the response body (as a slice of bytes) or an error. -// -// Note: this function returns ErrHTTPRequestFailed if the HTTP status code is -// greater or equal than 400. You could use errors.As to obtain a copy of the -// error that was returned and see for yourself the actual status code. -func Call(ctx context.Context, desc *Descriptor, endpoint *Endpoint) ([]byte, error) { - _, rawResponseBody, err := call(ctx, desc, endpoint) - return rawResponseBody, err -} - -// goodContentTypeForJSON tracks known-good content-types for JSON. If the content-type -// is not in this map, |CallWithJSONResponse| emits a warning message. -var goodContentTypeForJSON = map[string]bool{ - applicationJSON: true, -} - -// CallWithJSONResponse is like Call but also assumes that the response is a -// JSON body and attempts to parse it into the |response| field. -// -// Note: this function returns ErrHTTPRequestFailed if the HTTP status code is -// greater or equal than 400. You could use errors.As to obtain a copy of the -// error that was returned and see for yourself the actual status code. -func CallWithJSONResponse(ctx context.Context, desc *Descriptor, endpoint *Endpoint, response any) error { - httpResp, rawRespBody, err := call(ctx, desc, endpoint) - if err != nil { - return err - } - if ctype := httpResp.Header.Get("Content-Type"); !goodContentTypeForJSON[ctype] { - desc.Logger.Warnf("httpapi: unexpected content-type: %s", ctype) - // fallthrough - } - return json.Unmarshal(rawRespBody, response) -} diff --git a/internal/httpapi/call_test.go b/internal/httpapi/call_test.go deleted file mode 100644 index 8fd45d7..0000000 --- a/internal/httpapi/call_test.go +++ /dev/null @@ -1,1163 +0,0 @@ -package httpapi - -import ( - "context" - "errors" - "fmt" - "io" - "net/http" - "net/http/httptest" - "net/url" - "strings" - "syscall" - "testing" - - "github.com/google/go-cmp/cmp" - "github.com/ooni/probe-cli/v3/internal/model" - "github.com/ooni/probe-cli/v3/internal/model/mocks" - "github.com/ooni/probe-cli/v3/internal/netxlite" - "github.com/ooni/probe-cli/v3/internal/runtimex" -) - -func Test_joinURLPath(t *testing.T) { - tests := []struct { - name string - urlPath string - resourcePath string - want string - }{{ - name: "whole path inside urlPath and empty resourcePath", - urlPath: "/robots.txt", - resourcePath: "", - want: "/robots.txt", - }, { - name: "empty urlPath and slash-prefixed resourcePath", - urlPath: "", - resourcePath: "/foo", - want: "/foo", - }, { - name: "slash urlPath and slash-prefixed resourcePath", - urlPath: "/", - resourcePath: "/foo", - want: "/foo", - }, { - name: "empty urlPath and empty resourcePath", - urlPath: "", - resourcePath: "", - want: "/", - }, { - name: "non-slash-terminated urlPath and slash-prefixed resourcePath", - urlPath: "/foo", - resourcePath: "/bar", - want: "/foo/bar", - }, { - name: "slash-terminated urlPath and slash-prefixed resourcePath", - urlPath: "/foo/", - resourcePath: "/bar", - want: "/foo/bar", - }, { - name: "slash-terminated urlPath and non-slash-prefixed resourcePath", - urlPath: "/foo", - resourcePath: "bar", - want: "/foo/bar", - }} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := joinURLPath(tt.urlPath, tt.resourcePath) - if diff := cmp.Diff(tt.want, got); diff != "" { - t.Fatal(diff) - } - }) - } -} - -func Test_newRequest(t *testing.T) { - type args struct { - ctx context.Context - endpoint *Endpoint - desc *Descriptor - } - tests := []struct { - name string - args args - wantFn func(*testing.T, *http.Request) - wantErr error - }{{ - name: "url.Parse fails", - args: args{ - ctx: nil, - endpoint: &Endpoint{ - BaseURL: "\t\t\t", // does not parse! - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: "", - RequestBody: nil, - Timeout: 0, - URLPath: "", - URLQuery: nil, - }, - }, - wantFn: nil, - wantErr: errors.New(`parse "\t\t\t": net/url: invalid control character in URL`), - }, { - name: "http.NewRequestWithContext fails", - args: args{ - ctx: nil, // causes http.NewRequestWithContext to fail - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: "", - RequestBody: nil, - Timeout: 0, - URLPath: "", - URLQuery: nil, - }, - }, - wantFn: nil, - wantErr: errors.New("net/http: nil Context"), - }, { - name: "successful case with GET method, no body, and no extra headers", - args: args{ - ctx: context.Background(), - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: http.MethodGet, - RequestBody: nil, - Timeout: 0, - URLPath: "", - URLQuery: nil, - }, - }, - wantFn: func(t *testing.T, req *http.Request) { - if req == nil { - t.Fatal("expected non-nil request") - } - if req.Method != http.MethodGet { - t.Fatal("invalid method") - } - if req.URL.String() != "https://example.com/" { - t.Fatal("invalid URL") - } - if req.Body != nil { - t.Fatal("invalid body", req.Body) - } - }, - wantErr: nil, - }, { - name: "successful case with POST method and body", - args: args{ - ctx: context.Background(), - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: model.DiscardLogger, - MaxBodySize: 0, - Method: http.MethodPost, - RequestBody: []byte("deadbeef"), - Timeout: 0, - URLPath: "", - URLQuery: nil, - }, - }, - wantFn: func(t *testing.T, req *http.Request) { - if req == nil { - t.Fatal("expected non-nil request") - } - if req.Method != http.MethodPost { - t.Fatal("invalid method") - } - if req.URL.String() != "https://example.com/" { - t.Fatal("invalid URL") - } - data, err := netxlite.ReadAllContext(context.Background(), req.Body) - if err != nil { - t.Fatal(err) - } - if diff := cmp.Diff([]byte("deadbeef"), data); diff != "" { - t.Fatal(diff) - } - }, - wantErr: nil, - }, { - name: "with GET method and custom headers", - args: args{ - ctx: context.Background(), - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - HTTPClient: nil, - Host: "antani.org", - UserAgent: "httpclient/1.0.1", - }, - desc: &Descriptor{ - Accept: "application/json", - Authorization: "deafbeef", - ContentType: "text/plain", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: http.MethodPut, - RequestBody: nil, - Timeout: 0, - URLPath: "", - URLQuery: nil, - }, - }, - wantFn: func(t *testing.T, req *http.Request) { - if req == nil { - t.Fatal("expected non-nil request") - } - if req.Method != http.MethodPut { - t.Fatal("invalid method") - } - if req.Host != "antani.org" { - t.Fatal("invalid request host") - } - if req.URL.String() != "https://example.com/" { - t.Fatal("invalid URL") - } - if req.Header.Get("Authorization") != "deafbeef" { - t.Fatal("invalid authorization") - } - if req.Header.Get("Content-Type") != "text/plain" { - t.Fatal("invalid content-type") - } - if req.Header.Get("Accept") != "application/json" { - t.Fatal("invalid accept") - } - if req.Header.Get("User-Agent") != "httpclient/1.0.1" { - t.Fatal("invalid user-agent") - } - }, - wantErr: nil, - }, { - name: "we join the urlPath with the resourcePath", - args: args{ - ctx: context.Background(), - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/api/v1", - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: http.MethodGet, - RequestBody: nil, - Timeout: 0, - URLPath: "/test-list/urls", - URLQuery: nil, - }, - }, - wantFn: func(t *testing.T, req *http.Request) { - if req == nil { - t.Fatal("expected non-nil request") - } - if req.Method != http.MethodGet { - t.Fatal("invalid method") - } - if req.URL.String() != "https://www.example.com/api/v1/test-list/urls" { - t.Fatal("invalid URL") - } - }, - wantErr: nil, - }, { - name: "we discard any query element inside the Endpoint.BaseURL", - args: args{ - ctx: context.Background(), - endpoint: &Endpoint{ - BaseURL: "https://example.org/api/v1/?probe_cc=IT", - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: http.MethodGet, - RequestBody: nil, - Timeout: 0, - URLPath: "", - URLQuery: nil, - }, - }, - wantFn: func(t *testing.T, req *http.Request) { - if req == nil { - t.Fatal("expected non-nil request") - } - if req.Method != http.MethodGet { - t.Fatal("invalid method") - } - if req.URL.String() != "https://example.org/api/v1/" { - t.Fatal("invalid URL") - } - }, - wantErr: nil, - }, { - name: "we include query elements from Descriptor.URLQuery", - args: args{ - ctx: context.Background(), - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/api/v1/", - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: http.MethodGet, - RequestBody: nil, - Timeout: 0, - URLPath: "test-list/urls", - URLQuery: map[string][]string{ - "probe_cc": {"IT"}, - }, - }, - }, - wantFn: func(t *testing.T, req *http.Request) { - if req == nil { - t.Fatal("expected non-nil request") - } - if req.Method != http.MethodGet { - t.Fatal("invalid method") - } - if req.URL.String() != "https://www.example.com/api/v1/test-list/urls?probe_cc=IT" { - t.Fatal("invalid URL") - } - }, - wantErr: nil, - }, { - name: "with as many implicitly-initialized fields as possible", - args: args{ - ctx: context.Background(), - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - }, - desc: &Descriptor{}, - }, - wantFn: func(t *testing.T, req *http.Request) { - if req == nil { - t.Fatal("expected non-nil request") - } - if req.Method != http.MethodGet { - t.Fatal("invalid method") - } - if req.URL.String() != "https://example.com/" { - t.Fatal("invalid URL") - } - }, - wantErr: nil, - }} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got, err := newRequest(tt.args.ctx, tt.args.endpoint, tt.args.desc) - switch { - case err == nil && tt.wantErr == nil: - // nothing - case err != nil && tt.wantErr == nil: - t.Fatalf("expected error but got %s", err.Error()) - case err == nil && tt.wantErr != nil: - t.Fatalf("expected %s but got ", tt.wantErr.Error()) - case err.Error() == tt.wantErr.Error(): - // nothing - default: - t.Fatalf("expected %s but got %s", err.Error(), tt.wantErr.Error()) - } - if tt.wantFn != nil { - tt.wantFn(t, got) - return - } - if got != nil { - t.Fatal("got response with nil tt.wantFn") - } - }) - } -} - -func TestCall(t *testing.T) { - type args struct { - ctx context.Context - desc *Descriptor - endpoint *Endpoint - } - tests := []struct { - name string - args args - want []byte - wantErr error - errfn func(t *testing.T, err error) - }{{ - name: "newRequest fails", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: nil, - MaxBodySize: 0, - Method: "", - RequestBody: nil, - Timeout: 0, - URLPath: "", - URLQuery: nil, - }, - endpoint: &Endpoint{ - BaseURL: "\t\t\t", // causes newRequest to fail - HTTPClient: nil, - Host: "", - UserAgent: "", - }, - }, - want: nil, - wantErr: errors.New(`parse "\t\t\t": net/url: invalid control character in URL`), - errfn: nil, - }, { - name: "endpoint.HTTPClient.Do fails", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - }, - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF - }, - }, - }, - }, - want: nil, - wantErr: io.EOF, - errfn: func(t *testing.T, err error) { - var expect *errMaybeCensorship - if !errors.As(err, &expect) { - t.Fatal("unexpected error type") - } - }, - }, { - name: "reading body fails", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - }, - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Body: io.NopCloser(&mocks.Reader{ - MockRead: func(b []byte) (int, error) { - return 0, netxlite.ECONNRESET - }, - }), - } - return resp, nil - }, - }, - }, - }, - want: nil, - wantErr: errors.New(netxlite.FailureConnectionReset), - errfn: func(t *testing.T, err error) { - var expect *errMaybeCensorship - if !errors.As(err, &expect) { - t.Fatal("unexpected error type") - } - }, - }, { - name: "status code indicates failure", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - }, - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Body: io.NopCloser(strings.NewReader("deadbeef")), - StatusCode: 403, - } - return resp, nil - }, - }, - }, - }, - want: nil, - wantErr: errors.New("httpapi: http request failed: 403"), - errfn: func(t *testing.T, err error) { - var expect *ErrHTTPRequestFailed - if !errors.As(err, &expect) { - t.Fatal("invalid error type") - } - }, - }, { - name: "success with log body flag", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - LogBody: true, // as documented by this test's name - Logger: model.DiscardLogger, - Method: http.MethodGet, - }, - endpoint: &Endpoint{ - BaseURL: "https://example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Body: io.NopCloser(strings.NewReader("deadbeef")), - StatusCode: 200, - } - return resp, nil - }, - }, - }, - }, - want: []byte("deadbeef"), - wantErr: nil, - errfn: nil, - }} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got, err := Call(tt.args.ctx, tt.args.desc, tt.args.endpoint) - switch { - case err == nil && tt.wantErr == nil: - // nothing - case err != nil && tt.wantErr == nil: - t.Fatalf("expected error but got %s", err.Error()) - case err == nil && tt.wantErr != nil: - t.Fatalf("expected %s but got ", tt.wantErr.Error()) - case err.Error() == tt.wantErr.Error(): - // nothing - default: - t.Fatalf("expected %s but got %s", err.Error(), tt.wantErr.Error()) - } - if diff := cmp.Diff(tt.want, got); diff != "" { - t.Fatal(diff) - } - }) - } -} - -func TestCallWithJSONResponse(t *testing.T) { - type response struct { - Name string - Age int64 - } - expectedResponse := response{ - Name: "sbs", - Age: 99, - } - type args struct { - ctx context.Context - desc *Descriptor - endpoint *Endpoint - } - tests := []struct { - name string - args args - wantErr error - errfn func(*testing.T, error) - }{{ - name: "call fails", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - }, - endpoint: &Endpoint{ - BaseURL: "\t\t\t\t", // causes failure - }, - }, - wantErr: errors.New(`parse "\t\t\t\t": net/url: invalid control character in URL`), - errfn: nil, - }, { - name: "with error during httpClient.Do", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - }, - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/a", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF - }, - }, - }, - }, - wantErr: io.EOF, - errfn: func(t *testing.T, err error) { - var expect *errMaybeCensorship - if !errors.As(err, &expect) { - t.Fatal("invalid error type") - } - }, - }, { - name: "with error when reading the response body", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - }, - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/a", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Body: io.NopCloser(&mocks.Reader{ - MockRead: func(b []byte) (int, error) { - return 0, netxlite.ECONNRESET - }, - }), - StatusCode: 200, - } - return resp, nil - }, - }, - }, - }, - wantErr: errors.New(netxlite.FailureConnectionReset), - errfn: func(t *testing.T, err error) { - var expect *errMaybeCensorship - if !errors.As(err, &expect) { - t.Fatal("invalid error type") - } - }, - }, { - name: "with HTTP failure", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - }, - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/a", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Body: io.NopCloser(strings.NewReader(`{"Name": "sbs", "Age": 99}`)), - StatusCode: 400, - } - return resp, nil - }, - }, - }, - }, - wantErr: errors.New("httpapi: http request failed: 400"), - errfn: func(t *testing.T, err error) { - var expect *ErrHTTPRequestFailed - if !errors.As(err, &expect) { - t.Fatal("invalid error type") - } - }, - }, { - name: "with good response and missing header", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - }, - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/a", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Body: io.NopCloser(strings.NewReader(`{"Name": "sbs", "Age": 99}`)), - StatusCode: 200, - } - return resp, nil - }, - }, - }, - }, - wantErr: nil, - errfn: nil, - }, { - name: "with good response and good header", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - Logger: model.DiscardLogger, - }, - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/a", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Header: http.Header{ - "Content-Type": {"application/json"}, - }, - Body: io.NopCloser(strings.NewReader(`{"Name": "sbs", "Age": 99}`)), - StatusCode: 200, - } - return resp, nil - }, - }, - }, - }, - wantErr: nil, - errfn: nil, - }, { - name: "response is not JSON", - args: args{ - ctx: context.Background(), - desc: &Descriptor{ - LogBody: false, - Logger: model.DiscardLogger, - Method: http.MethodGet, - }, - endpoint: &Endpoint{ - BaseURL: "https://www.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - Header: http.Header{ - "Content-Type": {"application/json"}, - }, - Body: io.NopCloser(strings.NewReader(`{`)), // invalid JSON - StatusCode: 200, - } - return resp, nil - }, - }, - }, - }, - wantErr: errors.New("unexpected end of JSON input"), - errfn: nil, - }} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - var response response - err := CallWithJSONResponse(tt.args.ctx, tt.args.desc, tt.args.endpoint, &response) - switch { - case err == nil && tt.wantErr == nil: - if diff := cmp.Diff(expectedResponse, response); err != nil { - t.Fatal(diff) - } - case err != nil && tt.wantErr == nil: - t.Fatalf("expected error but got %s", err.Error()) - case err == nil && tt.wantErr != nil: - t.Fatalf("expected %s but got ", tt.wantErr.Error()) - case err.Error() == tt.wantErr.Error(): - // nothing - default: - t.Fatalf("expected %s but got %s", err.Error(), tt.wantErr.Error()) - } - if tt.errfn != nil { - tt.errfn(t, err) - } - }) - } -} - -func TestCallHonoursContext(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - cancel() // should fail HTTP request immediately - desc := &Descriptor{ - LogBody: false, - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/robots.txt", - } - endpoint := &Endpoint{ - BaseURL: "https://www.example.com/", - HTTPClient: http.DefaultClient, - UserAgent: model.HTTPHeaderUserAgent, - } - body, err := Call(ctx, desc, endpoint) - if !errors.Is(err, context.Canceled) { - t.Fatal("unexpected err", err) - } - if len(body) > 0 { - t.Fatal("expected zero-length body") - } -} - -func TestCallWithJSONResponseHonoursContext(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - cancel() // should fail HTTP request immediately - desc := &Descriptor{ - LogBody: false, - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/robots.txt", - } - endpoint := &Endpoint{ - BaseURL: "https://www.example.com/", - HTTPClient: http.DefaultClient, - UserAgent: model.HTTPHeaderUserAgent, - } - var resp url.URL - err := CallWithJSONResponse(ctx, desc, endpoint, &resp) - if !errors.Is(err, context.Canceled) { - t.Fatal("unexpected err", err) - } -} - -func TestDescriptorLogging(t *testing.T) { - - // This test was originally written for the httpx package and we have adapted it - // by keeping the ~same implementation with a custom callx function that converts - // the previous semantics of httpx to the new semantics of httpapi. - callx := func(baseURL string, logBody bool, logger model.Logger, request, response any) error { - desc := MustNewPOSTJSONWithJSONResponseDescriptor(logger, "/", request).WithBodyLogging(logBody) - runtimex.Assert(desc.LogBody == logBody, "desc.LogBody should be equal to logBody here") - endpoint := &Endpoint{ - BaseURL: baseURL, - HTTPClient: http.DefaultClient, - } - return CallWithJSONResponse(context.Background(), desc, endpoint, response) - } - - // we also needed to create a constructor for the logger - newlogger := func(logs chan string) model.Logger { - return &mocks.Logger{ - MockDebugf: func(format string, v ...interface{}) { - logs <- fmt.Sprintf(format, v...) - }, - MockWarnf: func(format string, v ...interface{}) { - logs <- fmt.Sprintf(format, v...) - }, - } - } - - t.Run("body logging enabled, 200 Ok, and without content-type", func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc( - func(w http.ResponseWriter, r *http.Request) { - w.Write([]byte("[]")) - }, - )) - logs := make(chan string, 1024) - defer server.Close() - var ( - input []string - output []string - ) - logger := newlogger(logs) - err := callx(server.URL, true, logger, input, &output) - var found int - close(logs) - for entry := range logs { - if strings.HasPrefix(entry, "httpapi: request body: ") { - // we expect this because body logging is enabled - found |= 1 << 0 - continue - } - if strings.HasPrefix(entry, "httpapi: response body: ") { - // we expect this because body logging is enabled - found |= 1 << 1 - continue - } - if strings.HasPrefix(entry, "httpapi: unexpected content-type: ") { - // we would expect this because the server does not send us any content-type - found |= 1 << 2 - continue - } - if strings.HasPrefix(entry, "httpapi: request body length: ") { - // we should see this because we sent a body - found |= 1 << 3 - continue - } - if strings.HasPrefix(entry, "httpapi: response body length: ") { - // we should see this because we receive a body - found |= 1 << 4 - continue - } - } - if found != (1<<0 | 1<<1 | 1<<2 | 1<<3 | 1<<4) { - t.Fatal("did not find the expected logs") - } - if err != nil { - t.Fatal(err) - } - }) - - t.Run("body logging enabled, 200 Ok, and with content-type", func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc( - func(w http.ResponseWriter, r *http.Request) { - w.Header().Add("content-type", "application/json") - w.Write([]byte("[]")) - }, - )) - logs := make(chan string, 1024) - defer server.Close() - var ( - input []string - output []string - ) - logger := newlogger(logs) - err := callx(server.URL, true, logger, input, &output) - var found int - close(logs) - for entry := range logs { - if strings.HasPrefix(entry, "httpapi: request body: ") { - // we expect this because body logging is enabled - found |= 1 << 0 - continue - } - if strings.HasPrefix(entry, "httpapi: response body: ") { - // we expect this because body logging is enabled - found |= 1 << 1 - continue - } - if strings.HasPrefix(entry, "httpapi: unexpected content-type: ") { - // we do not expect this because the server sends us a content-type - found |= 1 << 2 - continue - } - if strings.HasPrefix(entry, "httpapi: request body length: ") { - // we should see this because we sent a body - found |= 1 << 3 - continue - } - if strings.HasPrefix(entry, "httpapi: response body length: ") { - // we should see this because we receive a body - found |= 1 << 4 - continue - } - } - if found != (1<<0 | 1<<1 | 1<<3 | 1<<4) { - t.Fatal("did not find the expected logs") - } - if err != nil { - t.Fatal(err) - } - }) - - t.Run("body logging enabled and 401 Unauthorized", func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc( - func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(401) - w.Write([]byte("[]")) - }, - )) - logs := make(chan string, 1024) - defer server.Close() - var ( - input []string - output []string - ) - logger := newlogger(logs) - err := callx(server.URL, true, logger, input, &output) - var found int - close(logs) - for entry := range logs { - if strings.HasPrefix(entry, "httpapi: request body: ") { - // should occur because body logging is enabled - found |= 1 << 0 - continue - } - if strings.HasPrefix(entry, "httpapi: response body: ") { - // should occur because body logging is enabled - found |= 1 << 1 - continue - } - if strings.HasPrefix(entry, "httpapi: unexpected content-type: ") { - // note: this one should not occur because the code is 401 so we're not - // actually going to parse the JSON document - found |= 1 << 2 - continue - } - if strings.HasPrefix(entry, "httpapi: request body length: ") { - // we should see this because we send a body - found |= 1 << 3 - continue - } - if strings.HasPrefix(entry, "httpapi: response body length: ") { - // we should see this because we receive a body - found |= 1 << 4 - continue - } - } - if found != (1<<0 | 1<<1 | 1<<3 | 1<<4) { - t.Fatal("did not find the expected logs") - } - var failure *ErrHTTPRequestFailed - if !errors.As(err, &failure) || failure.StatusCode != 401 { - t.Fatal("unexpected err", err) - } - }) - - t.Run("body logging NOT enabled and 200 Ok", func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc( - func(w http.ResponseWriter, r *http.Request) { - w.Write([]byte("[]")) - }, - )) - logs := make(chan string, 1024) - defer server.Close() - var ( - input []string - output []string - ) - logger := newlogger(logs) - err := callx(server.URL, false, logger, input, &output) // no logging - var found int - close(logs) - for entry := range logs { - if strings.HasPrefix(entry, "httpapi: request body: ") { - // should not see it: body logging is disabled - found |= 1 << 0 - continue - } - if strings.HasPrefix(entry, "httpapi: response body: ") { - // should not see it: body logging is disabled - found |= 1 << 1 - continue - } - if strings.HasPrefix(entry, "httpapi: unexpected content-type: ") { - // this one should be logged ANYWAY because it's orthogonal to the - // body logging so we should see it also in this case. - found |= 1 << 2 - continue - } - if strings.HasPrefix(entry, "httpapi: request body length: ") { - // should see this because we send a body - found |= 1 << 3 - continue - } - if strings.HasPrefix(entry, "httpapi: response body length: ") { - // should see this because we're receiving a body - found |= 1 << 4 - continue - } - } - if found != (1<<2 | 1<<3 | 1<<4) { - t.Fatal("did not find the expected logs") - } - if err != nil { - t.Fatal(err) - } - }) - - t.Run("body logging NOT enabled and 401 Unauthorized", func(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc( - func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(401) - w.Write([]byte("[]")) - }, - )) - logs := make(chan string, 1024) - defer server.Close() - var ( - input []string - output []string - ) - logger := newlogger(logs) - err := callx(server.URL, false, logger, input, &output) // no logging - var found int - close(logs) - for entry := range logs { - if strings.HasPrefix(entry, "httpapi: request body: ") { - // should not see it: body logging is disabled - found |= 1 << 0 - continue - } - if strings.HasPrefix(entry, "httpapi: response body: ") { - // should not see it: body logging is disabled - found |= 1 << 1 - continue - } - if strings.HasPrefix(entry, "httpapi: unexpected content-type: ") { - // should not see it because we don't parse the body on 401 errors - found |= 1 << 2 - continue - } - if strings.HasPrefix(entry, "httpapi: request body length: ") { - // we send a body so we should see it - found |= 1 << 3 - continue - } - if strings.HasPrefix(entry, "httpapi: response body length: ") { - // we receive a body so we should see it - found |= 1 << 4 - continue - } - } - if found != (1<<3 | 1<<4) { - t.Fatal("did not find the expected logs") - } - var failure *ErrHTTPRequestFailed - if !errors.As(err, &failure) || failure.StatusCode != 401 { - t.Fatal("unexpected err", err) - } - }) -} - -func Test_errMaybeCensorship_Unwrap(t *testing.T) { - t.Run("for errors.Is", func(t *testing.T) { - var err error = &errMaybeCensorship{io.EOF} - if !errors.Is(err, io.EOF) { - t.Fatal("cannot unwrap") - } - }) - - t.Run("for errors.As", func(t *testing.T) { - var err error = &errMaybeCensorship{netxlite.ECONNRESET} - var syserr syscall.Errno - if !errors.As(err, &syserr) || syserr != netxlite.ECONNRESET { - t.Fatal("cannot unwrap") - } - }) -} diff --git a/internal/httpapi/descriptor.go b/internal/httpapi/descriptor.go deleted file mode 100644 index ed35e85..0000000 --- a/internal/httpapi/descriptor.go +++ /dev/null @@ -1,155 +0,0 @@ -package httpapi - -// -// HTTP API descriptor (e.g., GET /api/v1/test-list/urls) -// - -import ( - "encoding/json" - "net/http" - "net/url" - "time" - - "github.com/ooni/probe-cli/v3/internal/model" - "github.com/ooni/probe-cli/v3/internal/runtimex" -) - -// Descriptor contains the parameters for calling a given HTTP -// API (e.g., GET /api/v1/test-list/urls). -// -// The zero value of this struct is invalid. Please, fill all the -// fields marked as MANDATORY for correct initialization. -type Descriptor struct { - // Accept contains the OPTIONAL accept header. - Accept string - - // Authorization is the OPTIONAL authorization. - Authorization string - - // ContentType is the OPTIONAL content-type header. - ContentType string - - // LogBody OPTIONALLY enables logging bodies. - LogBody bool - - // Logger is the MANDATORY logger to use. - // - // For example, model.DiscardLogger. - Logger model.Logger - - // MaxBodySize is the OPTIONAL maximum response body size. If - // not set, we use the |DefaultMaxBodySize| constant. - MaxBodySize int64 - - // Method is the MANDATORY request method. - Method string - - // RequestBody is the OPTIONAL request body. - RequestBody []byte - - // Timeout is the OPTIONAL timeout for this call. If no timeout - // is specified we will use the |DefaultCallTimeout| const. - Timeout time.Duration - - // URLPath is the MANDATORY URL path. - URLPath string - - // URLQuery is the OPTIONAL query. - URLQuery url.Values -} - -// WithBodyLogging returns a SHALLOW COPY of |Descriptor| with LogBody set to |value|. You SHOULD -// only use this method when initializing the descriptor you want to use. -func (desc *Descriptor) WithBodyLogging(value bool) *Descriptor { - out := &Descriptor{} - *out = *desc - out.LogBody = value - return out -} - -// DefaultMaxBodySize is the default value for the maximum -// body size you can fetch using the httpapi package. -const DefaultMaxBodySize = 1 << 22 - -// DefaultCallTimeout is the default timeout for an httpapi call. -const DefaultCallTimeout = 60 * time.Second - -// NewGETJSONDescriptor is a convenience factory for creating a new descriptor -// that uses the GET method and expects a JSON response. -func NewGETJSONDescriptor(logger model.Logger, urlPath string) *Descriptor { - return NewGETJSONWithQueryDescriptor(logger, urlPath, url.Values{}) -} - -// applicationJSON is the content-type for JSON -const applicationJSON = "application/json" - -// NewGETJSONWithQueryDescriptor is like NewGETJSONDescriptor but it also -// allows you to provide |query| arguments. Leaving |query| nil or empty -// is equivalent to calling NewGETJSONDescriptor directly. -func NewGETJSONWithQueryDescriptor(logger model.Logger, urlPath string, query url.Values) *Descriptor { - return &Descriptor{ - Accept: applicationJSON, - Authorization: "", - ContentType: "", - LogBody: false, - Logger: logger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodGet, - RequestBody: nil, - Timeout: DefaultCallTimeout, - URLPath: urlPath, - URLQuery: query, - } -} - -// NewPOSTJSONWithJSONResponseDescriptor creates a descriptor that POSTs a JSON document -// and expects to receive back a JSON document from the API. -// -// This function ONLY fails if we cannot serialize the |request| to JSON. So, if you know -// that |request| is JSON-serializable, you can safely call MustNewPostJSONWithJSONResponseDescriptor instead. -func NewPOSTJSONWithJSONResponseDescriptor(logger model.Logger, urlPath string, request any) (*Descriptor, error) { - rawRequest, err := json.Marshal(request) - if err != nil { - return nil, err - } - desc := &Descriptor{ - Accept: applicationJSON, - Authorization: "", - ContentType: applicationJSON, - LogBody: false, - Logger: logger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodPost, - RequestBody: rawRequest, - Timeout: DefaultCallTimeout, - URLPath: urlPath, - URLQuery: nil, - } - return desc, nil -} - -// MustNewPOSTJSONWithJSONResponseDescriptor is like NewPOSTJSONWithJSONResponseDescriptor except that -// it panics in case it's not possible to JSON serialize the |request|. -func MustNewPOSTJSONWithJSONResponseDescriptor(logger model.Logger, urlPath string, request any) *Descriptor { - desc, err := NewPOSTJSONWithJSONResponseDescriptor(logger, urlPath, request) - runtimex.PanicOnError(err, "NewPOSTJSONWithJSONResponseDescriptor failed") - return desc -} - -// NewGETResourceDescriptor creates a generic descriptor for GETting a -// resource of unspecified type using the given |urlPath|. -func NewGETResourceDescriptor(logger model.Logger, urlPath string) *Descriptor { - return &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: logger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodGet, - RequestBody: nil, - Timeout: DefaultCallTimeout, - URLPath: urlPath, - URLQuery: url.Values{}, - } -} diff --git a/internal/httpapi/descriptor_test.go b/internal/httpapi/descriptor_test.go deleted file mode 100644 index c3c58c0..0000000 --- a/internal/httpapi/descriptor_test.go +++ /dev/null @@ -1,248 +0,0 @@ -package httpapi - -import ( - "log" - "net/http" - "net/url" - "testing" - "time" - - "github.com/google/go-cmp/cmp" - "github.com/ooni/probe-cli/v3/internal/model" -) - -func TestDescriptor_WithBodyLogging(t *testing.T) { - type fields struct { - Accept string - Authorization string - ContentType string - LogBody bool - Logger model.Logger - MaxBodySize int64 - Method string - RequestBody []byte - Timeout time.Duration - URLPath string - URLQuery url.Values - } - tests := []struct { - name string - fields fields - want *Descriptor - }{{ - name: "with empty fields", - fields: fields{}, // LogBody defaults to false - want: &Descriptor{ - LogBody: true, - }, - }, { - name: "with nonempty fields", - fields: fields{ - Accept: "xx", - Authorization: "y", - ContentType: "zzz", - LogBody: false, // obviously must be false - Logger: model.DiscardLogger, - MaxBodySize: 123, - Method: "POST", - RequestBody: []byte("123"), - Timeout: 15555, - URLPath: "/", - URLQuery: map[string][]string{ - "a": {"b"}, - }, - }, - want: &Descriptor{ - Accept: "xx", - Authorization: "y", - ContentType: "zzz", - LogBody: true, - Logger: model.DiscardLogger, - MaxBodySize: 123, - Method: "POST", - RequestBody: []byte("123"), - Timeout: 15555, - URLPath: "/", - URLQuery: map[string][]string{ - "a": {"b"}, - }, - }, - }} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - desc := &Descriptor{ - Accept: tt.fields.Accept, - Authorization: tt.fields.Authorization, - ContentType: tt.fields.ContentType, - LogBody: tt.fields.LogBody, - Logger: tt.fields.Logger, - MaxBodySize: tt.fields.MaxBodySize, - Method: tt.fields.Method, - RequestBody: tt.fields.RequestBody, - Timeout: tt.fields.Timeout, - URLPath: tt.fields.URLPath, - URLQuery: tt.fields.URLQuery, - } - got := desc.WithBodyLogging(true) - if diff := cmp.Diff(tt.want, got); diff != "" { - t.Fatal(diff) - } - }) - } -} - -func TestNewGetJSONDescriptor(t *testing.T) { - expected := &Descriptor{ - Accept: "application/json", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: model.DiscardLogger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodGet, - RequestBody: nil, - Timeout: DefaultCallTimeout, - URLPath: "/robots.txt", - URLQuery: url.Values{}, - } - got := NewGETJSONDescriptor(model.DiscardLogger, "/robots.txt") - if diff := cmp.Diff(expected, got); diff != "" { - t.Fatal(diff) - } -} - -func TestNewGetJSONWithQueryDescriptor(t *testing.T) { - query := url.Values{ - "a": {"b"}, - "c": {"d"}, - } - expected := &Descriptor{ - Accept: "application/json", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: model.DiscardLogger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodGet, - RequestBody: nil, - Timeout: DefaultCallTimeout, - URLPath: "/robots.txt", - URLQuery: query, - } - got := NewGETJSONWithQueryDescriptor(model.DiscardLogger, "/robots.txt", query) - if diff := cmp.Diff(expected, got); diff != "" { - t.Fatal(diff) - } -} - -func TestNewPOSTJSONWithJSONResponseDescriptor(t *testing.T) { - type request struct { - Name string - Age int64 - } - - t.Run("with failure", func(t *testing.T) { - request := make(chan int64) - got, err := NewPOSTJSONWithJSONResponseDescriptor(model.DiscardLogger, "/robots.txt", request) - if err == nil || err.Error() != "json: unsupported type: chan int64" { - log.Fatal("unexpected err", err) - } - if got != nil { - log.Fatal("expected to get a nil Descriptor") - } - }) - - t.Run("with success", func(t *testing.T) { - request := request{ - Name: "sbs", - Age: 99, - } - expected := &Descriptor{ - Accept: "application/json", - Authorization: "", - ContentType: "application/json", - LogBody: false, - Logger: model.DiscardLogger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodPost, - RequestBody: []byte(`{"Name":"sbs","Age":99}`), - Timeout: DefaultCallTimeout, - URLPath: "/robots.txt", - URLQuery: nil, - } - got, err := NewPOSTJSONWithJSONResponseDescriptor(model.DiscardLogger, "/robots.txt", request) - if err != nil { - log.Fatal(err) - } - if diff := cmp.Diff(expected, got); diff != "" { - t.Fatal(diff) - } - }) -} - -func TestMustNewPOSTJSONWithJSONResponseDescriptor(t *testing.T) { - type request struct { - Name string - Age int64 - } - - t.Run("with failure", func(t *testing.T) { - var panicked bool - func() { - defer func() { - if r := recover(); r != nil { - panicked = true - } - }() - request := make(chan int64) - _ = MustNewPOSTJSONWithJSONResponseDescriptor(model.DiscardLogger, "/robots.txt", request) - }() - if !panicked { - t.Fatal("did not panic") - } - }) - - t.Run("with success", func(t *testing.T) { - request := request{ - Name: "sbs", - Age: 99, - } - expected := &Descriptor{ - Accept: "application/json", - Authorization: "", - ContentType: "application/json", - LogBody: false, - Logger: model.DiscardLogger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodPost, - RequestBody: []byte(`{"Name":"sbs","Age":99}`), - Timeout: DefaultCallTimeout, - URLPath: "/robots.txt", - URLQuery: nil, - } - got := MustNewPOSTJSONWithJSONResponseDescriptor(model.DiscardLogger, "/robots.txt", request) - if diff := cmp.Diff(expected, got); diff != "" { - t.Fatal(diff) - } - }) -} - -func TestNewGetResourceDescriptor(t *testing.T) { - expected := &Descriptor{ - Accept: "", - Authorization: "", - ContentType: "", - LogBody: false, - Logger: model.DiscardLogger, - MaxBodySize: DefaultMaxBodySize, - Method: http.MethodGet, - RequestBody: nil, - Timeout: DefaultCallTimeout, - URLPath: "/robots.txt", - URLQuery: url.Values{}, - } - got := NewGETResourceDescriptor(model.DiscardLogger, "/robots.txt") - if diff := cmp.Diff(expected, got); diff != "" { - t.Fatal(diff) - } -} diff --git a/internal/httpapi/doc.go b/internal/httpapi/doc.go deleted file mode 100644 index 0c1361c..0000000 --- a/internal/httpapi/doc.go +++ /dev/null @@ -1,15 +0,0 @@ -// Package httpapi contains code for calling HTTP APIs. -// -// We model HTTP APIs as follows: -// -// 1. |Endpoint| is an API endpoint (e.g., https://api.ooni.io); -// -// 2. |Descriptor| describes the specific API you want to use (e.g., -// GET /api/v1/test-list/urls with JSON response body). -// -// Generally, you use |Call| to call the API identified by a |Descriptor| -// on the specified |Endpoint|. However, there are cases where you -// need more complex calling patterns. For example, with |SequenceCaller| -// you can invoke the same API |Descriptor| with multiple equivalent -// API |Endpoint|s until one of them succeeds or all fail. -package httpapi diff --git a/internal/httpapi/endpoint.go b/internal/httpapi/endpoint.go deleted file mode 100644 index acfc4d8..0000000 --- a/internal/httpapi/endpoint.go +++ /dev/null @@ -1,76 +0,0 @@ -package httpapi - -// -// HTTP API Endpoint (e.g., https://api.ooni.io) -// - -import "github.com/ooni/probe-cli/v3/internal/model" - -// Endpoint models an HTTP endpoint on which you can call -// several HTTP APIs (e.g., https://api.ooni.io) using a -// given HTTP client potentially using a circumvention tunnel -// mechanism such as psiphon or torsf. -// -// The zero value of this struct is invalid. Please, fill all the -// fields marked as MANDATORY for correct initialization. -type Endpoint struct { - // BaseURL is the MANDATORY endpoint base URL. We will honour the - // path of this URL and prepend it to the actual path specified inside - // a |Descriptor.URLPath|. However, we will always discard any query - // that may have been set inside the BaseURL. The only query string - // will be composed from the |Descriptor.URLQuery| values. - // - // For example, https://api.ooni.io. - BaseURL string - - // HTTPClient is the MANDATORY HTTP client to use. - // - // For example, http.DefaultClient. You can introduce circumvention - // here by using an HTTPClient bound to a specific tunnel. - HTTPClient model.HTTPClient - - // Host is the OPTIONAL host header to use. - // - // If this field is empty we use the BaseURL's hostname. A specific - // host header may be needed when using cloudfronting. - Host string - - // User-Agent is the OPTIONAL user-agent to use. If empty, - // we'll use the stdlib's default user-agent string. - UserAgent string -} - -// NewEndpointList constructs a list of API endpoints from |services| -// returned by the OONI backend (or known in advance). -// -// Arguments: -// -// - httpClient is the HTTP client to use for accessing the endpoints; -// -// - userAgent is the user agent you would like to use; -// -// - service is the list of services gathered from the backend. -func NewEndpointList(httpClient model.HTTPClient, - userAgent string, services ...model.OOAPIService) (out []*Endpoint) { - for _, svc := range services { - switch svc.Type { - case "https": - out = append(out, &Endpoint{ - BaseURL: svc.Address, - HTTPClient: httpClient, - Host: "", - UserAgent: userAgent, - }) - case "cloudfront": - out = append(out, &Endpoint{ - BaseURL: svc.Address, - HTTPClient: httpClient, - Host: svc.Front, - UserAgent: userAgent, - }) - default: - // nothing! - } - } - return -} diff --git a/internal/httpapi/endpoint_test.go b/internal/httpapi/endpoint_test.go deleted file mode 100644 index 7077e14..0000000 --- a/internal/httpapi/endpoint_test.go +++ /dev/null @@ -1,69 +0,0 @@ -package httpapi - -import ( - "testing" - - "github.com/google/go-cmp/cmp" - "github.com/ooni/probe-cli/v3/internal/model" - "github.com/ooni/probe-cli/v3/internal/model/mocks" -) - -func TestNewEndpointList(t *testing.T) { - type args struct { - httpClient model.HTTPClient - userAgent string - services []model.OOAPIService - } - defaultHTTPClient := &mocks.HTTPClient{} - tests := []struct { - name string - args args - wantOut []*Endpoint - }{{ - name: "with no services", - args: args{ - httpClient: defaultHTTPClient, - userAgent: model.HTTPHeaderUserAgent, - services: nil, - }, - wantOut: nil, - }, { - name: "common cases", - args: args{ - httpClient: defaultHTTPClient, - userAgent: model.HTTPHeaderUserAgent, - services: []model.OOAPIService{{ - Address: "https://www.example.com/", - Type: "https", - Front: "", - }, { - Address: "https://www.example.org/", - Type: "cloudfront", - Front: "example.org.it", - }, { - Address: "https://nonexistent.onion/", - Type: "onion", - Front: "", - }}, - }, - wantOut: []*Endpoint{{ - BaseURL: "https://www.example.com/", - HTTPClient: defaultHTTPClient, - Host: "", - UserAgent: model.HTTPHeaderUserAgent, - }, { - BaseURL: "https://www.example.org/", - HTTPClient: defaultHTTPClient, - Host: "example.org.it", - UserAgent: model.HTTPHeaderUserAgent, - }}, - }} - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - gotOut := NewEndpointList(tt.args.httpClient, tt.args.userAgent, tt.args.services...) - if diff := cmp.Diff(tt.wantOut, gotOut); diff != "" { - t.Fatal(diff) - } - }) - } -} diff --git a/internal/httpapi/sequence.go b/internal/httpapi/sequence.go deleted file mode 100644 index da1f11d..0000000 --- a/internal/httpapi/sequence.go +++ /dev/null @@ -1,92 +0,0 @@ -package httpapi - -// -// Sequentially call available API endpoints until one succeed -// or all of them fail. A future implementation of this code may -// (probably should?) take into account knowledge of what is -// working and what is not working to optimize the order with -// which to try different alternatives. -// - -import ( - "context" - "errors" - - "github.com/ooni/probe-cli/v3/internal/multierror" -) - -// SequenceCaller calls the API specified by |Descriptor| once for each of -// the available |Endpoints| until one of them succeeds. -// -// CAVEAT: this code will ONLY retry API calls with subsequent endpoints when -// the error originates in the HTTP round trip or while reading the body. -type SequenceCaller struct { - // Descriptor is the API |Descriptor|. - Descriptor *Descriptor - - // Endpoints is the list of |Endpoint| to use. - Endpoints []*Endpoint -} - -// NewSequenceCaller is a factory for creating a |SequenceCaller|. -func NewSequenceCaller(desc *Descriptor, endpoints ...*Endpoint) *SequenceCaller { - return &SequenceCaller{ - Descriptor: desc, - Endpoints: endpoints, - } -} - -// ErrAllEndpointsFailed indicates that all endpoints failed. -var ErrAllEndpointsFailed = errors.New("httpapi: all endpoints failed") - -// shouldRetry returns true when we should try with another endpoint given the -// value of |err| which could (obviously) be nil in case of success. -func (sc *SequenceCaller) shouldRetry(err error) bool { - var kind *errMaybeCensorship - belongs := errors.As(err, &kind) - return belongs -} - -// Call calls |Call| for each |Endpoint| and |Descriptor| until one endpoint succeeds. The -// return value is the response body and the selected endpoint index or the error. -// -// CAVEAT: this code will ONLY retry API calls with subsequent endpoints when -// the error originates in the HTTP round trip or while reading the body. -func (sc *SequenceCaller) Call(ctx context.Context) ([]byte, int, error) { - var selected int - merr := multierror.New(ErrAllEndpointsFailed) - for _, epnt := range sc.Endpoints { - respBody, err := Call(ctx, sc.Descriptor, epnt) - if sc.shouldRetry(err) { - merr.Add(err) - selected++ - continue - } - // Note: some errors will lead us to return - // early as documented for this method - return respBody, selected, err - } - return nil, -1, merr -} - -// CallWithJSONResponse is like |SequenceCaller.Call| except that it invokes the -// underlying |CallWithJSONResponse| rather than invoking |Call|. -// -// CAVEAT: this code will ONLY retry API calls with subsequent endpoints when -// the error originates in the HTTP round trip or while reading the body. -func (sc *SequenceCaller) CallWithJSONResponse(ctx context.Context, response any) (int, error) { - var selected int - merr := multierror.New(ErrAllEndpointsFailed) - for _, epnt := range sc.Endpoints { - err := CallWithJSONResponse(ctx, sc.Descriptor, epnt, response) - if sc.shouldRetry(err) { - merr.Add(err) - selected++ - continue - } - // Note: some errors will lead us to return - // early as documented for this method - return selected, err - } - return -1, merr -} diff --git a/internal/httpapi/sequence_test.go b/internal/httpapi/sequence_test.go deleted file mode 100644 index 13cc50f..0000000 --- a/internal/httpapi/sequence_test.go +++ /dev/null @@ -1,358 +0,0 @@ -package httpapi - -import ( - "context" - "errors" - "io" - "net/http" - "strings" - "testing" - - "github.com/google/go-cmp/cmp" - "github.com/ooni/probe-cli/v3/internal/model" - "github.com/ooni/probe-cli/v3/internal/model/mocks" -) - -func TestSequenceCaller(t *testing.T) { - t.Run("Call", func(t *testing.T) { - t.Run("first success", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - StatusCode: 200, - Body: io.NopCloser(strings.NewReader("deadbeef")), - } - return resp, nil - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF - }, - }, - }, - ) - data, idx, err := sc.Call(context.Background()) - if err != nil { - t.Fatal(err) - } - if idx != 0 { - t.Fatal("invalid idx") - } - if diff := cmp.Diff([]byte("deadbeef"), data); diff != "" { - t.Fatal(diff) - } - }) - - t.Run("first HTTP failure and we immediately stop", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - StatusCode: 403, // should cause us to return early - Body: io.NopCloser(strings.NewReader("deadbeef")), - } - return resp, nil - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF - }, - }, - }, - ) - data, idx, err := sc.Call(context.Background()) - var failure *ErrHTTPRequestFailed - if !errors.As(err, &failure) || failure.StatusCode != 403 { - t.Fatal("unexpected err", err) - } - if idx != 0 { - t.Fatal("invalid idx") - } - if len(data) > 0 { - t.Fatal("expected to see no response body") - } - }) - - t.Run("first network failure, second success", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF // should cause us to cycle to the second entry - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - StatusCode: 200, - Body: io.NopCloser(strings.NewReader("abad1dea")), - } - return resp, nil - }, - }, - }, - ) - data, idx, err := sc.Call(context.Background()) - if err != nil { - t.Fatal(err) - } - if idx != 1 { - t.Fatal("invalid idx") - } - if diff := cmp.Diff([]byte("abad1dea"), data); diff != "" { - t.Fatal(diff) - } - }) - - t.Run("all network failure", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF // should cause us to cycle to the next entry - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF // should cause us to cycle to the next entry - }, - }, - }, - ) - data, idx, err := sc.Call(context.Background()) - if !errors.Is(err, ErrAllEndpointsFailed) { - t.Fatal("unexpected err", err) - } - if idx != -1 { - t.Fatal("invalid idx") - } - if len(data) > 0 { - t.Fatal("expected zero-length data") - } - }) - }) - - t.Run("CallWithJSONResponse", func(t *testing.T) { - type response struct { - Name string - Age int64 - } - - t.Run("first success", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - StatusCode: 200, - Body: io.NopCloser(strings.NewReader(`{"Name":"sbs","Age":99}`)), - } - return resp, nil - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - StatusCode: 200, - Body: io.NopCloser(strings.NewReader(`{}`)), // different - } - return resp, nil - }, - }, - }, - ) - expect := response{ - Name: "sbs", - Age: 99, - } - var got response - idx, err := sc.CallWithJSONResponse(context.Background(), &got) - if err != nil { - t.Fatal(err) - } - if idx != 0 { - t.Fatal("invalid idx") - } - if diff := cmp.Diff(expect, got); diff != "" { - t.Fatal(diff) - } - }) - - t.Run("first HTTP failure and we immediately stop", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - StatusCode: 403, // should be enough to cause us fail immediately - Body: io.NopCloser(strings.NewReader(`{"Age": 155, "Name": "sbs"}`)), - } - return resp, nil - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF - }, - }, - }, - ) - // even though there is a JSON body we don't care about reading it - // and so we expect to see in output the zero-value struct - expect := response{ - Name: "", - Age: 0, - } - var got response - idx, err := sc.CallWithJSONResponse(context.Background(), &got) - var failure *ErrHTTPRequestFailed - if !errors.As(err, &failure) || failure.StatusCode != 403 { - t.Fatal("unexpected err", err) - } - if idx != 0 { - t.Fatal("invalid idx") - } - if diff := cmp.Diff(expect, got); diff != "" { - t.Fatal(diff) - } - }) - - t.Run("first network failure, second success", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF // should cause us to try the next entry - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - resp := &http.Response{ - StatusCode: 200, - Body: io.NopCloser(strings.NewReader(`{"Age":155}`)), - } - return resp, nil - }, - }, - }, - ) - expect := response{ - Name: "", - Age: 155, - } - var got response - idx, err := sc.CallWithJSONResponse(context.Background(), &got) - if err != nil { - t.Fatal(err) - } - if idx != 1 { - t.Fatal("invalid idx") - } - if diff := cmp.Diff(expect, got); diff != "" { - t.Fatal(diff) - } - }) - - t.Run("all network failure", func(t *testing.T) { - sc := NewSequenceCaller( - &Descriptor{ - Logger: model.DiscardLogger, - Method: http.MethodGet, - URLPath: "/", - }, - &Endpoint{ - BaseURL: "https://a.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF // should cause us to try the next entry - }, - }, - }, - &Endpoint{ - BaseURL: "https://b.example.com/", - HTTPClient: &mocks.HTTPClient{ - MockDo: func(req *http.Request) (*http.Response, error) { - return nil, io.EOF // should cause us to try the next entry - }, - }, - }, - ) - var got response - idx, err := sc.CallWithJSONResponse(context.Background(), &got) - if !errors.Is(err, ErrAllEndpointsFailed) { - t.Fatal("unexpected err", err) - } - if idx != -1 { - t.Fatal("invalid idx") - } - }) - }) -} diff --git a/internal/httpx/httpx.go b/internal/httpx/httpx.go index 8083f10..f330745 100644 --- a/internal/httpx/httpx.go +++ b/internal/httpx/httpx.go @@ -1,8 +1,4 @@ // Package httpx contains http extensions. -// -// Deprecated: new code should use httpapi instead. While this package and httpapi -// are basically using the same implementation, the API exposed by httpapi allows -// us to try the same request with multiple HTTP endpoints. package httpx import ( diff --git a/internal/model/experiment.go b/internal/model/experiment.go index 1cba6c1..acaadde 100644 --- a/internal/model/experiment.go +++ b/internal/model/experiment.go @@ -117,19 +117,6 @@ func (d PrinterCallbacks) OnProgress(percentage float64, message string) { d.Logger.Infof("[%5.1f%%] %s", percentage*100, message) } -// ExperimentArgs contains the arguments passed to an experiment. -type ExperimentArgs struct { - // Callbacks contains MANDATORY experiment callbacks. - Callbacks ExperimentCallbacks - - // Measurement is the MANDATORY measurement in which the experiment - // must write the results of the measurement. - Measurement *Measurement - - // Session is the MANDATORY session the experiment can use. - Session ExperimentSession -} - // ExperimentMeasurer is the interface that allows to run a // measurement for a specific experiment. type ExperimentMeasurer interface { @@ -146,7 +133,10 @@ type ExperimentMeasurer interface { // set the relevant OONI error inside of the measurement and // return nil. This is important because the caller WILL NOT submit // the measurement if this method returns an error. - Run(ctx context.Context, args *ExperimentArgs) error + Run( + ctx context.Context, sess ExperimentSession, + measurement *Measurement, callbacks ExperimentCallbacks, + ) error // GetSummaryKeys returns summary keys expected by ooni/probe-cli. GetSummaryKeys(*Measurement) (interface{}, error) diff --git a/internal/registry/smtp.go b/internal/registry/smtp.go new file mode 100644 index 0000000..cb67e45 --- /dev/null +++ b/internal/registry/smtp.go @@ -0,0 +1,22 @@ +package registry + +// +// Registers the `dnsping' experiment. +// + +import ( + "github.com/ooni/probe-cli/v3/internal/engine/experiment/smtp" + "github.com/ooni/probe-cli/v3/internal/model" +) + +func init() { + AllExperiments["smtp"] = &Factory{ + build: func(config interface{}) model.ExperimentMeasurer { + return smtp.NewExperimentMeasurer( + *config.(*smtp.Config), + ) + }, + config: &smtp.Config{}, + inputPolicy: model.InputOrStaticDefault, + } +} diff --git a/internal/tutorial/experiment/torsf/chapter01/README.md b/internal/tutorial/experiment/torsf/chapter01/README.md index eebb606..d357321 100644 --- a/internal/tutorial/experiment/torsf/chapter01/README.md +++ b/internal/tutorial/experiment/torsf/chapter01/README.md @@ -211,12 +211,7 @@ need any fancy context and we pass a `context.Background` to `Run`. ```Go ctx := context.Background() - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err = m.Run(ctx, args); err != nil { + if err = m.Run(ctx, sess, measurement, callbacks); err != nil { log.WithError(err).Fatal("torsf experiment failed") } ``` diff --git a/internal/tutorial/experiment/torsf/chapter01/main.go b/internal/tutorial/experiment/torsf/chapter01/main.go index d96c5f3..aef3716 100644 --- a/internal/tutorial/experiment/torsf/chapter01/main.go +++ b/internal/tutorial/experiment/torsf/chapter01/main.go @@ -212,12 +212,7 @@ func main() { // // ```Go ctx := context.Background() - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err = m.Run(ctx, args); err != nil { + if err = m.Run(ctx, sess, measurement, callbacks); err != nil { log.WithError(err).Fatal("torsf experiment failed") } // ``` diff --git a/internal/tutorial/experiment/torsf/chapter02/README.md b/internal/tutorial/experiment/torsf/chapter02/README.md index 36f3323..c9a2fbd 100644 --- a/internal/tutorial/experiment/torsf/chapter02/README.md +++ b/internal/tutorial/experiment/torsf/chapter02/README.md @@ -117,10 +117,10 @@ chapters, finally, we will modify this function until it is a minimal implementation of the `torsf` experiment. ```Go -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - _ = args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { ``` As you can see, this is just a stub implementation that sleeps for one second and prints a logging message. diff --git a/internal/tutorial/experiment/torsf/chapter02/main.go b/internal/tutorial/experiment/torsf/chapter02/main.go index efded79..4464bdc 100644 --- a/internal/tutorial/experiment/torsf/chapter02/main.go +++ b/internal/tutorial/experiment/torsf/chapter02/main.go @@ -54,12 +54,7 @@ func main() { MockableLogger: log.Log, MockableTempDir: tempdir, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err = m.Run(ctx, args); err != nil { + if err = m.Run(ctx, sess, measurement, callbacks); err != nil { log.WithError(err).Fatal("torsf experiment failed") } data, err := json.Marshal(measurement) diff --git a/internal/tutorial/experiment/torsf/chapter02/torsf.go b/internal/tutorial/experiment/torsf/chapter02/torsf.go index 2547cb4..af3e0f6 100644 --- a/internal/tutorial/experiment/torsf/chapter02/torsf.go +++ b/internal/tutorial/experiment/torsf/chapter02/torsf.go @@ -93,10 +93,10 @@ func (m *Measurer) ExperimentVersion() string { // minimal implementation of the `torsf` experiment. // // ```Go -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - _ = args.Callbacks - _ = args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { // ``` // As you can see, this is just a stub implementation that sleeps // for one second and prints a logging message. diff --git a/internal/tutorial/experiment/torsf/chapter03/README.md b/internal/tutorial/experiment/torsf/chapter03/README.md index 3ae6719..b6b9f85 100644 --- a/internal/tutorial/experiment/torsf/chapter03/README.md +++ b/internal/tutorial/experiment/torsf/chapter03/README.md @@ -32,10 +32,10 @@ print periodic updates via the `callbacks`. We will defer the real work to a private function called `run`. ```Go -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { ``` Let's create an instance of `TestKeys` and let's modify diff --git a/internal/tutorial/experiment/torsf/chapter03/main.go b/internal/tutorial/experiment/torsf/chapter03/main.go index 1e9c8ec..a8dad28 100644 --- a/internal/tutorial/experiment/torsf/chapter03/main.go +++ b/internal/tutorial/experiment/torsf/chapter03/main.go @@ -28,12 +28,7 @@ func main() { MockableLogger: log.Log, MockableTempDir: tempdir, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err = m.Run(ctx, args); err != nil { + if err = m.Run(ctx, sess, measurement, callbacks); err != nil { log.WithError(err).Fatal("torsf experiment failed") } data, err := json.Marshal(measurement) diff --git a/internal/tutorial/experiment/torsf/chapter03/torsf.go b/internal/tutorial/experiment/torsf/chapter03/torsf.go index e48ca3a..97b7ae4 100644 --- a/internal/tutorial/experiment/torsf/chapter03/torsf.go +++ b/internal/tutorial/experiment/torsf/chapter03/torsf.go @@ -65,10 +65,10 @@ type TestKeys struct { // real work to a private function called `run`. // // ```Go -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { // ``` // // Let's create an instance of `TestKeys` and let's modify diff --git a/internal/tutorial/experiment/torsf/chapter04/main.go b/internal/tutorial/experiment/torsf/chapter04/main.go index 1e9c8ec..a8dad28 100644 --- a/internal/tutorial/experiment/torsf/chapter04/main.go +++ b/internal/tutorial/experiment/torsf/chapter04/main.go @@ -28,12 +28,7 @@ func main() { MockableLogger: log.Log, MockableTempDir: tempdir, } - args := &model.ExperimentArgs{ - Callbacks: callbacks, - Measurement: measurement, - Session: sess, - } - if err = m.Run(ctx, args); err != nil { + if err = m.Run(ctx, sess, measurement, callbacks); err != nil { log.WithError(err).Fatal("torsf experiment failed") } data, err := json.Marshal(measurement) diff --git a/internal/tutorial/experiment/torsf/chapter04/torsf.go b/internal/tutorial/experiment/torsf/chapter04/torsf.go index d350ee9..d6998b9 100644 --- a/internal/tutorial/experiment/torsf/chapter04/torsf.go +++ b/internal/tutorial/experiment/torsf/chapter04/torsf.go @@ -99,10 +99,10 @@ type TestKeys struct { } // Run implements ExperimentMeasurer.Run. -func (m *Measurer) Run(ctx context.Context, args *model.ExperimentArgs) error { - callbacks := args.Callbacks - measurement := args.Measurement - sess := args.Session +func (m *Measurer) Run( + ctx context.Context, sess model.ExperimentSession, + measurement *model.Measurement, callbacks model.ExperimentCallbacks, +) error { testkeys := &TestKeys{} measurement.TestKeys = testkeys start := time.Now()