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
69 changes: 57 additions & 12 deletions internal/api/handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -229,6 +229,10 @@ const apiKeyAdminUserID = "admin-api-key"
// permissions and is denied (fail closed). Returns the session on success so
// callers can read session.UserID for account filtering.
func (h *Handler) requirePermission(ctx context.Context, req *events.LambdaFunctionURLRequest, action, resource string) (*Session, error) {
if principal, ok := principalFromContext(ctx); ok {
return h.requirePrincipalPermission(ctx, req, principal, action, resource)
}

apiKey := extractAPIKey(req)
if h.checkAdminAPIKey(apiKey) {
return &Session{UserID: apiKeyAdminUserID}, nil
Expand Down Expand Up @@ -261,6 +265,33 @@ func (h *Handler) requirePermission(ctx context.Context, req *events.LambdaFunct
return nil, NewClientError(401, "invalid session")
}

return h.requireSessionPermission(ctx, session, action, resource)
}

func (h *Handler) requirePrincipalPermission(ctx context.Context, req *events.LambdaFunctionURLRequest, principal *Principal, action, resource string) (*Session, error) {
switch principal.Kind {
case PrincipalAdminAPIKey:
return &Session{UserID: apiKeyAdminUserID}, nil
case PrincipalUserAPIKey:
session, err := h.authorizeAPIKey(ctx, extractAPIKey(req), action, resource)
if err != nil {
return nil, err
}
if session == nil {
return nil, NewClientError(401, "API key permission validation failed")
}
return session, nil
case PrincipalSession:
return h.requireSessionPermission(ctx, principal.Session, action, resource)
default:
return nil, NewClientError(401, "unsupported authentication principal")
}
}

func (h *Handler) requireSessionPermission(ctx context.Context, session *Session, action, resource string) (*Session, error) {
if session == nil || session.UserID == "" {
return nil, NewClientError(401, "invalid session")
}
has, err := h.auth.HasPermissionAPI(ctx, session.UserID, action, resource)
if err != nil {
return nil, fmt.Errorf("permission check failed: %w", err)
Expand Down Expand Up @@ -423,12 +454,13 @@ func (h *Handler) HandleRequest(ctx context.Context, req *events.LambdaFunctionU
logging.Debugf("API Request: %s %s", method, redactURL(path))

// Validate request
if response := h.validateRequest(ctx, req, method, path, corsHeaders); response != nil {
requestCtx, response := h.validateRequestContext(ctx, req, method, path, corsHeaders)
if response != nil {
return response, nil
}

// Route and execute request
return h.executeRequest(ctx, method, path, req, corsHeaders)
return h.executeRequest(requestCtx, method, path, req, corsHeaders)
}

// buildResponseHeaders creates response headers with security and CORS settings.
Expand All @@ -451,43 +483,56 @@ func (h *Handler) buildResponseHeaders() map[string]string {

// validateRequest validates the incoming request and returns error response if validation fails.
func (h *Handler) validateRequest(ctx context.Context, req *events.LambdaFunctionURLRequest, method, path string, corsHeaders map[string]string) *events.LambdaFunctionURLResponse {
_, response := h.validateRequestContext(ctx, req, method, path, corsHeaders)
return response
}

func (h *Handler) validateRequestContext(ctx context.Context, req *events.LambdaFunctionURLRequest, method, path string, corsHeaders map[string]string) (context.Context, *events.LambdaFunctionURLResponse) {
// Validate request body size
if err := validateRequestBodySize(req.Body); err != nil {
logging.Warnf("Request body size exceeded: %d bytes", len(req.Body))
return h.buildResponse(413, corsHeaders, map[string]string{"error": "Request body too large"}, nil)
return ctx, h.buildResponse(413, corsHeaders, map[string]string{"error": "Request body too large"}, nil)
}

// Validate Content-Type
if err := validateContentType(req); err != nil {
return h.buildResponse(415, corsHeaders, map[string]string{"error": err.Error()}, nil)
return ctx, h.buildResponse(415, corsHeaders, map[string]string{"error": err.Error()}, nil)
}

// Validate authentication and CSRF
if response := h.validateSecurity(ctx, req, method, path, corsHeaders); response != nil {
return response
requestCtx, response := h.validateSecurityContext(ctx, req, method, path, corsHeaders)
if response != nil {
return ctx, response
}

return nil
return requestCtx, nil
}

// validateSecurity validates authentication and CSRF token.
func (h *Handler) validateSecurity(ctx context.Context, req *events.LambdaFunctionURLRequest, method, path string, corsHeaders map[string]string) *events.LambdaFunctionURLResponse {
_, response := h.validateSecurityContext(ctx, req, method, path, corsHeaders)
return response
}

func (h *Handler) validateSecurityContext(ctx context.Context, req *events.LambdaFunctionURLRequest, method, path string, corsHeaders map[string]string) (context.Context, *events.LambdaFunctionURLResponse) {
if h.isPublicEndpoint(path) {
return nil
return ctx, nil
}

if !h.authenticate(ctx, req) {
return h.buildResponse(401, corsHeaders, map[string]string{"error": "Unauthorized"}, nil)
principal, err := h.authenticatePrincipal(ctx, req)
if err != nil {
return ctx, h.buildResponse(401, corsHeaders, map[string]string{"error": "Unauthorized"}, nil)
}
ctx = contextWithPrincipal(ctx, principal)

if h.requiresCSRFValidation(method, path, req) {
if err := h.validateCSRF(ctx, req); err != nil {
logging.Warnf("CSRF validation failed: %v", err)
return h.buildResponse(403, corsHeaders, map[string]string{"error": "CSRF validation failed"}, nil)
return ctx, h.buildResponse(403, corsHeaders, map[string]string{"error": "CSRF validation failed"}, nil)
}
}

return nil
return ctx, nil
}

// executeRequest routes and executes the API request.
Expand Down
42 changes: 19 additions & 23 deletions internal/api/handler_auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -82,15 +82,9 @@ func (h *Handler) getCurrentUser(ctx context.Context, req *events.LambdaFunction
return nil, fmt.Errorf("authentication service not configured")
}

// Get token from Authorization header
token := h.extractBearerToken(req)
if token == "" {
return nil, NewClientError(401, "no authorization token provided")
}

session, err := h.auth.ValidateSession(ctx, token)
session, err := h.requireSessionPrincipal(ctx, req)
if err != nil {
return nil, NewClientError(401, "invalid session")
return nil, err
}

user, err := h.auth.GetUser(ctx, session.UserID)
Expand Down Expand Up @@ -176,6 +170,19 @@ func (h *Handler) getCurrentUserPermissions(ctx context.Context, req *events.Lam
// special-case it (see getCurrentUserPermissions for the {admin, *}
// short-circuit).
func (h *Handler) resolveAuthenticatedUserID(ctx context.Context, req *events.LambdaFunctionURLRequest) (string, error) {
if principal, ok := principalFromContext(ctx); ok {
if principal.Kind == PrincipalAdminAPIKey {
return apiKeyAdminUserID, nil
}
if principal.UserID == "" {
return "", NewClientError(401, "authenticated principal has no user ID")
}
return principal.UserID, nil
}
return h.resolveAuthenticatedUserIDFromRequest(ctx, req)
}

func (h *Handler) resolveAuthenticatedUserIDFromRequest(ctx context.Context, req *events.LambdaFunctionURLRequest) (string, error) {
// Admin API key first (stateless, no per-user lookup).
apiKey := extractAPIKey(req)
if h.checkAdminAPIKey(apiKey) {
Expand Down Expand Up @@ -370,15 +377,9 @@ func (h *Handler) updateProfile(ctx context.Context, req *events.LambdaFunctionU
return nil, fmt.Errorf("authentication service not configured")
}

// Get current user from token
token := h.extractBearerToken(req)
if token == "" {
return nil, NewClientError(401, "no authorization token provided")
}

session, err := h.auth.ValidateSession(ctx, token)
session, err := h.requireSessionPrincipal(ctx, req)
if err != nil {
return nil, NewClientError(401, "invalid session")
return nil, err
}

// Parse request body
Expand Down Expand Up @@ -475,14 +476,9 @@ func (h *Handler) changePassword(ctx context.Context, req *events.LambdaFunction
return nil, err
}

token := h.extractBearerToken(req)
if token == "" {
return nil, NewClientError(401, "no authorization token provided")
}

session, err := h.auth.ValidateSession(ctx, token)
session, err := h.requireSessionPrincipal(ctx, req)
if err != nil {
return nil, NewClientError(401, "invalid session")
return nil, err
}

var pwdReq ChangePasswordRequest
Expand Down
28 changes: 6 additions & 22 deletions internal/api/handler_purchases.go
Original file line number Diff line number Diff line change
Expand Up @@ -334,20 +334,15 @@ func (h *Handler) runPlannedPurchase(ctx context.Context, req *events.LambdaFunc
// (creator-scope cancel, issue #1400). Ownership is enforced separately by
// authorizePlannedPurchaseCancel; this gate is the minimum-verb check.
func (h *Handler) requireDeleteOrCancelPurchasePermission(ctx context.Context, req *events.LambdaFunctionURLRequest) (*Session, error) {
apiKey := extractAPIKey(req)
if h.checkAdminAPIKey(apiKey) {
if principal, ok := principalFromContext(ctx); ok && principal.Kind == PrincipalAdminAPIKey {
return &Session{UserID: apiKeyAdminUserID}, nil
}
if h.auth == nil {
return nil, fmt.Errorf("authentication service not configured")
}
token := h.extractBearerToken(req)
if token == "" {
return nil, NewClientError(401, "no authorization token provided")
if h.checkAdminAPIKey(extractAPIKey(req)) {
return &Session{UserID: apiKeyAdminUserID}, nil
}
session, err := h.auth.ValidateSession(ctx, token)
session, err := h.requireSessionPrincipal(ctx, req)
if err != nil {
return nil, NewClientError(401, "invalid session")
return nil, err
}
for _, verb := range []string{auth.ActionDelete, auth.ActionCancelAny, auth.ActionCancelOwn} {
has, checkErr := h.auth.HasPermissionAPI(ctx, session.UserID, verb, auth.ResourcePurchases)
Expand Down Expand Up @@ -1465,18 +1460,7 @@ func (h *Handler) authorizeSessionCancel(ctx context.Context, session *Session,
// cancel-own check; falling through to the API-key admin role would let
// a key impersonate ownership we cannot verify.
func (h *Handler) requireSession(ctx context.Context, req *events.LambdaFunctionURLRequest) (*Session, error) {
if h.auth == nil {
return nil, fmt.Errorf("authentication service not configured")
}
token := h.extractBearerToken(req)
if token == "" {
return nil, NewClientError(401, "no authorization token provided")
}
session, err := h.auth.ValidateSession(ctx, token)
if err != nil || session == nil {
return nil, NewClientError(401, "invalid session")
}
return session, nil
return h.requireSessionPrincipal(ctx, req)
}

// retryThreshold is the number of attempts after which a retry is
Expand Down
32 changes: 16 additions & 16 deletions internal/api/handler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -379,7 +379,7 @@ func TestHandler_HandleRequest_PutConfig(t *testing.T) {
}, nil)
mockAuth.On("ValidateSession", ctx, "test-token").Return(adminSession, nil)
mockAuth.grantAdmin()
mockAuth.On("ValidateCSRFToken", ctx, mock.Anything, mock.Anything).Return(nil)
mockAuth.On("ValidateCSRFToken", mock.Anything, mock.Anything, mock.Anything).Return(nil)

handler := &Handler{config: mockStore, auth: mockAuth, apiKey: "test-key"}

Expand Down Expand Up @@ -918,11 +918,11 @@ func TestHandler_HandleRequest_GetDashboardSummary(t *testing.T) {
DefaultCoverage: 80.0,
}

mockScheduler.On("ListRecommendations", ctx, mock.Anything).Return(recommendations, nil)
mockStore.On("GetGlobalConfig", ctx).Return(globalCfg, nil)
mockScheduler.On("ListRecommendations", mock.Anything, mock.Anything).Return(recommendations, nil)
mockStore.On("GetGlobalConfig", mock.Anything).Return(globalCfg, nil)
// No account_id / account_ids filter → calculateCommitmentMetrics fetches the
// uncapped active set across all accounts via GetActivePurchaseHistory.
mockStore.On("GetActivePurchaseHistory", ctx, mock.Anything, mock.Anything, mock.Anything).Return([]config.PurchaseHistoryRecord{}, nil)
mockStore.On("GetActivePurchaseHistory", mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return([]config.PurchaseHistoryRecord{}, nil)

handler := &Handler{
scheduler: mockScheduler,
Expand Down Expand Up @@ -975,11 +975,11 @@ func TestHandler_HandleRequest_GetUpcomingPurchases(t *testing.T) {
},
}

mockStore.On("ListPurchasePlans", ctx, config.PurchasePlanFilter{}).Return(plans, nil)
mockStore.On("ListPurchasePlans", mock.Anything, config.PurchasePlanFilter{}).Return(plans, nil)
// New: handler now enumerates pending executions per PR #213. Fixture
// supplies one pending exec for the plan above so the integration test
// still observes a single upcoming row.
mockStore.On("GetPendingExecutions", ctx).Return([]config.PurchaseExecution{
mockStore.On("GetPendingExecutions", mock.Anything).Return([]config.PurchaseExecution{
{
ExecutionID: "exec-int-1",
PlanID: plans[0].ID,
Expand Down Expand Up @@ -1034,8 +1034,8 @@ func TestHandler_HandleRequest_GetPlannedPurchases(t *testing.T) {
{ID: "11111111-1111-1111-1111-111111111111", Name: "Test Plan", Services: map[string]config.ServiceConfig{"aws/rds": {Provider: "aws", Service: "rds"}}},
}

mockStore.On("GetPlannedExecutions", ctx, mock.Anything, mock.Anything).Return(executions, nil)
mockStore.On("ListPurchasePlans", ctx, config.PurchasePlanFilter{}).Return(plans, nil)
mockStore.On("GetPlannedExecutions", mock.Anything, mock.Anything, mock.Anything).Return(executions, nil)
mockStore.On("ListPurchasePlans", mock.Anything, config.PurchasePlanFilter{}).Return(plans, nil)

handler := &Handler{config: mockStore, auth: mockAuth, corsAllowedOrigin: "*", apiKey: "test-key"}

Expand Down Expand Up @@ -1068,7 +1068,7 @@ func TestHandler_HandleRequest_PausePlannedPurchase(t *testing.T) {
mockAuth.On("ValidateCSRFToken", ctx, mock.Anything, mock.Anything).Return(nil)

paused := &config.PurchaseExecution{ExecutionID: "11111111-1111-1111-1111-111111111111", Status: "paused"}
mockStore.On("TransitionExecutionStatus", ctx, "11111111-1111-1111-1111-111111111111", []string{"pending", "running"}, "paused", mock.Anything).Return(paused, nil)
mockStore.On("TransitionExecutionStatus", mock.Anything, "11111111-1111-1111-1111-111111111111", []string{"pending", "running"}, "paused", mock.Anything).Return(paused, nil)

handler := &Handler{config: mockStore, auth: mockAuth, corsAllowedOrigin: "*", apiKey: "test-key"}

Expand Down Expand Up @@ -1103,7 +1103,7 @@ func TestHandler_HandleRequest_ResumePlannedPurchase(t *testing.T) {
mockAuth.On("ValidateCSRFToken", ctx, mock.Anything, mock.Anything).Return(nil)

resumed := &config.PurchaseExecution{ExecutionID: "11111111-1111-1111-1111-111111111111", Status: "pending"}
mockStore.On("TransitionExecutionStatus", ctx, "11111111-1111-1111-1111-111111111111", []string{"paused"}, "pending", mock.Anything).Return(resumed, nil)
mockStore.On("TransitionExecutionStatus", mock.Anything, "11111111-1111-1111-1111-111111111111", []string{"paused"}, "pending", mock.Anything).Return(resumed, nil)

handler := &Handler{config: mockStore, auth: mockAuth, corsAllowedOrigin: "*", apiKey: "test-key"}

Expand Down Expand Up @@ -1138,7 +1138,7 @@ func TestHandler_HandleRequest_RunPlannedPurchase(t *testing.T) {
mockAuth.On("ValidateCSRFToken", ctx, mock.Anything, mock.Anything).Return(nil)

transitioned := &config.PurchaseExecution{ExecutionID: "11111111-1111-1111-1111-111111111111", Status: "running"}
mockStore.On("TransitionExecutionStatus", ctx, "11111111-1111-1111-1111-111111111111", []string{"pending", "paused"}, "running", mock.Anything).Return(transitioned, nil)
mockStore.On("TransitionExecutionStatus", mock.Anything, "11111111-1111-1111-1111-111111111111", []string{"pending", "paused"}, "running", mock.Anything).Return(transitioned, nil)

handler := &Handler{config: mockStore, auth: mockAuth, corsAllowedOrigin: "*", apiKey: "test-key"}

Expand Down Expand Up @@ -1173,7 +1173,7 @@ func TestHandler_HandleRequest_DeletePlannedPurchase(t *testing.T) {
mockAuth.On("ValidateCSRFToken", ctx, mock.Anything, mock.Anything).Return(nil)

canceled := &config.PurchaseExecution{ExecutionID: "11111111-1111-1111-1111-111111111111", Status: "canceled"}
mockStore.On("TransitionExecutionStatus", ctx, "11111111-1111-1111-1111-111111111111", []string{"pending", "paused"}, "canceled", mock.Anything).Return(canceled, nil)
mockStore.On("TransitionExecutionStatus", mock.Anything, "11111111-1111-1111-1111-111111111111", []string{"pending", "paused"}, "canceled", mock.Anything).Return(canceled, nil)

handler := &Handler{config: mockStore, auth: mockAuth, corsAllowedOrigin: "*", apiKey: "test-key"}

Expand Down Expand Up @@ -1212,9 +1212,9 @@ func TestHandler_HandleRequest_CreatePlannedPurchases(t *testing.T) {
RampSchedule: config.RampSchedule{StepIntervalDays: 7},
}

mockStore.On("GetPurchasePlan", ctx, "11111111-1111-1111-1111-111111111111").Return(plan, nil)
mockStore.On("SavePurchaseExecution", ctx, mock.Anything).Return(nil)
mockStore.On("UpdatePurchasePlan", ctx, mock.Anything).Return(nil)
mockStore.On("GetPurchasePlan", mock.Anything, "11111111-1111-1111-1111-111111111111").Return(plan, nil)
mockStore.On("SavePurchaseExecution", mock.Anything, mock.Anything).Return(nil)
mockStore.On("UpdatePurchasePlan", mock.Anything, mock.Anything).Return(nil)

handler := &Handler{config: mockStore, auth: mockAuth, corsAllowedOrigin: "*", apiKey: "test-key"}

Expand Down Expand Up @@ -1281,7 +1281,7 @@ func TestHandler_HandleRequest_DeleteUser_SelfDeletion(t *testing.T) {
adminSession := &Session{UserID: adminUserID, Email: "admin@example.com"}
mockAuth.On("ValidateSession", ctx, "test-token").Return(adminSession, nil)
mockAuth.grantAdmin()
mockAuth.On("ValidateCSRFToken", ctx, mock.Anything, mock.Anything).Return(nil)
mockAuth.On("ValidateCSRFToken", mock.Anything, mock.Anything, mock.Anything).Return(nil)

handler := &Handler{auth: mockAuth, apiKey: "test-key"}

Expand Down
Loading
Loading