From 09dcc796a3dbd163ad0c8b7018a02e2bfe36a218 Mon Sep 17 00:00:00 2001 From: Aloento <11802769+Aloento@users.noreply.github.com> Date: Sun, 4 Oct 2026 22:41:46 +0200 Subject: [PATCH 1/2] Thread request and task context through the DB layer Add a ctx context.Context parameter as the first argument to the DB methods that previously hardcoded context.Background() (28 in db.go, plus getEventsByType and EnsureNotificationSchema), and pass the request or task context down from the API handlers, the RSS feeds, the middleware, and the checker. Behavior is unchanged: no new timeouts, no new configuration, no changes to business logic or API output. Request cancellation and timeouts can now propagate to database queries. --- internal/api/middleware.go | 4 +- internal/api/rss/rss.go | 13 ++-- internal/api/v2/v2.go | 112 ++++++++++++++------------- internal/app/app.go | 2 +- internal/checker/checker.go | 9 +-- internal/checker/info.go | 7 +- internal/checker/maintenance.go | 16 ++-- internal/db/db.go | 114 ++++++++++++---------------- internal/db/event_types.go | 4 +- internal/db/info.go | 6 +- internal/db/maintenances.go | 6 +- internal/db/notification_ops.go | 4 +- internal/rss/rss.go | 13 ++-- tests/checker_notifications_test.go | 9 ++- tests/db_tx_test.go | 12 +-- tests/notifications_ops_test.go | 2 +- tests/notifications_test.go | 2 +- 17 files changed, 164 insertions(+), 171 deletions(-) diff --git a/internal/api/middleware.go b/internal/api/middleware.go index 28cb348..734622b 100644 --- a/internal/api/middleware.go +++ b/internal/api/middleware.go @@ -38,7 +38,7 @@ func ValidateComponentsMW(dbInst *db.DB, logger *zap.Logger) gin.HandlerFunc { // TODO: move this list to the memory cache // We should check, that all components are presented in our db. - dbComps, err := dbInst.GetComponentsAsMap() + dbComps, err := dbInst.GetComponentsAsMap(c.Request.Context()) if err != nil { apiErrors.RaiseInternalErr(c, err) return @@ -228,7 +228,7 @@ func CheckEventExistenceMW(dbInst *db.DB, logger *zap.Logger) gin.HandlerFunc { return } - event, err := dbInst.GetIncident(incID.ID) + event, err := dbInst.GetIncident(c.Request.Context(), incID.ID) if err != nil { if errors.Is(err, db.ErrDBIncidentDSNotExist) { apiErrors.RaiseStatusNotFoundErr(c, apiErrors.ErrIncidentDSNotExist) diff --git a/internal/api/rss/rss.go b/internal/api/rss/rss.go index 7a84bf8..78f7627 100644 --- a/internal/api/rss/rss.go +++ b/internal/api/rss/rss.go @@ -1,6 +1,7 @@ package rss import ( + "context" "fmt" "net/http" "sort" @@ -45,7 +46,7 @@ func HandleRSS(dbInst *db.DB, logger *zap.Logger) gin.HandlerFunc { baseURL: baseURL, } - incidents, err := getIncidents(dbInst, logger, params, maxIncidents) + incidents, err := getIncidents(c.Request.Context(), dbInst, logger, params, maxIncidents) if err != nil { if componentName != "" { apiErrors.RaiseStatusNotFoundErr(c, err) @@ -102,7 +103,7 @@ type feedParams struct { baseURL string } -func getIncidents(dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int) ([]*db.Incident, error) { +func getIncidents(ctx context.Context, dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int) ([]*db.Incident, error) { var incidents []*db.Incident var err error @@ -116,13 +117,13 @@ func getIncidents(dbInstance *db.DB, log *zap.Logger, params feedParams, maxInci } var component *db.Component - component, err = dbInstance.GetComponentFromNameAttrs(params.componentName, attr) + component, err = dbInstance.GetComponentFromNameAttrs(ctx, params.componentName, attr) if err != nil { log.Error("failed to get component", zap.Error(err)) return nil, err } - incidents, err = dbInstance.GetEventsByComponentID(component.ID, incParams) + incidents, err = dbInstance.GetEventsByComponentID(ctx, component.ID, incParams) if err != nil { return nil, err } @@ -132,12 +133,12 @@ func getIncidents(dbInstance *db.DB, log *zap.Logger, params feedParams, maxInci Value: params.region, } - incidents, err = dbInstance.GetIncidentsByComponentAttr(attr, incParams) + incidents, err = dbInstance.GetIncidentsByComponentAttr(ctx, attr, incParams) if err != nil { return nil, err } default: - incidents, err = dbInstance.GetEvents(db.PublicAccess, incParams) + incidents, err = dbInstance.GetEvents(ctx, db.PublicAccess, incParams) if err != nil { return nil, err } diff --git a/internal/api/v2/v2.go b/internal/api/v2/v2.go index 9ace5b4..7f8be23 100644 --- a/internal/api/v2/v2.go +++ b/internal/api/v2/v2.go @@ -196,7 +196,7 @@ func GetIncidentsHandler(dbInst *db.DB, logger *zap.Logger, svc *rbac.Service) g isAuth := hasExtendedView(c, svc) logger.Debug("retrieve incidents with params", zap.Any("params", params)) - r, err := dbInst.GetEvents(isAuth, params) + r, err := dbInst.GetEvents(c.Request.Context(), isAuth, params) if err != nil { logger.Error("failed to retrieve incidents", zap.Error(err)) apiErrors.RaiseInternalErr(c, err) @@ -235,7 +235,7 @@ func GetEventsHandler(dbInst *db.DB, logger *zap.Logger, svc *rbac.Service) gin. isAuth := hasExtendedView(c, svc) logger.Debug("retrieve events with params", zap.Any("params", params)) - r, total, err := dbInst.GetEventsWithCount(isAuth, params) + r, total, err := dbInst.GetEventsWithCount(c.Request.Context(), isAuth, params) if err != nil { logger.Error("failed to retrieve incidents", zap.Error(err)) apiErrors.RaiseInternalErr(c, err) @@ -296,7 +296,7 @@ func GetIncidentHandler(dbInst *db.DB, logger *zap.Logger, svc *rbac.Service) gi return } - r, err := dbInst.GetIncident(incID.ID) + r, err := dbInst.GetIncident(c.Request.Context(), incID.ID) if err != nil { if errors.Is(err, db.ErrDBIncidentDSNotExist) { apiErrors.RaiseStatusNotFoundErr(c, apiErrors.ErrIncidentDSNotExist) @@ -418,17 +418,18 @@ func PostIncidentHandler(dbInst *db.DB, logger *zap.Logger, pub ...*notification func routeIncidentCreation( c *gin.Context, dbInst *db.DB, log *zap.Logger, incData IncidentData, pub *notification.Publisher, ) ([]*ProcessComponentResp, error) { + ctx := c.Request.Context() if *incData.System { log.Info("system incident detected, using system incident creation logic") - return handleSystemIncidentCreation(dbInst, log, incData, pub) + return handleSystemIncidentCreation(ctx, dbInst, log, incData, pub) } log.Info("regular incident detected, using regular incident creation logic") userID := getUserIDFromContext(c) - return handleRegularIncidentCreation(dbInst, log, incData, userID, pub) + return handleRegularIncidentCreation(ctx, dbInst, log, incData, userID, pub) } func handleSystemIncidentCreation( - dbInst *db.DB, log *zap.Logger, incData IncidentData, _ *notification.Publisher, + ctx context.Context, dbInst *db.DB, log *zap.Logger, incData IncidentData, _ *notification.Publisher, ) ([]*ProcessComponentResp, error) { if incData.Type != event.TypeIncident { log.Info("system incident must be of type 'incident'") @@ -439,14 +440,14 @@ func handleSystemIncidentCreation( incData.Description = "System-wide incident affecting multiple components. Created automatically." } - components, err := fetchComponents(dbInst, incData.Components) + components, err := fetchComponents(ctx, dbInst, incData.Components) if err != nil { return nil, err } result := make([]*ProcessComponentResp, 0, len(components)) for _, comp := range components { - compResult, errProc := processSystemIncidentComponent(dbInst, log, comp, incData) + compResult, errProc := processSystemIncidentComponent(ctx, dbInst, log, comp, incData) if errProc != nil { return nil, errProc } @@ -456,10 +457,10 @@ func handleSystemIncidentCreation( return result, nil } -func fetchComponents(dbInst *db.DB, componentIDs []int) ([]db.Component, error) { +func fetchComponents(ctx context.Context, dbInst *db.DB, componentIDs []int) ([]db.Component, error) { components := make([]db.Component, len(componentIDs)) for i, compID := range componentIDs { - component, err := dbInst.GetComponent(compID) + component, err := dbInst.GetComponent(ctx, compID) if err != nil { return nil, err } @@ -469,38 +470,38 @@ func fetchComponents(dbInst *db.DB, componentIDs []int) ([]db.Component, error) } func processSystemIncidentComponent( - dbInst *db.DB, log *zap.Logger, comp db.Component, incData IncidentData, + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp db.Component, incData IncidentData, ) (*ProcessComponentResp, error) { log.Info("start to process component", zap.Any("component", comp)) log.Info("find events with target component", zap.Uint("componentID", comp.ID)) - events, err := getActiveEventsForComponent(dbInst, comp.ID) + events, err := getActiveEventsForComponent(ctx, dbInst, comp.ID) if err != nil { return nil, err } log.Info("found events for component", zap.Uint("componentID", comp.ID), zap.Int("eventsCount", len(events))) if len(events) == 0 { - return handleComponentWithNoEvents(dbInst, log, &comp, incData) + return handleComponentWithNoEvents(ctx, dbInst, log, &comp, incData) } - return handleComponentWithExistingEvents(dbInst, log, &comp, incData, events) + return handleComponentWithExistingEvents(ctx, dbInst, log, &comp, incData, events) } -func getActiveEventsForComponent(dbInst *db.DB, componentID uint) ([]*db.Incident, error) { +func getActiveEventsForComponent(ctx context.Context, dbInst *db.DB, componentID uint) ([]*db.Incident, error) { active := true params := &db.IncidentsParams{ IsActive: &active, Types: []string{event.TypeIncident, event.TypeMaintenance}, } - return dbInst.GetEventsByComponentID(componentID, params) + return dbInst.GetEventsByComponentID(ctx, componentID, params) } func handleComponentWithNoEvents( - dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, ) (*ProcessComponentResp, error) { log.Info("no events found for component, check and process all system incidents", zap.Uint("componentID", comp.ID)) - sysInc, err := addComponentToSystemIncident(dbInst, log, comp, incData) + sysInc, err := addComponentToSystemIncident(ctx, dbInst, log, comp, incData) if err != nil { return nil, err } @@ -515,7 +516,7 @@ func handleComponentWithNoEvents( } func handleComponentWithExistingEvents( - dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, events []*db.Incident, + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, events []*db.Incident, ) (*ProcessComponentResp, error) { log.Info("checking events for the component", zap.Uint("componentID", comp.ID), zap.Int("eventsCount", len(events))) @@ -554,7 +555,7 @@ func handleComponentWithExistingEvents( // If we found a system incident, handle it if firstSystemIncident != nil { - return handleSystemIncidentWithImpactComparison(dbInst, log, comp, incData, firstSystemIncident) + return handleSystemIncidentWithImpactComparison(ctx, dbInst, log, comp, incData, firstSystemIncident) } // This should not be reached - if we have events, one of the conditions above should handle it @@ -563,7 +564,7 @@ func handleComponentWithExistingEvents( } func handleSystemIncidentWithImpactComparison( - dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, evnt *db.Incident, + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, evnt *db.Incident, ) (*ProcessComponentResp, error) { log.Info( "found system incident for the component, compare impact", @@ -587,7 +588,7 @@ func handleSystemIncidentWithImpactComparison( zap.Uint("componentID", comp.ID), zap.Uint("fromIncidentID", evnt.ID), ) - sysInc, err := moveComponentFromToSystemIncidents(dbInst, log, comp, incData, evnt) + sysInc, err := moveComponentFromToSystemIncidents(ctx, dbInst, log, comp, incData, evnt) if err != nil { return nil, err } @@ -602,7 +603,7 @@ func handleSystemIncidentWithImpactComparison( } func addComponentToSystemIncident( - dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, ) (*db.Incident, error) { system := true active := true @@ -612,7 +613,7 @@ func addComponentToSystemIncident( IsSystem: &system, IsActive: &active, } - sysIncidents, errEvents := dbInst.GetEventsInternal(params) + sysIncidents, errEvents := dbInst.GetEventsInternal(ctx, params) if errEvents != nil { return nil, errEvents } @@ -632,7 +633,7 @@ func addComponentToSystemIncident( Text: fmt.Sprintf("%s added to the incident by system.", comp.PrintAttrs()), Timestamp: time.Now().UTC(), } - err := dbInst.AddComponentToIncident(sysInc, comp, status) + err := dbInst.AddComponentToIncident(ctx, sysInc, comp, status) if err != nil { return nil, err } @@ -657,7 +658,7 @@ func addComponentToSystemIncident( Components: []db.Component{*comp}, } - if err := createEvent(dbInst, log, &incIn, nil, nil); err != nil { + if err := createEvent(ctx, dbInst, log, &incIn, nil, nil); err != nil { return nil, err } @@ -665,7 +666,7 @@ func addComponentToSystemIncident( } func moveComponentFromToSystemIncidents( - dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, oldInc *db.Incident, + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, incData IncidentData, oldInc *db.Incident, ) (*db.Incident, error) { system := true active := true @@ -675,7 +676,7 @@ func moveComponentFromToSystemIncidents( IsSystem: &system, IsActive: &active, } - sysIncidents, errEvents := dbInst.GetEventsInternal(params) + sysIncidents, errEvents := dbInst.GetEventsInternal(ctx, params) if errEvents != nil { return nil, errEvents } @@ -694,7 +695,7 @@ func moveComponentFromToSystemIncidents( closeOld = true } - inc, err := dbInst.MoveComponentFromOldToAnotherIncident(comp, oldInc, sysInc, closeOld) + inc, err := dbInst.MoveComponentFromOldToAnotherIncident(ctx, comp, oldInc, sysInc, closeOld) if err != nil { return nil, err } @@ -712,7 +713,7 @@ func moveComponentFromToSystemIncidents( "the source incident has only 1 target component with the lower impact, we can just update its impact", zap.Uint("componentID", comp.ID), zap.Uint("incidentID", oldInc.ID), ) - inc, err := dbInst.IncreaseIncidentImpact(oldInc, *incData.Impact) + inc, err := dbInst.IncreaseIncidentImpact(ctx, oldInc, *incData.Impact) if err != nil { return nil, err } @@ -726,6 +727,7 @@ func moveComponentFromToSystemIncidents( ) inc, err := dbInst.ExtractComponentsToNewIncident( + ctx, []db.Component{*comp}, oldInc, *incData.Impact, @@ -738,7 +740,7 @@ func moveComponentFromToSystemIncidents( // Update the new incident to mark it as a system incident inc.System = true - if err = dbInst.ModifyIncident(inc); err != nil { + if err = dbInst.ModifyIncident(ctx, inc); err != nil { return nil, err } @@ -746,7 +748,7 @@ func moveComponentFromToSystemIncidents( } func handleRegularIncidentCreation( - dbInst *db.DB, log *zap.Logger, incData IncidentData, userID *string, pub *notification.Publisher, + ctx context.Context, dbInst *db.DB, log *zap.Logger, incData IncidentData, userID *string, pub *notification.Publisher, ) ([]*ProcessComponentResp, error) { components := make([]db.Component, len(incData.Components)) for i, comp := range incData.Components { @@ -776,14 +778,14 @@ func handleRegularIncidentCreation( log.Info("get active events from the database") isActive := true - openedIncidents, err := dbInst.GetEventsInternal(&db.IncidentsParams{IsActive: &isActive}) + openedIncidents, err := dbInst.GetEventsInternal(ctx, &db.IncidentsParams{IsActive: &isActive}) if err != nil { return nil, err } log.Info("opened incidents and maintenances retrieved", zap.Any("openedIncidents", openedIncidents)) - if err = createEvent(dbInst, log, &incIn, userID, pub); err != nil { + if err = createEvent(ctx, dbInst, log, &incIn, userID, pub); err != nil { return nil, err } @@ -793,7 +795,7 @@ func handleRegularIncidentCreation( } // Process component movement for complex cases - return processComponentMovement(dbInst, log, &incIn, openedIncidents) + return processComponentMovement(ctx, dbInst, log, &incIn, openedIncidents) } func shouldSkipComponentMovement(openedIncidents []*db.Incident, incData IncidentData) bool { @@ -818,13 +820,13 @@ func createSimpleIncidentResult(log *zap.Logger, incIn *db.Incident, incData Inc } func processComponentMovement( - dbInst *db.DB, log *zap.Logger, incIn *db.Incident, openedIncidents []*db.Incident, + ctx context.Context, dbInst *db.DB, log *zap.Logger, incIn *db.Incident, openedIncidents []*db.Incident, ) ([]*ProcessComponentResp, error) { log.Info("start to analyse component movement") result := make([]*ProcessComponentResp, 0, len(incIn.Components)) for _, comp := range incIn.Components { - compResult, err := processComponentInOpenedIncidents(dbInst, log, &comp, incIn, openedIncidents) + compResult, err := processComponentInOpenedIncidents(ctx, dbInst, log, &comp, incIn, openedIncidents) if err != nil { return nil, err } @@ -836,7 +838,7 @@ func processComponentMovement( // processComponentInOpenedIncidents processes a single component against all opened incidents. func processComponentInOpenedIncidents( - dbInst *db.DB, log *zap.Logger, comp *db.Component, incIn *db.Incident, openedIncidents []*db.Incident, + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, incIn *db.Incident, openedIncidents []*db.Incident, ) (*ProcessComponentResp, error) { compResult := &ProcessComponentResp{ ComponentID: int(comp.ID), @@ -851,7 +853,7 @@ func processComponentInOpenedIncidents( continue } - moved, err := tryMoveComponentIfFound(dbInst, log, comp, inc, incIn, compResult) + moved, err := tryMoveComponentIfFound(ctx, dbInst, log, comp, inc, incIn, compResult) if err != nil { return nil, err } @@ -875,6 +877,7 @@ func shouldSkipIncident(inc *db.Incident) bool { // tryMoveComponentIfFound attempts to move a component if it's found in the given incident. func tryMoveComponentIfFound( + ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, @@ -887,7 +890,7 @@ func tryMoveComponentIfFound( log.Info("found the component in the opened incident", zap.Any("component", comp), zap.Any("incident", inc)) closeInc := len(inc.Components) == 1 - incident, err := dbInst.MoveComponentFromOldToAnotherIncident(comp, inc, incIn, closeInc) + incident, err := dbInst.MoveComponentFromOldToAnotherIncident(ctx, comp, inc, incIn, closeInc) if err != nil { return false, err } @@ -966,11 +969,11 @@ func validateEventCreationTimes(incData IncidentData) error { return nil } -func createEvent(dbInst *db.DB, log *zap.Logger, inc *db.Incident, userID *string, pub *notification.Publisher) error { +func createEvent(ctx context.Context, dbInst *db.DB, log *zap.Logger, inc *db.Incident, userID *string, pub *notification.Publisher) error { log.Info("start to save an event to the database") - err := dbInst.WithTx(context.Background(), func(tx *db.Tx) error { - id, err := dbInst.SaveIncidentTx(tx, inc) + err := dbInst.WithTx(ctx, func(tx *db.Tx) error { + id, err := dbInst.SaveIncidentTx(ctx, tx, inc) if err != nil { return err } @@ -1017,12 +1020,12 @@ func createEvent(dbInst *db.DB, log *zap.Logger, inc *db.Incident, userID *strin }) inc.Status = status - if err = dbInst.ModifyIncidentTx(tx, inc); err != nil { + if err = dbInst.ModifyIncidentTx(ctx, tx, inc); err != nil { return err } // A newly created maintenance has no previous status. - return publishMaintenanceChange(context.Background(), tx, pub, inc, "", userID) + return publishMaintenanceChange(ctx, tx, pub, inc, "", userID) }) if err == nil && inc.Type == event.TypeMaintenance { pub.Notify() // wake the worker after the commit @@ -1075,7 +1078,7 @@ func persistIncidentPatch( ) bool { statusChanged := storedIncident.Status != oldStatus err := dbInst.WithTx(c.Request.Context(), func(tx *db.Tx) error { - if e := dbInst.ModifyIncidentTx(tx, storedIncident); e != nil { + if e := dbInst.ModifyIncidentTx(c.Request.Context(), tx, storedIncident); e != nil { return e } if !statusChanged { @@ -1182,7 +1185,7 @@ func PatchIncidentHandler(dbInst *db.DB, logger *zap.Logger, pub ...*notificatio return } - inc, errDB := dbInst.GetIncident(int(storedIncident.ID)) + inc, errDB := dbInst.GetIncident(c.Request.Context(), int(storedIncident.ID)) if errDB != nil { logger.Error("incident patch: failed to retrieve updated event", zap.Uint("event_id", storedIncident.ID), zap.Error(errDB)) @@ -1204,7 +1207,7 @@ func reopenIncident( logger.Info("reopening incident", zap.Uint("event_id", storedIncident.ID), ) - err := dbInst.ReOpenIncident(storedIncident) + err := dbInst.ReOpenIncident(c.Request.Context(), storedIncident) if err != nil { logger.Error("incident reopen failed: database error", zap.Uint("event_id", storedIncident.ID), zap.Error(err)) @@ -1411,6 +1414,7 @@ func PostIncidentExtractHandler(dbInst *db.DB, logger *zap.Logger) gin.HandlerFu } inc, err := dbInst.ExtractComponentsToNewIncident( + c.Request.Context(), movedComponents, storedInc, *storedInc.Impact, @@ -1469,7 +1473,7 @@ func GetComponentsHandler(dbInst *db.DB, logger *zap.Logger) gin.HandlerFunc { return func(c *gin.Context) { logger.Debug("retrieve components") - r, err := dbInst.GetComponentsWithValues() + r, err := dbInst.GetComponentsWithValues(c.Request.Context()) if err != nil { apiErrors.RaiseInternalErr(c, err) return @@ -1489,7 +1493,7 @@ func GetComponentHandler(dbInst *db.DB, logger *zap.Logger) gin.HandlerFunc { return } - r, err := dbInst.GetComponent(compID.ID) + r, err := dbInst.GetComponent(c.Request.Context(), compID.ID) if err != nil { if errors.Is(err, db.ErrDBComponentDSNotExist) { apiErrors.RaiseStatusNotFoundErr(c, apiErrors.ErrComponentDSNotExist) @@ -1536,7 +1540,7 @@ func PostComponentHandler(dbInst *db.DB, logger *zap.Logger) gin.HandlerFunc { Attrs: attrs, } - componentID, err := dbInst.SaveComponent(compDB) + componentID, err := dbInst.SaveComponent(c.Request.Context(), compDB) if err != nil { if errors.Is(err, db.ErrDBComponentExists) { apiErrors.RaiseBadRequestErr(c, apiErrors.ErrComponentExist) @@ -1592,7 +1596,7 @@ func GetComponentsAvailabilityHandler(dbInst *db.DB, logger *zap.Logger) gin.Han return func(c *gin.Context) { logger.Debug("retrieve availability of components") - components, err := dbInst.GetComponentsWithIncidents() + components, err := dbInst.GetComponentsWithIncidents(c.Request.Context()) if err != nil { apiErrors.RaiseInternalErr(c, err) return @@ -1821,7 +1825,7 @@ func PatchEventUpdateTextHandler(dbInst *db.DB, logger *zap.Logger) gin.HandlerF } // Update existence check. - updates, err := dbInst.GetEventUpdates(uint(incID)) + updates, err := dbInst.GetEventUpdates(c.Request.Context(), uint(incID)) if err != nil { apiErrors.RaiseInternalErr(c, err) return @@ -1836,7 +1840,7 @@ func PatchEventUpdateTextHandler(dbInst *db.DB, logger *zap.Logger) gin.HandlerF targetUPD.Text = text targetUPD.ModifiedBy = getUserIDFromContext(c) - updated, err := dbInst.ModifyEventUpdate(targetUPD) + updated, err := dbInst.ModifyEventUpdate(c.Request.Context(), targetUPD) if err != nil { apiErrors.RaiseInternalErr(c, err) diff --git a/internal/app/app.go b/internal/app/app.go index 47b92f5..21fcf28 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -92,7 +92,7 @@ func buildWorker( return nil, nil, nil } - if err = dbNew.EnsureNotificationSchema(); err != nil { + if err = dbNew.EnsureNotificationSchema(context.Background()); err != nil { return nil, nil, err } diff --git a/internal/checker/checker.go b/internal/checker/checker.go index ba668fe..2d9f9d6 100644 --- a/internal/checker/checker.go +++ b/internal/checker/checker.go @@ -26,9 +26,8 @@ func New(database *db.DB, log *zap.Logger, notifier *notification.Publisher) *Ch // Check runs one full scan and returns the combined error of its two halves. It // is the body of the scheduler's scan task, which holds the advisory lock for the -// whole round. Cancellation is observed only before the round starts: the two -// scans do not take a context yet, so a caller must not close the pool while -// Check is running. +// whole round. Cancellation is observed only before the round starts, so a +// caller must not close the pool while Check is running. func (ch *Checker) Check(ctx context.Context) error { if err := ctx.Err(); err != nil { return err @@ -43,7 +42,7 @@ func (ch *Checker) Check(ctx context.Context) error { wg.Add(1) go func() { defer wg.Done() - if err := ch.CheckMaintenance(); err != nil { + if err := ch.CheckMaintenance(ctx); err != nil { ch.log.Error("error to check maintenances", zap.Error(err)) mntErr = err } @@ -52,7 +51,7 @@ func (ch *Checker) Check(ctx context.Context) error { wg.Add(1) go func() { defer wg.Done() - if err := ch.CheckInfoEvents(); err != nil { + if err := ch.CheckInfoEvents(ctx); err != nil { ch.log.Error("error to check info events", zap.Error(err)) infoErr = err } diff --git a/internal/checker/info.go b/internal/checker/info.go index a58c613..e709e02 100644 --- a/internal/checker/info.go +++ b/internal/checker/info.go @@ -1,6 +1,7 @@ package checker import ( + "context" "time" "go.uber.org/zap" @@ -45,10 +46,10 @@ func (st *InfoStatusHistory) setStatus(status event.Status) { } } -func (ch *Checker) CheckInfoEvents() error { +func (ch *Checker) CheckInfoEvents(ctx context.Context) error { ch.log.Info("check info event statuses") - infos, err := ch.db.GetInfoEvents() + infos, err := ch.db.GetInfoEvents(ctx) if err != nil { return err } @@ -72,7 +73,7 @@ func (ch *Checker) CheckInfoEvents() error { // Only update the incident if the status has actually changed if info.Status != actualStatus { info.Status = actualStatus - err = ch.db.ModifyIncident(info) + err = ch.db.ModifyIncident(ctx, info) if err != nil { return err } diff --git a/internal/checker/maintenance.go b/internal/checker/maintenance.go index 61f0630..bd10560 100644 --- a/internal/checker/maintenance.go +++ b/internal/checker/maintenance.go @@ -53,10 +53,10 @@ func (st *MntStatusHistory) setStatus(status event.Status) { } } -func (ch *Checker) CheckMaintenance() error { +func (ch *Checker) CheckMaintenance(ctx context.Context) error { ch.log.Info("check maintenances statuses") - maintenances, err := ch.db.GetMaintenances() + maintenances, err := ch.db.GetMaintenances(ctx) if err != nil { return err } @@ -68,7 +68,7 @@ func (ch *Checker) CheckMaintenance() error { continue } - if processErr := ch.processMaintenance(mn); processErr != nil { + if processErr := ch.processMaintenance(ctx, mn); processErr != nil { ch.log.Error("failed to process maintenance", zap.Uint("mntID", mn.ID), zap.Error(processErr)) continue @@ -88,7 +88,7 @@ func needsRefetch(mn *db.Incident) bool { return mn.Status != calculateCurrentMntStatus(calculateMntStatusHistory(mn), mn) } -func (ch *Checker) processMaintenance(mn *db.Incident) error { +func (ch *Checker) processMaintenance(ctx context.Context, mn *db.Incident) error { // Decide from the batch-loaded state whether the status will change. Only // then refetch: a fresh read immediately before the read-modify-write // shrinks the version-conflict window, and the write is the only place the @@ -98,7 +98,7 @@ func (ch *Checker) processMaintenance(mn *db.Incident) error { return nil } - fresh, err := ch.db.GetIncident(int(mn.ID)) + fresh, err := ch.db.GetIncident(ctx, int(mn.ID)) if err != nil { return fmt.Errorf("refetch maintenance %d: %w", mn.ID, err) } @@ -112,11 +112,11 @@ func (ch *Checker) processMaintenance(mn *db.Incident) error { fresh.Status = actualStatus // The modify + enqueue share one transaction: on a version conflict the // whole thing rolls back and no notification is published. - txErr := ch.db.WithTx(context.Background(), func(tx *db.Tx) error { - if modErr := ch.db.ModifyIncidentTx(tx, fresh); modErr != nil { + txErr := ch.db.WithTx(ctx, func(tx *db.Tx) error { + if modErr := ch.db.ModifyIncidentTx(ctx, tx, fresh); modErr != nil { return modErr } - return ch.notifier.PublishTx(context.Background(), tx, notification.Change{ + return ch.notifier.PublishTx(ctx, tx, notification.Change{ IncidentID: fresh.ID, Title: strDeref(fresh.Text), OldStatus: oldStatus, diff --git a/internal/db/db.go b/internal/db/db.go index 268cdf6..88ad52f 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -185,14 +185,12 @@ func noPublicStatus() predicate.Incident { } // GetEventsWithCount retrieves events based on the provided parameters, with pagination and total count. -func (db *DB) GetEventsWithCount(isAuth bool, params ...*IncidentsParams) ([]*Incident, int64, error) { +func (db *DB) GetEventsWithCount(ctx context.Context, isAuth bool, params ...*IncidentsParams) ([]*Incident, int64, error) { var param IncidentsParams if len(params) > 0 && params[0] != nil { param = *params[0] } - ctx := context.Background() - preds, err := applyEventsFilters(¶m, isAuth) if err != nil { return nil, 0, err @@ -244,19 +242,17 @@ func (db *DB) GetEventsWithCount(isAuth bool, params ...*IncidentsParams) ([]*In // GetEvents retrieves events based on the provided parameters. // This is a wrapper around GetEventsWithCount for backward compatibility. -func (db *DB) GetEvents(isAuth bool, params ...*IncidentsParams) ([]*Incident, error) { - events, _, err := db.GetEventsWithCount(isAuth, params...) +func (db *DB) GetEvents(ctx context.Context, isAuth bool, params ...*IncidentsParams) ([]*Incident, error) { + events, _, err := db.GetEventsWithCount(ctx, isAuth, params...) return events, err } // GetEventsInternal retrieves all events for internal use (no filtering by auth). -func (db *DB) GetEventsInternal(params ...*IncidentsParams) ([]*Incident, error) { - return db.GetEvents(AuthorizedAccess, params...) +func (db *DB) GetEventsInternal(ctx context.Context, params ...*IncidentsParams) ([]*Incident, error) { + return db.GetEvents(ctx, AuthorizedAccess, params...) } -func (db *DB) GetIncident(id int) (*Incident, error) { - ctx := context.Background() - +func (db *DB) GetIncident(ctx context.Context, id int) (*Incident, error) { e, err := db.e.Incident.Query(). Where(incident.IDEQ(id)). WithComponents(func(q *ent.ComponentQuery) { @@ -297,12 +293,11 @@ func (db *DB) WithTx(ctx context.Context, fn func(tx *Tx) error) error { } // SaveIncidentTx creates an incident using the provided transaction. -func (db *DB) SaveIncidentTx(tx *Tx, inc *Incident) (uint, error) { +func (db *DB) SaveIncidentTx(ctx context.Context, tx *Tx, inc *Incident) (uint, error) { if inc.Text == nil || *inc.Text == "" { return 0, ErrIncidentTextRequired } - ctx := context.Background() c := db.clientFor(tx) now := time.Now().UTC() @@ -379,19 +374,19 @@ func (db *DB) SaveIncidentTx(tx *Tx, inc *Incident) (uint, error) { return inc.ID, nil } -func (db *DB) SaveIncident(inc *Incident) (uint, error) { - return db.SaveIncidentTx(nil, inc) +func (db *DB) SaveIncident(ctx context.Context, inc *Incident) (uint, error) { + return db.SaveIncidentTx(ctx, nil, inc) } // ModifyIncidentTx applies a modification (with maintenance optimistic locking and // new status inserts) using the provided transaction. -func (db *DB) ModifyIncidentTx(tx *Tx, inc *Incident) error { - return db.modifyIncident(context.Background(), db.clientFor(tx), inc) +func (db *DB) ModifyIncidentTx(ctx context.Context, tx *Tx, inc *Incident) error { + return db.modifyIncident(ctx, db.clientFor(tx), inc) } -func (db *DB) ModifyIncident(inc *Incident) error { - return db.execWithTx(context.Background(), nil, func(client *ent.Client, _ entsql.ExecQuerier) error { - return db.modifyIncident(context.Background(), client, inc) +func (db *DB) ModifyIncident(ctx context.Context, inc *Incident) error { + return db.execWithTx(ctx, nil, func(client *ent.Client, _ entsql.ExecQuerier) error { + return db.modifyIncident(ctx, client, inc) }) } @@ -475,14 +470,13 @@ func applyIncidentPatch(update *ent.IncidentUpdate, inc *Incident) { } // AddComponentToIncident adds a component and a status update to an incident using optimistic locking. -func (db *DB) AddComponentToIncident(inc *Incident, comp *Component, status IncidentStatus) error { +func (db *DB) AddComponentToIncident(ctx context.Context, inc *Incident, comp *Component, status IncidentStatus) error { if inc.Version == nil { return errors.New("version is required for incident modification") } expectedVersion := *inc.Version newVersion := expectedVersion + 1 - ctx := context.Background() err := db.execWithTx(ctx, nil, func(client *ent.Client, _ entsql.ExecQuerier) error { affected, err := client.Incident.Update(). @@ -517,10 +511,10 @@ func (db *DB) AddComponentToIncident(inc *Incident, comp *Component, status Inci } // ReOpenIncident the special function if you need to NULL your end_date. -func (db *DB) ReOpenIncident(inc *Incident) error { +func (db *DB) ReOpenIncident(ctx context.Context, inc *Incident) error { err := db.e.Incident.UpdateOneID(int(inc.ID)). ClearEndDate(). - Exec(context.Background()) + Exec(ctx) if ent.IsNotFound(err) { return nil } @@ -532,14 +526,12 @@ func (db *DB) ReOpenIncident(inc *Incident) error { // Not affected to getActiveEventsForComponent (v2.go) because IsActive filter already contains // exceptions for "event.TypeMaintenance, event.MaintenancePendingReview, event.MaintenanceReviewed". // Supports optional filtering parameters: isActive, Types, LastCount. -func (db *DB) GetEventsByComponentID(componentID uint, params ...*IncidentsParams) ([]*Incident, error) { +func (db *DB) GetEventsByComponentID(ctx context.Context, componentID uint, params ...*IncidentsParams) ([]*Incident, error) { var param IncidentsParams if params != nil && params[0] != nil { param = *params[0] } - ctx := context.Background() - preds := []predicate.Incident{ incident.HasComponentsWith(component.IDEQ(int(componentID))), } @@ -603,7 +595,7 @@ func (db *DB) GetEventsByComponentID(componentID uint, params ...*IncidentsParam return incidents, nil } -func (db *DB) GetIncidentsByComponentAttr(attr *ComponentAttr, params ...*IncidentsParams) ([]*Incident, error) { +func (db *DB) GetIncidentsByComponentAttr(ctx context.Context, attr *ComponentAttr, params ...*IncidentsParams) ([]*Incident, error) { // Get all public incidents for components with this attribute. // Maintenance events in pending_review/reviewed status are excluded (require authentication). var param IncidentsParams @@ -611,8 +603,6 @@ func (db *DB) GetIncidentsByComponentAttr(attr *ComponentAttr, params ...*Incide param = *params[0] } - ctx := context.Background() - // The previous ORM joined the relation and the attribute tables directly, so an // incident matched by several components appeared once per match. The raw id // query keeps that shape and its ordering. @@ -637,11 +627,11 @@ func (db *DB) GetIncidentsByComponentAttr(attr *ComponentAttr, params ...*Incide return db.incidentsByIDs(ctx, ids) } -func (db *DB) GetComponent(id int) (*Component, error) { +func (db *DB) GetComponent(ctx context.Context, id int) (*Component, error) { e, err := db.e.Component.Query(). Where(component.IDEQ(id)). WithAttributes(). - First(context.Background()) + First(ctx) if err != nil { if ent.IsNotFound(err) { return nil, ErrDBComponentDSNotExist @@ -653,8 +643,8 @@ func (db *DB) GetComponent(id int) (*Component, error) { return &comp, nil } -func (db *DB) GetComponentsAsMap() (map[int]*Component, error) { - rows, err := db.e.Component.Query().All(context.Background()) +func (db *DB) GetComponentsAsMap(ctx context.Context) (map[int]*Component, error) { + rows, err := db.e.Component.Query().All(ctx) if err != nil { return nil, err } @@ -668,10 +658,10 @@ func (db *DB) GetComponentsAsMap() (map[int]*Component, error) { return compMap, nil } -func (db *DB) GetComponentsWithValues() ([]Component, error) { +func (db *DB) GetComponentsWithValues(ctx context.Context) ([]Component, error) { rows, err := db.e.Component.Query(). WithAttributes(). - All(context.Background()) + All(ctx) if err != nil { return nil, err } @@ -684,9 +674,7 @@ func (db *DB) GetComponentsWithValues() ([]Component, error) { return components, nil } -func (db *DB) GetComponentsWithIncidents() ([]Component, error) { - ctx := context.Background() - +func (db *DB) GetComponentsWithIncidents(ctx context.Context) ([]Component, error) { rows, err := db.e.Component.Query(). WithAttributes(). WithIncidents(). @@ -716,7 +704,7 @@ func (db *DB) GetComponentsWithIncidents() ([]Component, error) { } // GetComponentFromNameAttrs returns the Component from its name and region attribute. -func (db *DB) GetComponentFromNameAttrs(name string, attr *ComponentAttr) (*Component, error) { +func (db *DB) GetComponentFromNameAttrs(ctx context.Context, name string, attr *ComponentAttr) (*Component, error) { e, err := db.e.Component.Query(). Where( component.NameEQ(name), @@ -724,7 +712,7 @@ func (db *DB) GetComponentFromNameAttrs(name string, attr *ComponentAttr) (*Comp ). WithAttributes(). Order(component.ByID(entsql.OrderAsc())). - First(context.Background()) + First(ctx) if err != nil { if ent.IsNotFound(err) { return nil, ErrDBComponentDSNotExist @@ -736,9 +724,7 @@ func (db *DB) GetComponentFromNameAttrs(name string, attr *ComponentAttr) (*Comp return &comp, nil } -func (db *DB) SaveComponent(comp *Component) (uint, error) { - ctx := context.Background() - +func (db *DB) SaveComponent(ctx context.Context, comp *Component) (uint, error) { // Validate required region attribute hasRegion := false for _, attr := range comp.Attrs { @@ -815,12 +801,12 @@ func (db *DB) SaveComponent(comp *Component) (uint, error) { } func (db *DB) MoveComponentFromOldToAnotherIncident( - comp *Component, incOld, incNew *Incident, closeOld bool, + ctx context.Context, comp *Component, incOld, incNew *Incident, closeOld bool, ) (*Incident, error) { timeNow := time.Now().UTC() if comp.Name == "" { - c, err := db.GetComponent(int(comp.ID)) + c, err := db.GetComponent(ctx, int(comp.ID)) if err != nil { return nil, err } @@ -855,18 +841,18 @@ func (db *DB) MoveComponentFromOldToAnotherIncident( incOld.EndDate = &timeNow } - err := db.execWithTx(context.Background(), nil, func(client *ent.Client, _ entsql.ExecQuerier) error { + err := db.execWithTx(ctx, nil, func(client *ent.Client, _ entsql.ExecQuerier) error { if !closeOld && comp.ID != 0 { - if errRemove := removeIncidentComponent(context.Background(), client, incOld.ID, comp.ID); errRemove != nil { + if errRemove := removeIncidentComponent(ctx, client, incOld.ID, comp.ID); errRemove != nil { return errRemove } dropIncidentComponent(incOld, comp.ID) } - if errSave := saveIncidentFull(context.Background(), client, incNew); errSave != nil { + if errSave := saveIncidentFull(ctx, client, incNew); errSave != nil { return errSave } - return saveIncidentFull(context.Background(), client, incOld) + return saveIncidentFull(ctx, client, incOld) }) if err != nil { return nil, err @@ -876,7 +862,7 @@ func (db *DB) MoveComponentFromOldToAnotherIncident( } func (db *DB) ExtractComponentsToNewIncident( - comp []Component, incOld *Incident, impact int, text string, description *string, + ctx context.Context, comp []Component, incOld *Incident, impact int, text string, description *string, ) (*Incident, error) { if len(comp) == 0 { return nil, fmt.Errorf("no components to extract") @@ -897,7 +883,7 @@ func (db *DB) ExtractComponentsToNewIncident( Components: comp, } - id, err := db.SaveIncident(inc) + id, err := db.SaveIncident(ctx, inc) if err != nil { return nil, err } @@ -923,22 +909,22 @@ func (db *DB) ExtractComponentsToNewIncident( } // Use a transaction to save both incidents with their statuses and update associations - err = db.execWithTx(context.Background(), nil, func(client *ent.Client, _ entsql.ExecQuerier) error { + err = db.execWithTx(ctx, nil, func(client *ent.Client, _ entsql.ExecQuerier) error { // Remove component from old incident for i := range comp { if comp[i].ID == 0 { continue } - if errRemove := removeIncidentComponent(context.Background(), client, incOld.ID, comp[i].ID); errRemove != nil { + if errRemove := removeIncidentComponent(ctx, client, incOld.ID, comp[i].ID); errRemove != nil { return errRemove } dropIncidentComponent(incOld, comp[i].ID) } - if errSave := saveIncidentFull(context.Background(), client, inc); errSave != nil { + if errSave := saveIncidentFull(ctx, client, inc); errSave != nil { return errSave } - return saveIncidentFull(context.Background(), client, incOld) + return saveIncidentFull(ctx, client, incOld) }) if err != nil { return nil, err @@ -947,7 +933,7 @@ func (db *DB) ExtractComponentsToNewIncident( return inc, nil } -func (db *DB) IncreaseIncidentImpact(inc *Incident, impact int) (*Incident, error) { +func (db *DB) IncreaseIncidentImpact(ctx context.Context, inc *Incident, impact int) (*Incident, error) { timeNow := time.Now().UTC() text := fmt.Sprintf("impact changed from %d to %d", *inc.Impact, impact) inc.Statuses = append(inc.Statuses, IncidentStatus{ @@ -964,7 +950,6 @@ func (db *DB) IncreaseIncidentImpact(inc *Incident, impact int) (*Incident, erro // Only non-zero fields are written for the incident row, mirroring the // previous struct-based update. incident_status has no Ent edge, so the // appended status row is inserted directly in the same transaction. - ctx := context.Background() tx, err := db.e.Tx(ctx) if err != nil { return nil, err @@ -1032,12 +1017,12 @@ func (db *DB) IncreaseIncidentImpact(inc *Incident, impact int) (*Incident, erro return inc, nil } -func (db *DB) GetUniqueAttributeValues(attrName string) ([]string, error) { +func (db *DB) GetUniqueAttributeValues(ctx context.Context, attrName string) ([]string, error) { rows, err := db.e.ComponentAttr.Query(). Where(componentattr.NameEQ(attrName)). Select(componentattr.FieldValue). Order(componentattr.ByValue(entsql.OrderAsc())). - All(context.Background()) + All(ctx) if err != nil { return nil, err } @@ -1053,11 +1038,11 @@ func (db *DB) GetUniqueAttributeValues(attrName string) ([]string, error) { return values, nil } -func (db *DB) GetEventUpdates(incidentID uint) ([]IncidentStatus, error) { +func (db *DB) GetEventUpdates(ctx context.Context, incidentID uint) ([]IncidentStatus, error) { rows, err := db.e.IncidentStatus.Query(). Where(incidentstatus.IncidentID(int(incidentID))). Order(incidentstatus.ByID(entsql.OrderAsc())). - All(context.Background()) + All(ctx) if err != nil { return nil, err } @@ -1072,8 +1057,7 @@ func (db *DB) GetEventUpdates(incidentID uint) ([]IncidentStatus, error) { // ModifyEventUpdateTx patches an event status update's text using the provided // transaction and returns the updated row. -func (db *DB) ModifyEventUpdateTx(tx *Tx, update IncidentStatus) (IncidentStatus, error) { - ctx := context.Background() +func (db *DB) ModifyEventUpdateTx(ctx context.Context, tx *Tx, update IncidentStatus) (IncidentStatus, error) { c := db.clientFor(tx) now := time.Now().UTC() @@ -1102,6 +1086,6 @@ func (db *DB) ModifyEventUpdateTx(tx *Tx, update IncidentStatus) (IncidentStatus return incidentStatusFromEnt(row), nil } -func (db *DB) ModifyEventUpdate(update IncidentStatus) (IncidentStatus, error) { - return db.ModifyEventUpdateTx(nil, update) +func (db *DB) ModifyEventUpdate(ctx context.Context, update IncidentStatus) (IncidentStatus, error) { + return db.ModifyEventUpdateTx(ctx, nil, update) } diff --git a/internal/db/event_types.go b/internal/db/event_types.go index d574873..d4c8935 100644 --- a/internal/db/event_types.go +++ b/internal/db/event_types.go @@ -9,9 +9,7 @@ import ( ) // getEventsByType lists events of a single type with their update history. -func (db *DB) getEventsByType(eventType incident.Type, order entsql.OrderTermOption) ([]*Incident, error) { - ctx := context.Background() - +func (db *DB) getEventsByType(ctx context.Context, eventType incident.Type, order entsql.OrderTermOption) ([]*Incident, error) { query := db.e.Incident.Query(). Where(incident.TypeEQ(eventType)). Order(incident.ByID(order)) diff --git a/internal/db/info.go b/internal/db/info.go index d05ab53..0ebeb07 100644 --- a/internal/db/info.go +++ b/internal/db/info.go @@ -1,11 +1,13 @@ package db import ( + "context" + entsql "entgo.io/ent/dialect/sql" "github.com/stackmon/otc-status-dashboard/ent/incident" ) -func (db *DB) GetInfoEvents() ([]*Incident, error) { - return db.getEventsByType(incident.TypeInfo, entsql.OrderDesc()) +func (db *DB) GetInfoEvents(ctx context.Context) ([]*Incident, error) { + return db.getEventsByType(ctx, incident.TypeInfo, entsql.OrderDesc()) } diff --git a/internal/db/maintenances.go b/internal/db/maintenances.go index b0ba160..4495f66 100644 --- a/internal/db/maintenances.go +++ b/internal/db/maintenances.go @@ -1,11 +1,13 @@ package db import ( + "context" + entsql "entgo.io/ent/dialect/sql" "github.com/stackmon/otc-status-dashboard/ent/incident" ) -func (db *DB) GetMaintenances() ([]*Incident, error) { - return db.getEventsByType(incident.TypeMaintenance, entsql.OrderAsc()) +func (db *DB) GetMaintenances(ctx context.Context) ([]*Incident, error) { + return db.getEventsByType(ctx, incident.TypeMaintenance, entsql.OrderAsc()) } diff --git a/internal/db/notification_ops.go b/internal/db/notification_ops.go index eab7309..53578a4 100644 --- a/internal/db/notification_ops.go +++ b/internal/db/notification_ops.go @@ -84,10 +84,10 @@ func (db *DB) ListNotificationsByStatus(ctx context.Context, status string, limi // EnsureNotificationSchema reports whether the outbox table exists. Migrations are // applied out of band, so without this check a stale database would let the app start // and only fail on the first maintenance change. -func (db *DB) EnsureNotificationSchema() error { +func (db *DB) EnsureNotificationSchema(ctx context.Context) error { var count int if err := db.sql.QueryRowContext( - context.Background(), + ctx, `SELECT count(*) FROM information_schema.tables WHERE table_schema = CURRENT_SCHEMA() AND table_name = 'notification_outbox' AND table_type = 'BASE TABLE'`, diff --git a/internal/rss/rss.go b/internal/rss/rss.go index 087e579..3058df7 100644 --- a/internal/rss/rss.go +++ b/internal/rss/rss.go @@ -1,6 +1,7 @@ package rss import ( + "context" "fmt" "net/http" "sort" @@ -65,7 +66,7 @@ func HandleRSS(dbInst *db.DB, logger *zap.Logger) gin.HandlerFunc { componentName: componentName, } - events, err := getEvents(dbInst, logger, params, maxEvents) + events, err := getEvents(c.Request.Context(), dbInst, logger, params, maxEvents) if err != nil { if componentName != "" { apiErrors.RaiseStatusNotFoundErr(c, err) @@ -121,7 +122,7 @@ type feedParams struct { componentName string } -func getEvents(dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int) ([]*db.Incident, error) { +func getEvents(ctx context.Context, dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int) ([]*db.Incident, error) { var incidents []*db.Incident var err error @@ -135,13 +136,13 @@ func getEvents(dbInstance *db.DB, log *zap.Logger, params feedParams, maxInciden } var component *db.Component - component, err = dbInstance.GetComponentFromNameAttrs(params.componentName, attr) + component, err = dbInstance.GetComponentFromNameAttrs(ctx, params.componentName, attr) if err != nil { log.Error("failed to get component", zap.Error(err)) return nil, err } - incidents, err = dbInstance.GetEventsByComponentID(component.ID, incParams) + incidents, err = dbInstance.GetEventsByComponentID(ctx, component.ID, incParams) if err != nil { return nil, err } @@ -151,12 +152,12 @@ func getEvents(dbInstance *db.DB, log *zap.Logger, params feedParams, maxInciden Value: params.region, } - incidents, err = dbInstance.GetIncidentsByComponentAttr(attr, incParams) + incidents, err = dbInstance.GetIncidentsByComponentAttr(ctx, attr, incParams) if err != nil { return nil, err } default: - incidents, err = dbInstance.GetEvents(false, incParams) + incidents, err = dbInstance.GetEvents(ctx, false, incParams) if err != nil { return nil, err } diff --git a/tests/checker_notifications_test.go b/tests/checker_notifications_test.go index 3026a5d..06646d3 100644 --- a/tests/checker_notifications_test.go +++ b/tests/checker_notifications_test.go @@ -1,6 +1,7 @@ package tests import ( + "context" "testing" "github.com/stretchr/testify/assert" @@ -58,7 +59,7 @@ func TestChecker_ReviewedToPlanned_EnqueuesStatusChangedToCreator(t *testing.T) chk := newTestChecker(t) - require.NoError(t, chk.CheckMaintenance()) // reviewed -> planned + require.NoError(t, chk.CheckMaintenance(context.Background())) // reviewed -> planned rows := queryOutbox(t, g, "incident_id = $1", eventID) require.Len(t, rows, 1, "one notification per real transition") @@ -79,7 +80,7 @@ func TestChecker_NoTransition_EnqueuesNothing(t *testing.T) { chk := newTestChecker(t) // Planned with a future start date: the checker computes planned again -> no change. - require.NoError(t, chk.CheckMaintenance()) + require.NoError(t, chk.CheckMaintenance(context.Background())) assert.Equal(t, int64(0), outboxCount(t, g, eventID), "no notification without a real transition") } @@ -100,8 +101,8 @@ func TestChecker_SteadyState_SkipsRefetch(t *testing.T) { chk := newTestChecker(t) - require.NoError(t, chk.CheckMaintenance()) - require.NoError(t, chk.CheckMaintenance()) + require.NoError(t, chk.CheckMaintenance(context.Background())) + require.NoError(t, chk.CheckMaintenance(context.Background())) after := getEventOK(t, r, eventID, adminToken) assert.Equal(t, initialVersion, eventVersion(after), "steady-state scan must not bump the version") diff --git a/tests/db_tx_test.go b/tests/db_tx_test.go index a39c957..b02dd7c 100644 --- a/tests/db_tx_test.go +++ b/tests/db_tx_test.go @@ -19,7 +19,7 @@ func TestWithTx_CommitsIncidentAndOutboxAtomically(t *testing.T) { var incID uint err := d.WithTx(ctx, func(tx *db.Tx) error { - id, e := d.SaveIncidentTx(tx, newMaintenanceIncident()) + id, e := d.SaveIncidentTx(ctx, tx, newMaintenanceIncident()) if e != nil { return e } @@ -42,7 +42,7 @@ func TestWithTx_RollsBackBothOnError(t *testing.T) { var incID uint var dedup string err := d.WithTx(ctx, func(tx *db.Tx) error { - id, e := d.SaveIncidentTx(tx, newMaintenanceIncident()) + id, e := d.SaveIncidentTx(ctx, tx, newMaintenanceIncident()) if e != nil { return e } @@ -67,20 +67,20 @@ func TestModifyIncidentTx_SharedTxWithEnqueue(t *testing.T) { d, g := newNotifDB(t) incID := seedIncident(t, d) - inc, err := d.GetIncident(int(incID)) + inc, err := d.GetIncident(ctx, int(incID)) require.NoError(t, err) inc.Status = event.MaintenanceReviewed row := newOutboxRow(incID, "creator@com.com") err = d.WithTx(ctx, func(tx *db.Tx) error { - if e := d.ModifyIncidentTx(tx, inc); e != nil { + if e := d.ModifyIncidentTx(ctx, tx, inc); e != nil { return e } return d.Enqueue(ctx, tx, row) }) require.NoError(t, err) - got, err := d.GetIncident(int(incID)) + got, err := d.GetIncident(ctx, int(incID)) require.NoError(t, err) assert.Equal(t, event.MaintenanceReviewed, got.Status) @@ -97,7 +97,7 @@ func TestModifyEventUpdateTx_UpdatesText(t *testing.T) { var updated db.IncidentStatus err := d.WithTx(context.Background(), func(tx *db.Tx) error { - u, e := d.ModifyEventUpdateTx(tx, db.IncidentStatus{ + u, e := d.ModifyEventUpdateTx(context.Background(), tx, db.IncidentStatus{ ID: statusID, IncidentID: incID, Text: "patched", }) if e != nil { diff --git a/tests/notifications_ops_test.go b/tests/notifications_ops_test.go index 121a21a..4cadf1e 100644 --- a/tests/notifications_ops_test.go +++ b/tests/notifications_ops_test.go @@ -84,7 +84,7 @@ func TestListNotificationsByStatus(t *testing.T) { func TestEnsureNotificationSchema(t *testing.T) { d, _ := newNotifDB(t) - require.NoError(t, d.EnsureNotificationSchema(), "migrations are applied in the test DB") + require.NoError(t, d.EnsureNotificationSchema(context.Background()), "migrations are applied in the test DB") } func TestRedriveFailed_AllAndByID(t *testing.T) { diff --git a/tests/notifications_test.go b/tests/notifications_test.go index 4f1fb2c..025ec03 100644 --- a/tests/notifications_test.go +++ b/tests/notifications_test.go @@ -35,7 +35,7 @@ func seedIncident(t *testing.T, d *db.DB) uint { text := "notif-test maintenance" start := time.Now().UTC() impact := 0 - id, err := d.SaveIncident(&db.Incident{ + id, err := d.SaveIncident(context.Background(), &db.Incident{ Text: &text, StartDate: &start, Impact: &impact, From 86ead53424a80383796ca279e17b67bb29008da5 Mon Sep 17 00:00:00 2001 From: Aloento <11802769+Aloento@users.noreply.github.com> Date: Sun, 4 Oct 2026 22:46:18 +0200 Subject: [PATCH 2/2] Wrap long signatures and correct the Check comment Adding the ctx parameter pushed several signatures past the 120-column limit; wrap them. The Check comment claimed cancellation is only observed before the round starts, but ctx is now threaded into the scan, so state that cancellation is observed throughout the round. --- internal/api/rss/rss.go | 4 +++- internal/api/v2/v2.go | 8 ++++++-- internal/checker/checker.go | 5 +++-- internal/db/db.go | 12 +++++++++--- internal/db/event_types.go | 4 +++- internal/rss/rss.go | 4 +++- 6 files changed, 27 insertions(+), 10 deletions(-) diff --git a/internal/api/rss/rss.go b/internal/api/rss/rss.go index 78f7627..0786a2c 100644 --- a/internal/api/rss/rss.go +++ b/internal/api/rss/rss.go @@ -103,7 +103,9 @@ type feedParams struct { baseURL string } -func getIncidents(ctx context.Context, dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int) ([]*db.Incident, error) { +func getIncidents( + ctx context.Context, dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int, +) ([]*db.Incident, error) { var incidents []*db.Incident var err error diff --git a/internal/api/v2/v2.go b/internal/api/v2/v2.go index 7f8be23..a032e42 100644 --- a/internal/api/v2/v2.go +++ b/internal/api/v2/v2.go @@ -838,7 +838,8 @@ func processComponentMovement( // processComponentInOpenedIncidents processes a single component against all opened incidents. func processComponentInOpenedIncidents( - ctx context.Context, dbInst *db.DB, log *zap.Logger, comp *db.Component, incIn *db.Incident, openedIncidents []*db.Incident, + ctx context.Context, dbInst *db.DB, log *zap.Logger, + comp *db.Component, incIn *db.Incident, openedIncidents []*db.Incident, ) (*ProcessComponentResp, error) { compResult := &ProcessComponentResp{ ComponentID: int(comp.ID), @@ -969,7 +970,10 @@ func validateEventCreationTimes(incData IncidentData) error { return nil } -func createEvent(ctx context.Context, dbInst *db.DB, log *zap.Logger, inc *db.Incident, userID *string, pub *notification.Publisher) error { +func createEvent( + ctx context.Context, dbInst *db.DB, log *zap.Logger, + inc *db.Incident, userID *string, pub *notification.Publisher, +) error { log.Info("start to save an event to the database") err := dbInst.WithTx(ctx, func(tx *db.Tx) error { diff --git a/internal/checker/checker.go b/internal/checker/checker.go index 2d9f9d6..01c8c10 100644 --- a/internal/checker/checker.go +++ b/internal/checker/checker.go @@ -26,8 +26,9 @@ func New(database *db.DB, log *zap.Logger, notifier *notification.Publisher) *Ch // Check runs one full scan and returns the combined error of its two halves. It // is the body of the scheduler's scan task, which holds the advisory lock for the -// whole round. Cancellation is observed only before the round starts, so a -// caller must not close the pool while Check is running. +// whole round. Cancellation is observed throughout the round, so a caller must +// not close the pool while Check is running; a round aborted mid-scan leaves the +// remaining events to the next tick. func (ch *Checker) Check(ctx context.Context) error { if err := ctx.Err(); err != nil { return err diff --git a/internal/db/db.go b/internal/db/db.go index 88ad52f..4ea5a13 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -185,7 +185,9 @@ func noPublicStatus() predicate.Incident { } // GetEventsWithCount retrieves events based on the provided parameters, with pagination and total count. -func (db *DB) GetEventsWithCount(ctx context.Context, isAuth bool, params ...*IncidentsParams) ([]*Incident, int64, error) { +func (db *DB) GetEventsWithCount( + ctx context.Context, isAuth bool, params ...*IncidentsParams, +) ([]*Incident, int64, error) { var param IncidentsParams if len(params) > 0 && params[0] != nil { param = *params[0] @@ -526,7 +528,9 @@ func (db *DB) ReOpenIncident(ctx context.Context, inc *Incident) error { // Not affected to getActiveEventsForComponent (v2.go) because IsActive filter already contains // exceptions for "event.TypeMaintenance, event.MaintenancePendingReview, event.MaintenanceReviewed". // Supports optional filtering parameters: isActive, Types, LastCount. -func (db *DB) GetEventsByComponentID(ctx context.Context, componentID uint, params ...*IncidentsParams) ([]*Incident, error) { +func (db *DB) GetEventsByComponentID( + ctx context.Context, componentID uint, params ...*IncidentsParams, +) ([]*Incident, error) { var param IncidentsParams if params != nil && params[0] != nil { param = *params[0] @@ -595,7 +599,9 @@ func (db *DB) GetEventsByComponentID(ctx context.Context, componentID uint, para return incidents, nil } -func (db *DB) GetIncidentsByComponentAttr(ctx context.Context, attr *ComponentAttr, params ...*IncidentsParams) ([]*Incident, error) { +func (db *DB) GetIncidentsByComponentAttr( + ctx context.Context, attr *ComponentAttr, params ...*IncidentsParams, +) ([]*Incident, error) { // Get all public incidents for components with this attribute. // Maintenance events in pending_review/reviewed status are excluded (require authentication). var param IncidentsParams diff --git a/internal/db/event_types.go b/internal/db/event_types.go index d4c8935..6257b54 100644 --- a/internal/db/event_types.go +++ b/internal/db/event_types.go @@ -9,7 +9,9 @@ import ( ) // getEventsByType lists events of a single type with their update history. -func (db *DB) getEventsByType(ctx context.Context, eventType incident.Type, order entsql.OrderTermOption) ([]*Incident, error) { +func (db *DB) getEventsByType( + ctx context.Context, eventType incident.Type, order entsql.OrderTermOption, +) ([]*Incident, error) { query := db.e.Incident.Query(). Where(incident.TypeEQ(eventType)). Order(incident.ByID(order)) diff --git a/internal/rss/rss.go b/internal/rss/rss.go index 3058df7..82fc132 100644 --- a/internal/rss/rss.go +++ b/internal/rss/rss.go @@ -122,7 +122,9 @@ type feedParams struct { componentName string } -func getEvents(ctx context.Context, dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int) ([]*db.Incident, error) { +func getEvents( + ctx context.Context, dbInstance *db.DB, log *zap.Logger, params feedParams, maxIncidents int, +) ([]*db.Incident, error) { var incidents []*db.Incident var err error