Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
68 changes: 59 additions & 9 deletions frametests/driver.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,11 @@ package frametests

import (
"context"
"fmt"
"net"
"net/http"
"net/http/httptest"
"strings"
"sync"

"github.com/pitabwire/util"
Expand All @@ -29,22 +31,59 @@ func GetFreePort(ctx context.Context) (int, error) {
}

type testDriver struct {
mu sync.RWMutex
srv *httptest.Server
mu sync.RWMutex
bindAddr bool
srv *httptest.Server
}

func (t *testDriver) ListenAndServe(_ string, h http.Handler) error {
t.mu.Lock()
defer t.mu.Unlock()
t.srv = httptest.NewServer(h)
func (t *testDriver) ListenAndServe(addr string, h http.Handler) error {
srv, err := t.newServer(addr, h)
if err != nil {
return err
}
srv.Start()
t.setServer(srv)
return nil
}

func (t *testDriver) ListenAndServeTLS(addr, _, _ string, h http.Handler) error {
srv, err := t.newServer(addr, h)
if err != nil {
return err
}
srv.StartTLS()
t.setServer(srv)
return nil
}
func (t *testDriver) ListenAndServeTLS(_, _, _ string, h http.Handler) error {

// newServer prepares an unstarted httptest server. Drivers created with
// WithBoundHTTPTestDriver bind the address frame derived from configuration;
// otherwise httptest picks a free loopback port. Binding happens here, on the
// caller's goroutine, so listen failures are returned from Service.Run.
func (t *testDriver) newServer(addr string, h http.Handler) (*httptest.Server, error) {
srv := httptest.NewUnstartedServer(h)
if !t.bindAddr {
return srv, nil
}

if !strings.Contains(addr, ":") {
addr = ":" + addr
}
// Driver construction has no request context; the listener lives for the service.
listener, err := (&net.ListenConfig{}).Listen(context.Background(), "tcp", addr)
if err != nil {
_ = srv.Listener.Close()
return nil, fmt.Errorf("test driver listen on %s: %w", addr, err)
}
_ = srv.Listener.Close()
srv.Listener = listener
return srv, nil
}

func (t *testDriver) setServer(srv *httptest.Server) {
t.mu.Lock()
defer t.mu.Unlock()
t.srv = httptest.NewTLSServer(h)
return nil
t.srv = srv
}

func (t *testDriver) Shutdown(_ context.Context) error {
Expand All @@ -64,11 +103,22 @@ func (t *testDriver) GetTestServer() *httptest.Server {
}

// WithHTTPTestDriver uses a driver, mostly useful when writing tests against the frame service.
// The server listens on a random loopback port; read it from the returned accessor.
func WithHTTPTestDriver() (frame.Option, func() *httptest.Server) {
driver := &testDriver{}
return frame.WithDriver(driver), driver.GetTestServer
}

// WithBoundHTTPTestDriver is WithHTTPTestDriver bound to the service's configured
// HTTP port (HTTPServerPort) instead of a random one. Use it when peers such as an
// OAuth2 server or webhook caller are configured with the service URL before the
// service starts. Service.Run returns once the listener is bound, so startup and
// listen errors surface synchronously — no goroutine or readiness polling needed.
func WithBoundHTTPTestDriver() (frame.Option, func() *httptest.Server) {
driver := &testDriver{bindAddr: true}
return frame.WithDriver(driver), driver.GetTestServer
}

type noopDriver struct {
}

Expand Down
79 changes: 79 additions & 0 deletions frametests/driver_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
package frametests_test

import (
"fmt"
"net"
"net/http"
"strconv"
"testing"

"github.com/pitabwire/frame/v2"
"github.com/pitabwire/frame/v2/config"
"github.com/pitabwire/frame/v2/frametests"
"github.com/stretchr/testify/require"
)

func okHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusAccepted)
})
}

func TestBoundHTTPTestDriver_ServesOnConfiguredPort(t *testing.T) {
port, err := frametests.GetFreePort(t.Context())
require.NoError(t, err)

driverOpt, testServer := frametests.WithBoundHTTPTestDriver()
cfg := config.ConfigurationDefault{HTTPServerPort: strconv.Itoa(port)}
ctx, svc := frame.NewServiceWithContext(t.Context(),
frame.WithName("bound-driver"),
frame.WithConfig(&cfg),
driverOpt,
frame.WithHTTPHandler(okHandler()),
)
defer svc.Stop(ctx)

require.NoError(t, svc.Run(ctx, ""), "Run returns once the listener is bound")
require.NotNil(t, testServer())

req, err := http.NewRequestWithContext(ctx, http.MethodGet, fmt.Sprintf("http://127.0.0.1:%d/", port), nil)
require.NoError(t, err)
resp, err := testServer().Client().Do(req)
require.NoError(t, err, "service must be reachable on the pre-allocated port")
_ = resp.Body.Close()
require.Equal(t, http.StatusAccepted, resp.StatusCode)
}

func TestBoundHTTPTestDriver_ReturnsListenErrorFromRun(t *testing.T) {
// Occupy a port so binding it must fail.
occupied, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer func() { _ = occupied.Close() }()
addr := occupied.Addr().String()

driverOpt, _ := frametests.WithBoundHTTPTestDriver()
ctx, svc := frame.NewServiceWithContext(t.Context(),
frame.WithName("bound-driver-conflict"),
driverOpt,
frame.WithHTTPHandler(okHandler()),
)
defer svc.Stop(ctx)

err = svc.Run(ctx, addr)
require.Error(t, err, "a listen failure must surface from Run, not be lost")
require.ErrorContains(t, err, addr)
}

func TestHTTPTestDriver_StillPicksRandomPort(t *testing.T) {
driverOpt, testServer := frametests.WithHTTPTestDriver()
ctx, svc := frame.NewServiceWithContext(t.Context(),
frame.WithName("random-driver"),
driverOpt,
frame.WithHTTPHandler(okHandler()),
)
defer svc.Stop(ctx)

require.NoError(t, svc.Run(ctx, ":41577"))
require.NotNil(t, testServer())
require.NotContains(t, testServer().URL, ":41577", "random-port driver must ignore the configured address")
}
Loading