diff --git a/async/async.go b/async/async.go index 36ee997..8ccef58 100644 --- a/async/async.go +++ b/async/async.go @@ -16,6 +16,52 @@ import ( // from the originating request so it survives handler return. type Task func(ctx context.Context) error +// Runner bounds fire-and-forget work so degraded downstreams cannot create +// unbounded goroutines under load. +type Runner struct { + sem chan struct{} + logger *slog.Logger +} + +// NewRunner creates a bounded background task runner. +func NewRunner(maxInFlight int, logger *slog.Logger) *Runner { + if maxInFlight <= 0 { + maxInFlight = 1 + } + if logger == nil { + logger = slog.Default() + } + return &Runner{sem: make(chan struct{}, maxInFlight), logger: logger} +} + +// TryGo launches task with a context detached from request cancellation when +// capacity is available. It returns false without launching a goroutine when +// the runner is full. +func (r *Runner) TryGo(ctx context.Context, op string, task Task) bool { + if r == nil || task == nil { + return false + } + select { + case r.sem <- struct{}{}: + default: + r.logger.Warn("background task rejected: capacity exhausted", "op", op) + return false + } + bgCtx := context.WithoutCancel(ctx) + go func() { + defer func() { + if recovered := recover(); recovered != nil { + r.logger.Error("background task panicked", "op", op, "panic", recovered, "stack", string(debug.Stack())) + } + <-r.sem + }() + if err := task(bgCtx); err != nil { + r.logger.Warn("background task failed", "op", op, "error", err) + } + }() + return true +} + // FireAndForget launches task in a background goroutine with a context // detached from the parent (via context.WithoutCancel). Errors are logged // at WARN level and panics are recovered and logged at ERROR level with the diff --git a/async/async_test.go b/async/async_test.go index 701d471..77fd979 100644 --- a/async/async_test.go +++ b/async/async_test.go @@ -80,6 +80,33 @@ func TestFireAndForgetRecoversPanics(t *testing.T) { } } +func TestRunnerRejectsWhenFull(t *testing.T) { + runner := async.NewRunner(1, slog.Default()) + started := make(chan struct{}) + release := make(chan struct{}) + + if !runner.TryGo(context.Background(), "test.block", func(_ context.Context) error { + close(started) + <-release + return nil + }) { + t.Fatal("expected first task to be accepted") + } + <-started + + var ran atomic.Bool + if runner.TryGo(context.Background(), "test.rejected", func(_ context.Context) error { + ran.Store(true) + return nil + }) { + t.Fatal("expected second task to be rejected while runner is full") + } + if ran.Load() { + t.Fatal("rejected task should not run") + } + close(release) +} + func TestCloneProtoPreventsMutationRace(t *testing.T) { original := &approvalsv1.ApprovalRequest{ Id: "req-1", diff --git a/authmw/authmw.go b/authmw/authmw.go index 1803439..8594c0f 100644 --- a/authmw/authmw.go +++ b/authmw/authmw.go @@ -3,6 +3,7 @@ package authmw import ( "context" "net/http" + "slices" "strings" "github.com/evalops/service-runtime/httpkit" @@ -76,7 +77,7 @@ func (middleware *Middleware) WithAuth(scopes ...string) func(http.Handler) http return } - actor, err := middleware.authenticate(request.Context(), token, scopes) + actor, availableScopes, err := middleware.authenticate(request.Context(), token, scopes) if err != nil { status := http.StatusUnauthorized if middleware.isForbidden(err) { @@ -87,6 +88,7 @@ func (middleware *Middleware) WithAuth(scopes ...string) func(http.Handler) http } ctx := context.WithValue(request.Context(), actorContextKey, actor) + ctx = ContextWithPrincipal(ctx, PrincipalFromActor(actor, availableScopes)) next.ServeHTTP(writer, request.WithContext(ctx)) }) } @@ -101,20 +103,18 @@ func ActorFromContext(ctx context.Context) (Actor, bool) { // HasAllScopes reports whether all required scopes are present in available. func HasAllScopes(available []string, required []string) bool { for _, requirement := range required { - found := false - for _, scope := range available { - if scope == requirement { - found = true - break - } - } - if !found { + if !HasScope(available, requirement) { return false } } return true } +// HasScope reports whether the required scope is present in available scopes. +func HasScope(available []string, required string) bool { + return required != "" && slices.Contains(available, required) +} + // BearerToken extracts the token from an Authorization: Bearer header. func BearerToken(header string) (string, bool) { parts := strings.SplitN(header, " ", 2) @@ -125,41 +125,41 @@ func BearerToken(header string) (string, bool) { return token, token != "" } -func (middleware *Middleware) authenticate(ctx context.Context, token string, scopes []string) (Actor, error) { +func (middleware *Middleware) authenticate(ctx context.Context, token string, scopes []string) (Actor, []string, error) { if strings.HasPrefix(token, "pk_") { return middleware.authenticateAPIKey(ctx, token, scopes) } return middleware.authenticateToken(ctx, token, scopes) } -func (middleware *Middleware) authenticateAPIKey(ctx context.Context, token string, scopes []string) (Actor, error) { +func (middleware *Middleware) authenticateAPIKey(ctx context.Context, token string, scopes []string) (Actor, []string, error) { if middleware == nil || middleware.apiKeyValidator == nil { - return Actor{}, errAPIKeyValidatorUnavailable + return Actor{}, nil, errAPIKeyValidatorUnavailable } key, err := middleware.apiKeyValidator.ValidateAPIKey(ctx, token) if err != nil { - return Actor{}, err + return Actor{}, nil, err } if !HasAllScopes(key.Scopes, scopes) { - return Actor{}, errMissingScopes + return Actor{}, nil, errMissingScopes } return Actor{ Type: "api_key", ID: key.ID, OrganizationID: key.OrganizationID, - }, nil + }, append([]string(nil), key.Scopes...), nil } -func (middleware *Middleware) authenticateToken(ctx context.Context, token string, scopes []string) (Actor, error) { +func (middleware *Middleware) authenticateToken(ctx context.Context, token string, scopes []string) (Actor, []string, error) { if middleware == nil || middleware.tokenVerifier == nil { - return Actor{}, errTokenVerifierUnavailable + return Actor{}, nil, errTokenVerifierUnavailable } verified, err := middleware.tokenVerifier.VerifyToken(ctx, token, scopes) if err != nil { - return Actor{}, err + return Actor{}, nil, err } - return verified.Actor, nil + return verified.Actor, append([]string(nil), verified.Scopes...), nil } func (middleware *Middleware) isForbidden(err error) bool { diff --git a/authmw/authmw_test.go b/authmw/authmw_test.go index 645db0f..beed8c8 100644 --- a/authmw/authmw_test.go +++ b/authmw/authmw_test.go @@ -40,12 +40,17 @@ func TestWithAuthAPIKey(t *testing.T) { middleware := New(Config{APIKeyValidator: validator}) var seenActor Actor + var seenPrincipal Principal handler := middleware.WithAuth("scope:read")(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { var ok bool seenActor, ok = ActorFromContext(request.Context()) if !ok { t.Fatal("expected actor in context") } + seenPrincipal, ok = PrincipalFromContext(request.Context()) + if !ok { + t.Fatal("expected principal in context") + } writer.WriteHeader(http.StatusNoContent) })) @@ -60,6 +65,12 @@ func TestWithAuthAPIKey(t *testing.T) { if seenActor.Type != "api_key" || seenActor.OrganizationID != "org-123" { t.Fatalf("unexpected actor %#v", seenActor) } + if seenPrincipal.OrganizationID != "org-123" || seenPrincipal.Subject != "integration-key" { + t.Fatalf("unexpected principal %#v", seenPrincipal) + } + if err := seenPrincipal.RequireScope("scope:read"); err != nil { + t.Fatalf("expected scope to be present: %v", err) + } } func TestWithAuthAPIKeyMissingScopes(t *testing.T) { @@ -118,17 +129,23 @@ func TestWithAuthServiceToken(t *testing.T) { ID: "pipeline", OrganizationID: "org-123", }, + Scopes: []string{"scope:write"}, }, }, }) var seenActor Actor + var seenPrincipal Principal handler := middleware.WithAuth("scope:write")(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { var ok bool seenActor, ok = ActorFromContext(request.Context()) if !ok { t.Fatal("expected actor in context") } + seenPrincipal, ok = PrincipalFromContext(request.Context()) + if !ok { + t.Fatal("expected principal in context") + } writer.WriteHeader(http.StatusAccepted) })) @@ -143,6 +160,12 @@ func TestWithAuthServiceToken(t *testing.T) { if seenActor.Type != "service" || seenActor.ID != "pipeline" { t.Fatalf("unexpected actor %#v", seenActor) } + if seenPrincipal.Service != "pipeline" || seenPrincipal.OrganizationID != "org-123" { + t.Fatalf("unexpected principal %#v", seenPrincipal) + } + if err := seenPrincipal.RequireScope("scope:write"); err != nil { + t.Fatalf("expected scope to be present: %v", err) + } } func TestWithAuthServiceTokenWithoutVerifier(t *testing.T) { diff --git a/authmw/principal.go b/authmw/principal.go new file mode 100644 index 0000000..02f7c1a --- /dev/null +++ b/authmw/principal.go @@ -0,0 +1,142 @@ +package authmw + +import ( + "context" + "errors" + "fmt" + "slices" + "strings" + + "connectrpc.com/connect" +) + +const principalContextKey contextKey = "principal" + +// Principal is the normalized authenticated caller context that service +// handlers should authorize against after transport authentication succeeds. +type Principal struct { + OrganizationID string `json:"organization_id"` + WorkspaceID string `json:"workspace_id,omitempty"` + Subject string `json:"subject,omitempty"` + UserSubject string `json:"user_subject,omitempty"` + Service string `json:"service,omitempty"` + TokenType string `json:"token_type,omitempty"` + Scopes []string `json:"scopes,omitempty"` + AgentID string `json:"agent_id,omitempty"` + IsHuman bool `json:"is_human,omitempty"` +} + +// PrincipalFromActor converts the lower-level authentication actor into the +// authorization principal shape used by service handlers. +func PrincipalFromActor(actor Actor, scopes []string) Principal { + principal := Principal{ + OrganizationID: strings.TrimSpace(actor.OrganizationID), + Subject: strings.TrimSpace(actor.ID), + TokenType: strings.TrimSpace(actor.Type), + Scopes: append([]string(nil), scopes...), + IsHuman: isHumanActorType(actor.Type), + } + if principal.TokenType == "service" { + principal.Service = principal.Subject + } + for key, value := range actor.Attributes { + trimmed := strings.TrimSpace(value) + switch strings.ToLower(strings.TrimSpace(key)) { + case "workspace_id", "workspace": + principal.WorkspaceID = trimmed + case "agent_id", "agent": + principal.AgentID = trimmed + case "user_subject": + principal.UserSubject = trimmed + case "service": + principal.Service = trimmed + case "token_type": + principal.TokenType = trimmed + case "is_human": + principal.IsHuman = principal.IsHuman || strings.EqualFold(trimmed, "true") + } + } + return principal +} + +// ContextWithPrincipal stores an authenticated principal in context. +func ContextWithPrincipal(ctx context.Context, principal Principal) context.Context { + principal.Scopes = append([]string(nil), principal.Scopes...) + return context.WithValue(ctx, principalContextKey, principal) +} + +// PrincipalFromContext retrieves the authenticated principal from context. +func PrincipalFromContext(ctx context.Context) (Principal, bool) { + principal, ok := ctx.Value(principalContextKey).(Principal) + if !ok { + return Principal{}, false + } + principal.Scopes = append([]string(nil), principal.Scopes...) + return principal, true +} + +// RequireOrganization returns PermissionDenied when the caller is not scoped to +// the requested organization. +func (p Principal) RequireOrganization(id string) error { + target := strings.TrimSpace(id) + if target == "" { + return connect.NewError(connect.CodeInvalidArgument, errors.New("organization_id is required")) + } + if strings.TrimSpace(p.OrganizationID) != target { + return connect.NewError(connect.CodePermissionDenied, fmt.Errorf("organization %s not authorized", target)) + } + return nil +} + +// RequireWorkspace returns PermissionDenied when the caller is not scoped to +// the requested workspace. +func (p Principal) RequireWorkspace(id string) error { + target := strings.TrimSpace(id) + if target == "" { + return connect.NewError(connect.CodeInvalidArgument, errors.New("workspace_id is required")) + } + if strings.TrimSpace(p.WorkspaceID) != target { + return connect.NewError(connect.CodePermissionDenied, fmt.Errorf("workspace %s not authorized", target)) + } + return nil +} + +// RequireScope returns PermissionDenied when the caller lacks scope. +func (p Principal) RequireScope(scope string) error { + scope = strings.TrimSpace(scope) + if scope == "" || slices.Contains(p.Scopes, scope) { + return nil + } + return connect.NewError(connect.CodePermissionDenied, fmt.Errorf("missing scope %q", scope)) +} + +// RequireHuman rejects service and API-key principals for human-only actions. +func (p Principal) RequireHuman() error { + if p.IsHuman || isHumanActorType(p.TokenType) { + return nil + } + return connect.NewError(connect.CodePermissionDenied, errors.New("human principal required")) +} + +// RejectSelfApproval rejects attempts by an agent to approve its own request. +func (p Principal) RejectSelfApproval(agentID string) error { + agentID = strings.TrimSpace(agentID) + if agentID == "" { + return nil + } + for _, candidate := range []string{p.AgentID, p.Subject, p.UserSubject} { + if strings.TrimSpace(candidate) == agentID { + return connect.NewError(connect.CodePermissionDenied, errors.New("principal cannot approve its own request")) + } + } + return nil +} + +func isHumanActorType(actorType string) bool { + switch strings.ToLower(strings.TrimSpace(actorType)) { + case "human", "user": + return true + default: + return false + } +} diff --git a/downstream/downstream.go b/downstream/downstream.go index a168328..173be31 100644 --- a/downstream/downstream.go +++ b/downstream/downstream.go @@ -5,6 +5,7 @@ package downstream import ( "context" + "fmt" "log/slog" "time" @@ -15,6 +16,10 @@ import ( // FailureMode controls what happens when a downstream call fails. type FailureMode int +// FailurePolicy is kept as a source-compatible alias for services that adopted +// the earlier platform package name before this module exported downstream. +type FailurePolicy = FailureMode + const ( // FailClosed returns the error to the caller. Use for safety-critical // services where a degraded response is worse than no response. @@ -62,7 +67,16 @@ type Metrics struct { } // New creates a downstream client with the given name and configuration. -func New(name string, cfg Config) *Client { +// +// New accepts both the current shape: +// +// downstream.New("meter", downstream.Config{FailureMode: downstream.FailOpen}) +// +// and the earlier platform shape: +// +// downstream.New("meter", downstream.FailOpen, downstream.Config{}) +func New(name string, cfgOrPolicy any, configs ...Config) *Client { + cfg := normalizeConfig(cfgOrPolicy, configs...) logger := cfg.Logger if logger == nil { logger = slog.Default() @@ -76,6 +90,36 @@ func New(name string, cfg Config) *Client { } } +func normalizeConfig(cfgOrPolicy any, configs ...Config) Config { + switch value := cfgOrPolicy.(type) { + case Config: + if len(configs) != 0 { + panic(fmt.Sprintf("downstream.New: unexpected %d extra config arguments", len(configs))) + } + return value + case *Config: + if len(configs) != 0 { + panic(fmt.Sprintf("downstream.New: unexpected %d extra config arguments", len(configs))) + } + if value == nil { + return Config{} + } + return *value + case FailureMode: + if len(configs) > 1 { + panic(fmt.Sprintf("downstream.New: expected at most one config with failure policy, got %d", len(configs))) + } + cfg := Config{FailureMode: value} + if len(configs) == 1 { + cfg = configs[0] + cfg.FailureMode = value + } + return cfg + default: + panic(fmt.Sprintf("downstream.New: unsupported config argument type %T", cfgOrPolicy)) + } +} + // Name returns the downstream service name. func (c *Client) Name() string { return c.name } diff --git a/downstream/downstream_test.go b/downstream/downstream_test.go index 3490f26..ce06bce 100644 --- a/downstream/downstream_test.go +++ b/downstream/downstream_test.go @@ -309,6 +309,40 @@ func TestFailureModeString(t *testing.T) { } } +func TestFailurePolicyAliasAndLegacyConstructor(t *testing.T) { + policy := downstream.FailurePolicy(downstream.FailOpen) + c := downstream.New("meter", policy, downstream.Config{}) + + if c.Name() != "meter" { + t.Fatalf("expected meter, got %s", c.Name()) + } + if c.Mode() != downstream.FailOpen { + t.Fatalf("expected FailOpen, got %s", c.Mode()) + } +} + +func TestLegacyConstructorPreservesConfig(t *testing.T) { + breaker := resilience.NewBreaker(resilience.BreakerConfig{ + FailureThreshold: 1, + ResetTimeout: time.Hour, + }) + c := downstream.New("governance", downstream.FailClosed, downstream.Config{ + Breaker: breaker, + }) + + _, _ = downstream.Call(context.Background(), c, func(_ context.Context) (string, error) { + return "", errSimulated + }) + + _, err := downstream.Call(context.Background(), c, func(_ context.Context) (string, error) { + t.Fatal("fn should not be called when breaker is open") + return "", nil + }) + if !errors.Is(err, resilience.ErrCircuitOpen) { + t.Fatalf("expected ErrCircuitOpen, got: %v", err) + } +} + // --- Nil logger safety --- func TestNilLoggerDefaultsToSlog(t *testing.T) { diff --git a/httpkit/security.go b/httpkit/security.go new file mode 100644 index 0000000..9e71894 --- /dev/null +++ b/httpkit/security.go @@ -0,0 +1,146 @@ +package httpkit + +import ( + "mime" + "net/http" + "strings" +) + +// ErrorCodeCrossSiteRequest and related constants are standard error codes for +// browser security guardrail responses. +const ( + ErrorCodeCrossSiteRequest = "cross_site_request" + ErrorCodeMissingContentType = "missing_content_type" + ErrorCodeUnsupportedContentType = "unsupported_content_type" + defaultContentSecurityPolicy = "default-src 'self'; object-src 'none'; base-uri 'none'; frame-ancestors 'none'; script-src 'self' 'unsafe-inline'; style-src 'self' 'unsafe-inline'; img-src 'self' data:; connect-src 'self'" + defaultPermissionsPolicy = "camera=(), microphone=(), geolocation=(), payment=(), usb=()" + defaultCrossOriginOpenerPolicy = "same-origin" + defaultCrossOriginResourcePolicy = "same-origin" +) + +// WithBrowserSecurityDefaults applies the shared safe-by-default HTTP guardrail +// set for platform APIs. It intentionally uses low-risk primitives that work +// for Connect, JSON REST handlers, OAuth form posts, and internal CLIs. +func WithBrowserSecurityDefaults(next http.Handler) http.Handler { + return WithSecurityHeaders(WithFetchMetadataProtection(WithRequestContentTypeValidation(next))) +} + +// WithSecurityHeaders sets browser hardening headers unless a handler has +// already supplied service-specific values. +func WithSecurityHeaders(next http.Handler) http.Handler { + next = nonNilHandler(next) + return http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + header := writer.Header() + setHeaderIfEmpty(header, "X-Content-Type-Options", "nosniff") + setHeaderIfEmpty(header, "X-Frame-Options", "DENY") + setHeaderIfEmpty(header, "Referrer-Policy", "no-referrer") + setHeaderIfEmpty(header, "Content-Security-Policy", defaultContentSecurityPolicy) + setHeaderIfEmpty(header, "Permissions-Policy", defaultPermissionsPolicy) + setHeaderIfEmpty(header, "Cross-Origin-Opener-Policy", defaultCrossOriginOpenerPolicy) + setHeaderIfEmpty(header, "Cross-Origin-Resource-Policy", defaultCrossOriginResourcePolicy) + next.ServeHTTP(writer, request) + }) +} + +// WithFetchMetadataProtection rejects cross-site browser writes. Non-browser +// clients and webhooks generally omit Sec-Fetch-Site, so they continue through +// the normal auth and signature checks. +func WithFetchMetadataProtection(next http.Handler) http.Handler { + next = nonNilHandler(next) + return http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if unsafeHTTPMethod(request.Method) && strings.EqualFold(strings.TrimSpace(request.Header.Get("Sec-Fetch-Site")), "cross-site") { + WriteError(writer, http.StatusForbidden, ErrorCodeCrossSiteRequest, "cross-site browser requests are not allowed") + return + } + next.ServeHTTP(writer, request) + }) +} + +// WithRequestContentTypeValidation rejects body-bearing unsafe requests without +// an allowed Content-Type. This catches ambiguous browser form/text posts while +// preserving JSON, Connect, protobuf, and OAuth form traffic. +func WithRequestContentTypeValidation(next http.Handler) http.Handler { + next = nonNilHandler(next) + return http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if !unsafeHTTPMethod(request.Method) || !requestHasBody(request) { + next.ServeHTTP(writer, request) + return + } + + rawContentType := strings.TrimSpace(request.Header.Get("Content-Type")) + if rawContentType == "" { + WriteError(writer, http.StatusUnsupportedMediaType, ErrorCodeMissingContentType, "request body requires a Content-Type header") + return + } + mediaType, _, err := mime.ParseMediaType(rawContentType) + if err != nil || !allowedRequestContentType(mediaType) { + WriteError(writer, http.StatusUnsupportedMediaType, ErrorCodeUnsupportedContentType, "request Content-Type is not supported") + return + } + next.ServeHTTP(writer, request) + }) +} + +func allowedRequestContentType(mediaType string) bool { + mediaType = strings.ToLower(strings.TrimSpace(mediaType)) + switch { + case mediaType == "application/json": + return true + case strings.HasSuffix(mediaType, "+json"): + return true + case mediaType == "application/x-www-form-urlencoded": + return true + case mediaType == "multipart/form-data": + return true + case mediaType == "application/protobuf", mediaType == "application/x-protobuf", mediaType == "application/proto": + return true + case strings.HasSuffix(mediaType, "+proto"): + return true + case mediaType == "application/grpc": + return true + case strings.HasPrefix(mediaType, "application/grpc+"): + return true + case mediaType == "application/grpc-web", mediaType == "application/grpc-web-text": + return true + case strings.HasPrefix(mediaType, "application/grpc-web+"), strings.HasPrefix(mediaType, "application/grpc-web-text+"): + return true + case mediaType == "application/octet-stream": + return true + case mediaType == "application/x-ndjson": + return true + default: + return false + } +} + +func setHeaderIfEmpty(header http.Header, key string, value string) { + if header.Get(key) == "" { + header.Set(key, value) + } +} + +func unsafeHTTPMethod(method string) bool { + switch method { + case http.MethodGet, http.MethodHead, http.MethodOptions, http.MethodTrace: + return false + default: + return true + } +} + +func requestHasBody(request *http.Request) bool { + if request == nil || request.Body == nil || request.Body == http.NoBody { + return false + } + if request.ContentLength > 0 || request.ContentLength == -1 { + return true + } + return len(request.TransferEncoding) > 0 +} + +func nonNilHandler(next http.Handler) http.Handler { + if next == nil { + return http.DefaultServeMux + } + return next +} diff --git a/startup/http.go b/startup/http.go new file mode 100644 index 0000000..5e9b01e --- /dev/null +++ b/startup/http.go @@ -0,0 +1,145 @@ +package startup + +import ( + "context" + "errors" + "fmt" + "log/slog" + "net" + "net/http" + "os" + "os/signal" + "strings" + "syscall" + "time" + + "github.com/evalops/service-runtime/httpkit" +) + +// HTTPServerConfig describes a managed HTTP service startup. +type HTTPServerConfig struct { + ServiceName string + Addr string + Version string + Environment string + Server *http.Server + Listener net.Listener + ShutdownTimeout time.Duration + Lifecycle *Lifecycle + Logger *slog.Logger + TLSCertFile string + TLSKeyFile string + TLSClientCAFile string + ConnectServiceNames []string + GRPCReflectionEnabled *bool +} + +// NotifyContext returns a signal-aware root context for long-running services. +func NotifyContext() (context.Context, context.CancelFunc) { + return signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) +} + +// RunHTTPServer starts server and coordinates graceful shutdown with lifecycle hooks. +func RunHTTPServer(ctx context.Context, cfg HTTPServerConfig) error { + if cfg.Server == nil { + return fmt.Errorf("startup: server is required") + } + if ctx == nil { + ctx = context.Background() + } + + timeout := cfg.ShutdownTimeout + if timeout <= 0 { + timeout = DefaultShutdownTimeout + } + + logger := cfg.Logger + if logger == nil { + logger = slog.Default() + } + + serviceName := strings.TrimSpace(cfg.ServiceName) + if serviceName == "" { + serviceName = "service" + } + + lifecycle := cfg.Lifecycle + if lifecycle == nil { + lifecycle = NewLifecycle() + } + if err := lifecycle.EnableTracingFromEnv(ctx, serviceName); err != nil { + return fmt.Errorf("startup tracing: %w", err) + } + + addr := strings.TrimSpace(cfg.Addr) + if addr == "" { + addr = cfg.Server.Addr + } + if addr == "" && cfg.Listener != nil { + addr = cfg.Listener.Addr().String() + } + + cfg.Server.Handler = httpkit.WithBrowserSecurityDefaults(nonNilHTTPHandler(cfg.Server.Handler)) + + errCh := make(chan error, 1) + go func() { + logger.Info("starting service", "service", serviceName, "addr", addr, "version", strings.TrimSpace(cfg.Version)) + + var err error + switch { + case strings.TrimSpace(cfg.TLSCertFile) != "" && strings.TrimSpace(cfg.TLSKeyFile) != "": + logger.Info("tls enabled", "service", serviceName) + if strings.TrimSpace(cfg.TLSClientCAFile) != "" { + logger.Info("verified client certificates required", "service", serviceName) + } + if cfg.Listener != nil { + err = cfg.Server.ServeTLS(cfg.Listener, "", "") + } else { + err = cfg.Server.ListenAndServeTLS("", "") + } + case cfg.Listener != nil: + err = cfg.Server.Serve(cfg.Listener) + default: + err = cfg.Server.ListenAndServe() + } + + if err != nil && !errors.Is(err, http.ErrServerClosed) { + errCh <- err + } + close(errCh) + }() + + select { + case err := <-errCh: + if err == nil { + return nil + } + shutdownCtx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + if lifecycleErr := lifecycle.Shutdown(shutdownCtx); lifecycleErr != nil { + return errors.Join(fmt.Errorf("listen failed: %w", err), lifecycleErr) + } + return fmt.Errorf("listen failed: %w", err) + case <-ctx.Done(): + logger.Info("shutting down service", "service", serviceName) + shutdownCtx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + + var errs []error + if err := cfg.Server.Shutdown(shutdownCtx); err != nil && !errors.Is(err, context.Canceled) { + errs = append(errs, fmt.Errorf("shutdown server: %w", err)) + } + if lifecycleErr := lifecycle.Shutdown(shutdownCtx); lifecycleErr != nil { + errs = append(errs, lifecycleErr) + } + return errors.Join(errs...) + } +} + +func nonNilHTTPHandler(handler http.Handler) http.Handler { + if handler == nil { + return http.DefaultServeMux + } + return handler +} diff --git a/startup/http_test.go b/startup/http_test.go new file mode 100644 index 0000000..b71ae25 --- /dev/null +++ b/startup/http_test.go @@ -0,0 +1,95 @@ +package startup + +import ( + "context" + "io" + "log/slog" + "net" + "net/http" + "testing" + "time" +) + +func TestRunHTTPServerAppliesHTTPDefenses(t *testing.T) { + t.Parallel() + + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen test server: %v", err) + } + t.Cleanup(func() { _ = listener.Close() }) + addr := listener.Addr().String() + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + mux := http.NewServeMux() + mux.HandleFunc("/", func(writer http.ResponseWriter, request *http.Request) { + writer.WriteHeader(http.StatusNoContent) + }) + server := &http.Server{ + Addr: addr, + Handler: mux, + } + + errCh := make(chan error, 1) + go func() { + errCh <- RunHTTPServer(ctx, HTTPServerConfig{ + ServiceName: "test-service", + Addr: addr, + Server: server, + Listener: listener, + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + }) + }() + waitForServer(t, addr, errCh) + + client := &http.Client{Timeout: time.Second} + response, err := client.Get("http://" + addr + "/") + if err != nil { + t.Fatalf("GET server: %v", err) + } + _ = response.Body.Close() + if response.StatusCode != http.StatusNoContent { + t.Fatalf("GET status = %d, want %d", response.StatusCode, http.StatusNoContent) + } + assertResponseHeader(t, response.Header, "X-Content-Type-Options", "nosniff") + assertResponseHeader(t, response.Header, "X-Frame-Options", "DENY") + assertResponseHeader(t, response.Header, "Referrer-Policy", "no-referrer") + + cancel() + select { + case err := <-errCh: + if err != nil { + t.Fatalf("RunHTTPServer() error = %v", err) + } + case <-time.After(5 * time.Second): + t.Fatal("RunHTTPServer() did not return after shutdown") + } +} + +func waitForServer(t *testing.T, addr string, errCh <-chan error) { + t.Helper() + deadline := time.After(5 * time.Second) + for { + select { + case err := <-errCh: + t.Fatalf("server exited before accepting connections: %v", err) + case <-deadline: + t.Fatal("timed out waiting for server") + default: + conn, err := net.DialTimeout("tcp", addr, 50*time.Millisecond) + if err == nil { + _ = conn.Close() + return + } + time.Sleep(10 * time.Millisecond) + } + } +} + +func assertResponseHeader(t *testing.T, header http.Header, key string, want string) { + t.Helper() + if got := header.Get(key); got != want { + t.Fatalf("%s = %q, want %q", key, got, want) + } +}