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..0786a2c 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,9 @@ 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 +119,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 +135,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..a032e42 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,8 @@ 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 +854,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 +878,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 +891,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 +970,14 @@ 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 +1024,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 +1082,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 +1189,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 +1211,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 +1418,7 @@ func PostIncidentExtractHandler(dbInst *db.DB, logger *zap.Logger) gin.HandlerFu } inc, err := dbInst.ExtractComponentsToNewIncident( + c.Request.Context(), movedComponents, storedInc, *storedInc.Impact, @@ -1469,7 +1477,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 +1497,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 +1544,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 +1600,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 +1829,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 +1844,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..01c8c10 100644 --- a/internal/checker/checker.go +++ b/internal/checker/checker.go @@ -26,9 +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: 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 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 @@ -43,7 +43,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 +52,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..4ea5a13 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -185,14 +185,14 @@ 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 +244,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 +295,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 +376,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 +472,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 +513,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 +528,14 @@ 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 +599,9 @@ 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 +609,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 +633,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 +649,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 +664,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 +680,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 +710,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 +718,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 +730,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 +807,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 +847,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 +868,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 +889,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 +915,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 +939,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 +956,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 +1023,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 +1044,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 +1063,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 +1092,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..6257b54 100644 --- a/internal/db/event_types.go +++ b/internal/db/event_types.go @@ -9,9 +9,9 @@ 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..82fc132 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,9 @@ 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 +138,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 +154,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,