From bb463be895213d3565b6545efa1ff6f4f3d67e76 Mon Sep 17 00:00:00 2001 From: gearnode Date: Sat, 1 Mar 2025 23:09:15 +0100 Subject: [PATCH] Refactor sql row can using collectable rows Signed-off-by: gearnode --- pkg/probo/coredata/control.go | 92 ++++++---------- .../coredata/control_state_transition.go | 36 ++----- pkg/probo/coredata/evidence.go | 67 ++++-------- .../coredata/evidence_state_transition.go | 34 ++---- pkg/probo/coredata/framework.go | 84 ++++++--------- pkg/probo/coredata/organization.go | 35 +++--- pkg/probo/coredata/people.go | 87 ++++++--------- pkg/probo/coredata/state_transition.go | 12 +-- pkg/probo/coredata/task.go | 65 ++++-------- pkg/probo/coredata/task_state_transition.go | 35 ++---- pkg/probo/coredata/vendor.go | 100 +++++++----------- pkg/usrmgr/coredata/session.go | 35 +++--- pkg/usrmgr/coredata/user.go | 57 +++++----- 13 files changed, 262 insertions(+), 477 deletions(-) diff --git a/pkg/probo/coredata/control.go b/pkg/probo/coredata/control.go index 31ed198d6..8498da7cd 100644 --- a/pkg/probo/coredata/control.go +++ b/pkg/probo/coredata/control.go @@ -16,7 +16,6 @@ package coredata import ( "context" - "errors" "fmt" "maps" "time" @@ -30,16 +29,16 @@ import ( type ( Control struct { - ID gid.GID - FrameworkID gid.GID - Category string - Name string - Description string - State ControlState - ContentRef string - CreatedAt time.Time - UpdatedAt time.Time - Version int + ID gid.GID `db:"id"` + FrameworkID gid.GID `db:"framework_id"` + Category string `db:"category"` + Name string `db:"name"` + Description string `db:"description"` + State ControlState `db:"state"` + ContentRef string `db:"content_ref"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` + Version int `db:"version"` } Controls []*Control @@ -57,22 +56,7 @@ func (c Control) CursorKey() page.CursorKey { return page.NewCursorKey(c.ID, c.CreatedAt) } -func (c *Control) scan(r pgx.Row) error { - return r.Scan( - &c.ID, - &c.FrameworkID, - &c.Category, - &c.Name, - &c.Description, - &c.State, - &c.ContentRef, - &c.CreatedAt, - &c.UpdatedAt, - &c.Version, - ) -} - -func (v *Control) LoadByID( +func (c *Control) LoadByID( ctx context.Context, conn pg.Conn, scope *Scope, @@ -119,14 +103,17 @@ LIMIT 1; args := pgx.NamedArgs{"control_id": controlID} maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - c2 := Control{} - if err := c2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query controls: %w", err) } - *v = c2 + control, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Control]) + if err != nil { + return fmt.Errorf("cannot collect controls: %w", err) + } + + *c = control return nil } @@ -222,24 +209,14 @@ WHERE maps.Copy(args, scope.SQLArguments()) maps.Copy(args, cursor.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - controls := Controls{} - for r.Next() { - control := &Control{} - if err := control.scan(r); err != nil { - return err - } - - controls = append(controls, control) + return fmt.Errorf("cannot query controls: %w", err) } - if err := r.Err(); err != nil { - return err + controls, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[Control]) + if err != nil { + return fmt.Errorf("cannot collect controls: %w", err) } *c = controls @@ -306,16 +283,17 @@ RETURNING maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - c2 := Control{} - if err := c2.scan(r); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return ErrConcurrentModification - } - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query controls: %w", err) } - *c = c2 + control, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Control]) + if err != nil { + return fmt.Errorf("cannot collect controls: %w", err) + } + + *c = control + return nil } diff --git a/pkg/probo/coredata/control_state_transition.go b/pkg/probo/coredata/control_state_transition.go index a966b4085..3d32b4d01 100644 --- a/pkg/probo/coredata/control_state_transition.go +++ b/pkg/probo/coredata/control_state_transition.go @@ -29,7 +29,7 @@ type ( ControlStateTransition struct { StateTransition[ControlState] - ControlID gid.GID + ControlID gid.GID `db:"control_id"` } ControlStateTransitions []*ControlStateTransition @@ -39,18 +39,6 @@ func (cst ControlStateTransition) CursorKey() page.CursorKey { return page.NewCursorKey(cst.ID, cst.CreatedAt) } -func (cst *ControlStateTransition) scan(r pgx.Row) error { - return r.Scan( - &cst.ID, - &cst.ControlID, - &cst.FromState, - &cst.ToState, - &cst.Reason, - &cst.CreatedAt, - &cst.UpdatedAt, - ) -} - func (cst ControlStateTransition) Insert( ctx context.Context, conn pg.Conn, @@ -119,27 +107,17 @@ WHERE args := pgx.NamedArgs{"control_id": controlID} maps.Copy(args, scope.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - controlStateTransitions := ControlStateTransitions{} - for r.Next() { - controlStateTransition := &ControlStateTransition{} - if err := controlStateTransition.scan(r); err != nil { - return err - } - - controlStateTransitions = append(controlStateTransitions, controlStateTransition) + return fmt.Errorf("cannot query control state transitions: %w", err) } - if err := r.Err(); err != nil { - return err + controlStateTransitions, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[ControlStateTransition]) + if err != nil { + return fmt.Errorf("cannot collect control state transitions: %w", err) } - *cst = controlStateTransitions + *cst = ControlStateTransitions(controlStateTransitions) return nil } diff --git a/pkg/probo/coredata/evidence.go b/pkg/probo/coredata/evidence.go index 9e1f9d6ee..046a240f7 100644 --- a/pkg/probo/coredata/evidence.go +++ b/pkg/probo/coredata/evidence.go @@ -28,15 +28,15 @@ import ( type ( Evidence struct { - ID gid.GID - TaskID gid.GID - State EvidenceState - ObjectKey string - MimeType string - Size uint64 - Filename string - CreatedAt time.Time - UpdatedAt time.Time + ID gid.GID `db:"id"` + TaskID gid.GID `db:"task_id"` + State EvidenceState `db:"state"` + ObjectKey string `db:"object_key"` + MimeType string `db:"mime_type"` + Size uint64 `db:"size"` + Filename string `db:"filename"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` } Evidences []*Evidence @@ -46,20 +46,6 @@ func (e Evidence) CursorKey() page.CursorKey { return page.NewCursorKey(e.ID, e.CreatedAt) } -func (e *Evidence) scan(r pgx.Row) error { - return r.Scan( - &e.ID, - &e.TaskID, - &e.State, - &e.ObjectKey, - &e.MimeType, - &e.Size, - &e.Filename, - &e.CreatedAt, - &e.UpdatedAt, - ) -} - func (e Evidence) Insert( ctx context.Context, conn pg.Conn, @@ -149,14 +135,17 @@ LIMIT 1; args := pgx.NamedArgs{"evidence_id": evidenceID} maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - e2 := Evidence{} - if err := e2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query evidence: %w", err) } - *e = e2 + evidence, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Evidence]) + if err != nil { + return fmt.Errorf("cannot collect evidence: %w", err) + } + + *e = evidence return nil } @@ -208,24 +197,14 @@ WHERE maps.Copy(args, scope.SQLArguments()) maps.Copy(args, cursor.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - evidences := Evidences{} - for r.Next() { - evidence := &Evidence{} - if err := evidence.scan(r); err != nil { - return err - } - - evidences = append(evidences, evidence) + return fmt.Errorf("cannot query evidence: %w", err) } - if err := r.Err(); err != nil { - return err + evidences, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[Evidence]) + if err != nil { + return fmt.Errorf("cannot collect evidence: %w", err) } *e = evidences diff --git a/pkg/probo/coredata/evidence_state_transition.go b/pkg/probo/coredata/evidence_state_transition.go index a76f7fd87..c8514be4d 100644 --- a/pkg/probo/coredata/evidence_state_transition.go +++ b/pkg/probo/coredata/evidence_state_transition.go @@ -29,7 +29,7 @@ type ( EvidenceStateTransition struct { StateTransition[EvidenceState] - EvidenceID gid.GID + EvidenceID gid.GID `db:"evidence_id"` } EvidenceStateTransitions []*EvidenceStateTransition @@ -39,18 +39,6 @@ func (cst EvidenceStateTransition) CursorKey() page.CursorKey { return page.NewCursorKey(cst.ID, cst.CreatedAt) } -func (cst *EvidenceStateTransition) scan(r pgx.Row) error { - return r.Scan( - &cst.ID, - &cst.EvidenceID, - &cst.FromState, - &cst.ToState, - &cst.Reason, - &cst.CreatedAt, - &cst.UpdatedAt, - ) -} - func (est EvidenceStateTransition) Insert( ctx context.Context, conn pg.Conn, @@ -119,24 +107,14 @@ WHERE args := pgx.NamedArgs{"evidence_id": evidenceID} maps.Copy(args, scope.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - evidenceStateTransitions := EvidenceStateTransitions{} - for r.Next() { - evidenceStateTransition := &EvidenceStateTransition{} - if err := evidenceStateTransition.scan(r); err != nil { - return err - } - - evidenceStateTransitions = append(evidenceStateTransitions, evidenceStateTransition) + return fmt.Errorf("cannot query evidence state transitions: %w", err) } - if err := r.Err(); err != nil { - return err + evidenceStateTransitions, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[EvidenceStateTransition]) + if err != nil { + return fmt.Errorf("cannot collect evidence state transitions: %w", err) } *cst = evidenceStateTransitions diff --git a/pkg/probo/coredata/framework.go b/pkg/probo/coredata/framework.go index b4fad9eb2..d93cf3987 100644 --- a/pkg/probo/coredata/framework.go +++ b/pkg/probo/coredata/framework.go @@ -16,7 +16,6 @@ package coredata import ( "context" - "errors" "fmt" "maps" "time" @@ -29,14 +28,14 @@ import ( type ( Framework struct { - ID gid.GID - OrganizationID gid.GID - Name string - Description string - ContentRef string - CreatedAt time.Time - UpdatedAt time.Time - Version int + ID gid.GID `db:"id"` + OrganizationID gid.GID `db:"organization_id"` + Name string `db:"name"` + Description string `db:"description"` + ContentRef string `db:"content_ref"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` + Version int `db:"version"` } Frameworks []*Framework @@ -52,19 +51,6 @@ func (f Framework) CursorKey() page.CursorKey { return page.NewCursorKey(f.ID, f.CreatedAt) } -func (f *Framework) scan(r pgx.Row) error { - return r.Scan( - &f.ID, - &f.OrganizationID, - &f.Name, - &f.Description, - &f.ContentRef, - &f.CreatedAt, - &f.UpdatedAt, - &f.Version, - ) -} - func (f *Frameworks) LoadByOrganizationID( ctx context.Context, conn pg.Conn, @@ -96,24 +82,14 @@ WHERE maps.Copy(args, cursor.SQLArguments()) maps.Copy(args, cursor.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - frameworks := Frameworks{} - for r.Next() { - framework := &Framework{} - if err := framework.scan(r); err != nil { - return err - } - - frameworks = append(frameworks, framework) + return fmt.Errorf("cannot query frameworks: %w", err) } - if err := r.Err(); err != nil { - return err + frameworks, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[Framework]) + if err != nil { + return fmt.Errorf("cannot collect frameworks: %w", err) } *f = frameworks @@ -149,14 +125,17 @@ LIMIT 1; args := pgx.NamedArgs{"framework_id": frameworkID} maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - f2 := Framework{} - if err := f2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query frameworks: %w", err) } - *f = f2 + framework, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Framework]) + if err != nil { + return fmt.Errorf("cannot collect framework: %w", err) + } + + *f = framework return nil } @@ -267,16 +246,17 @@ RETURNING maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - f2 := Framework{} - if err := f2.scan(r); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return ErrConcurrentModification - } - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query frameworks: %w", err) } - *f = f2 + framework, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Framework]) + if err != nil { + return fmt.Errorf("cannot collect framework: %w", err) + } + + *f = framework + return nil } diff --git a/pkg/probo/coredata/organization.go b/pkg/probo/coredata/organization.go index 65661544a..c6bab2474 100644 --- a/pkg/probo/coredata/organization.go +++ b/pkg/probo/coredata/organization.go @@ -27,24 +27,14 @@ import ( type ( Organization struct { - ID gid.GID - Name string - LogoURL string - CreatedAt time.Time - UpdatedAt time.Time + ID gid.GID `db:"id"` + Name string `db:"name"` + LogoURL string `db:"logo_url"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` } ) -func (o *Organization) scan(r pgx.Row) error { - return r.Scan( - &o.ID, - &o.Name, - &o.LogoURL, - &o.CreatedAt, - &o.UpdatedAt, - ) -} - func (o *Organization) LoadByID( ctx context.Context, conn pg.Conn, @@ -71,14 +61,17 @@ LIMIT 1; args := pgx.NamedArgs{"organization_id": organizationID} maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - o2 := Organization{} - if err := o2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query organizations: %w", err) } - *o = o2 + organization, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Organization]) + if err != nil { + return fmt.Errorf("cannot collect organization: %w", err) + } + + *o = organization return nil } diff --git a/pkg/probo/coredata/people.go b/pkg/probo/coredata/people.go index 291627738..b450b985c 100644 --- a/pkg/probo/coredata/people.go +++ b/pkg/probo/coredata/people.go @@ -16,7 +16,6 @@ package coredata import ( "context" - "errors" "fmt" "maps" "time" @@ -29,15 +28,15 @@ import ( type ( People struct { - ID gid.GID - OrganizationID gid.GID - Kind PeopleKind - FullName string - PrimaryEmailAddress string - AdditionalEmailAddresses []string - CreatedAt time.Time - UpdatedAt time.Time - Version int + ID gid.GID `db:"id"` + OrganizationID gid.GID `db:"organization_id"` + Kind PeopleKind `db:"kind"` + FullName string `db:"full_name"` + PrimaryEmailAddress string `db:"primary_email_address"` + AdditionalEmailAddresses []string `db:"additional_email_addresses"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` + Version int `db:"version"` } Peoples []*People @@ -55,20 +54,6 @@ func (p People) CursorKey() page.CursorKey { return page.NewCursorKey(p.ID, p.CreatedAt) } -func (p *People) scan(r pgx.Row) error { - return r.Scan( - &p.ID, - &p.OrganizationID, - &p.Kind, - &p.FullName, - &p.PrimaryEmailAddress, - &p.AdditionalEmailAddresses, - &p.CreatedAt, - &p.UpdatedAt, - &p.Version, - ) -} - func (p *People) LoadByID( ctx context.Context, conn pg.Conn, @@ -99,14 +84,17 @@ LIMIT 1; args := pgx.NamedArgs{"people_id": peopleID} maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - p2 := People{} - if err := p2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query people: %w", err) } - *p = p2 + people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People]) + if err != nil { + return fmt.Errorf("cannot collect people: %w", err) + } + + *p = people return nil } @@ -206,24 +194,14 @@ WHERE maps.Copy(args, cursor.SQLArguments()) maps.Copy(args, cursor.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - peoples := Peoples{} - for r.Next() { - people := &People{} - if err := people.scan(r); err != nil { - return err - } - - peoples = append(peoples, people) + return fmt.Errorf("cannot query people: %w", err) } - if err := r.Err(); err != nil { - return err + peoples, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[People]) + if err != nil { + return fmt.Errorf("cannot collect people: %w", err) } *p = peoples @@ -282,16 +260,17 @@ RETURNING maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - p2 := People{} - if err := p2.scan(r); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return ErrConcurrentModification - } - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query people: %w", err) } - *p = p2 + people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People]) + if err != nil { + return fmt.Errorf("cannot collect people: %w", err) + } + + *p = people + return nil } diff --git a/pkg/probo/coredata/state_transition.go b/pkg/probo/coredata/state_transition.go index 7310269d9..1317ce1c4 100644 --- a/pkg/probo/coredata/state_transition.go +++ b/pkg/probo/coredata/state_transition.go @@ -21,11 +21,11 @@ import ( type ( StateTransition[T any] struct { - ID gid.GID - FromState *T - ToState T - Reason *string - CreatedAt time.Time - UpdatedAt time.Time + ID gid.GID `db:"id"` + ToState T `db:"to_state"` + FromState *T `db:"from_state"` + Reason *string `db:"reason"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` } ) diff --git a/pkg/probo/coredata/task.go b/pkg/probo/coredata/task.go index bf03897b5..96b26e535 100644 --- a/pkg/probo/coredata/task.go +++ b/pkg/probo/coredata/task.go @@ -29,14 +29,14 @@ import ( type ( Task struct { - ID gid.GID - ControlID gid.GID - Name string - Description string - State TaskState - ContentRef string - CreatedAt time.Time - UpdatedAt time.Time + ID gid.GID `db:"id"` + ControlID gid.GID `db:"control_id"` + Name string `db:"name"` + Description string `db:"description"` + State TaskState `db:"state"` + ContentRef string `db:"content_ref"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` } Tasks []*Task @@ -46,19 +46,6 @@ func (t Task) CursorKey() page.CursorKey { return page.NewCursorKey(t.ID, t.CreatedAt) } -func (t *Task) scan(r pgx.Row) error { - return r.Scan( - &t.ID, - &t.ControlID, - &t.Name, - &t.Description, - &t.State, - &t.ContentRef, - &t.CreatedAt, - &t.UpdatedAt, - ) -} - func (t *Task) LoadByID( ctx context.Context, conn pg.Conn, @@ -122,14 +109,17 @@ LIMIT 1; args := pgx.NamedArgs{"task_id": taskID} maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - t2 := Task{} - if err := t2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query task: %w", err) } - *t = t2 + task, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Task]) + if err != nil { + return fmt.Errorf("cannot collect task: %w", err) + } + + *t = task return nil } @@ -245,24 +235,14 @@ WHERE maps.Copy(args, scope.SQLArguments()) maps.Copy(args, cursor.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - tasks := Tasks{} - for r.Next() { - task := &Task{} - if err := task.scan(r); err != nil { - return err - } - - tasks = append(tasks, task) + return fmt.Errorf("cannot query tasks: %w", err) } - if err := r.Err(); err != nil { - return err + tasks, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[Task]) + if err != nil { + return fmt.Errorf("cannot collect tasks: %w", err) } *t = tasks @@ -275,7 +255,6 @@ func (t *Task) Delete( conn pg.Conn, scope *Scope, ) error { - // Use a single transaction with conditional logic to handle both cases q := ` WITH control_count AS ( SELECT COUNT(*) AS count FROM controls_tasks WHERE task_id = @task_id diff --git a/pkg/probo/coredata/task_state_transition.go b/pkg/probo/coredata/task_state_transition.go index be3f99029..8d2d94eac 100644 --- a/pkg/probo/coredata/task_state_transition.go +++ b/pkg/probo/coredata/task_state_transition.go @@ -27,9 +27,8 @@ import ( type ( TaskStateTransition struct { - TaskID gid.GID - StateTransition[TaskState] + TaskID gid.GID `db:"task_id"` } TaskStateTransitions []*TaskStateTransition @@ -39,18 +38,6 @@ func (tst TaskStateTransition) CursorKey() page.CursorKey { return page.NewCursorKey(tst.ID, tst.CreatedAt) } -func (tst *TaskStateTransition) scan(r pgx.Row) error { - return r.Scan( - &tst.ID, - &tst.TaskID, - &tst.FromState, - &tst.ToState, - &tst.Reason, - &tst.CreatedAt, - &tst.UpdatedAt, - ) -} - func (tst TaskStateTransition) Insert( ctx context.Context, conn pg.Conn, @@ -119,24 +106,14 @@ WHERE args := pgx.NamedArgs{"task_id": taskID} maps.Copy(args, scope.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - taskStateTransitions := TaskStateTransitions{} - for r.Next() { - taskStateTransition := &TaskStateTransition{} - if err := taskStateTransition.scan(r); err != nil { - return err - } - - taskStateTransitions = append(taskStateTransitions, taskStateTransition) + return fmt.Errorf("cannot query task state transitions: %w", err) } - if err := r.Err(); err != nil { - return err + taskStateTransitions, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[TaskStateTransition]) + if err != nil { + return fmt.Errorf("cannot collect task state transitions: %w", err) } *tst = taskStateTransitions diff --git a/pkg/probo/coredata/vendor.go b/pkg/probo/coredata/vendor.go index 9b2ab74e3..7fd547480 100644 --- a/pkg/probo/coredata/vendor.go +++ b/pkg/probo/coredata/vendor.go @@ -31,20 +31,20 @@ var ErrConcurrentModification = errors.New("concurrent modification") type ( Vendor struct { - ID gid.GID - OrganizationID gid.GID - Name string - Description string - ServiceStartAt time.Time - ServiceTerminationAt *time.Time - ServiceCriticality ServiceCriticality - RiskTier RiskTier - StatusPageURL *string - TermsOfServiceURL *string - PrivacyPolicyURL *string - CreatedAt time.Time - UpdatedAt time.Time - Version int + ID gid.GID `db:"id"` + OrganizationID gid.GID `db:"organization_id"` + Name string `db:"name"` + Description string `db:"description"` + ServiceStartAt time.Time `db:"service_start_at"` + ServiceTerminationAt *time.Time `db:"service_termination_at"` + ServiceCriticality ServiceCriticality `db:"service_criticality"` + RiskTier RiskTier `db:"risk_tier"` + StatusPageURL *string `db:"status_page_url"` + TermsOfServiceURL *string `db:"terms_of_service_url"` + PrivacyPolicyURL *string `db:"privacy_policy_url"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` + Version int `db:"version"` } Vendors []*Vendor @@ -67,25 +67,6 @@ func (v Vendor) CursorKey() page.CursorKey { return page.NewCursorKey(v.ID, v.CreatedAt) } -func (v *Vendor) scan(r pgx.Row) error { - return r.Scan( - &v.ID, - &v.OrganizationID, - &v.Name, - &v.Description, - &v.ServiceStartAt, - &v.ServiceTerminationAt, - &v.ServiceCriticality, - &v.RiskTier, - &v.StatusPageURL, - &v.TermsOfServiceURL, - &v.PrivacyPolicyURL, - &v.CreatedAt, - &v.UpdatedAt, - &v.Version, - ) -} - func (v *Vendor) LoadByID( ctx context.Context, conn pg.Conn, @@ -121,14 +102,18 @@ LIMIT 1; args := pgx.NamedArgs{"vendor_id": vendorID} maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query vendor: %w", err) + } + defer rows.Close() - v2 := Vendor{} - if err := v2.scan(r); err != nil { - return err + vendor, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Vendor]) + if err != nil { + return fmt.Errorf("cannot collect vendor: %w", err) } - *v = v2 + *v = vendor return nil } @@ -247,24 +232,14 @@ WHERE maps.Copy(args, cursor.SQLArguments()) maps.Copy(args, scope.SQLArguments()) - r, err := conn.Query(ctx, q, args) + rows, err := conn.Query(ctx, q, args) if err != nil { - return err - } - defer r.Close() - - vendors := Vendors{} - for r.Next() { - vendor := &Vendor{} - if err := vendor.scan(r); err != nil { - return err - } - - vendors = append(vendors, vendor) + return fmt.Errorf("cannot query vendors: %w", err) } - if err := r.Err(); err != nil { - return err + vendors, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[Vendor]) + if err != nil { + return fmt.Errorf("cannot collect vendors: %w", err) } *v = vendors @@ -348,16 +323,17 @@ RETURNING maps.Copy(args, scope.SQLArguments()) - r := conn.QueryRow(ctx, q, args) - - v2 := Vendor{} - if err := v2.scan(r); err != nil { - if errors.Is(err, pgx.ErrNoRows) { - return ErrConcurrentModification - } - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query vendor: %w", err) } - *v = v2 + vendor, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Vendor]) + if err != nil { + return fmt.Errorf("cannot collect vendor: %w", err) + } + + *v = vendor + return nil } diff --git a/pkg/usrmgr/coredata/session.go b/pkg/usrmgr/coredata/session.go index 85ae59d7c..852a7aee5 100644 --- a/pkg/usrmgr/coredata/session.go +++ b/pkg/usrmgr/coredata/session.go @@ -16,6 +16,7 @@ package coredata import ( "context" + "fmt" "time" "github.com/getprobo/probo/pkg/gid" @@ -26,11 +27,11 @@ import ( type ( Session struct { - ID gid.GID - UserID gid.GID - ExpiredAt time.Time - CreatedAt time.Time - UpdatedAt time.Time + ID gid.GID `db:"id"` + UserID gid.GID `db:"user_id"` + ExpiredAt time.Time `db:"expired_at"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` } ) @@ -38,16 +39,6 @@ func (s Session) CursorKey() page.CursorKey { return page.NewCursorKey(s.ID, s.CreatedAt) } -func (s *Session) scan(r pgx.Row) error { - return r.Scan( - &s.ID, - &s.UserID, - &s.ExpiredAt, - &s.CreatedAt, - &s.UpdatedAt, - ) -} - func (s *Session) LoadByID( ctx context.Context, conn pg.Conn, @@ -69,14 +60,16 @@ LIMIT 1; args := pgx.NamedArgs{"session_id": sessionID} - r := conn.QueryRow(ctx, q, args) - - s2 := Session{} - if err := s2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query session: %w", err) } - *s = s2 + session, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Session]) + if err != nil { + return fmt.Errorf("cannot collect session: %w", err) + } + *s = session return nil } diff --git a/pkg/usrmgr/coredata/user.go b/pkg/usrmgr/coredata/user.go index 5df80c74a..12bb5a21f 100644 --- a/pkg/usrmgr/coredata/user.go +++ b/pkg/usrmgr/coredata/user.go @@ -16,6 +16,7 @@ package coredata import ( "context" + "fmt" "time" "github.com/getprobo/probo/pkg/gid" @@ -26,13 +27,13 @@ import ( type ( User struct { - ID gid.GID - EmailAddress string - HashedPassword []byte - FullName string - OrganizationID gid.GID - CreatedAt time.Time - UpdatedAt time.Time + ID gid.GID `db:"id"` + EmailAddress string `db:"email_address"` + HashedPassword []byte `db:"hashed_password"` + FullName string `db:"fullname"` + OrganizationID gid.GID `db:"organization_id"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` } ) @@ -40,18 +41,6 @@ func (u User) CursorKey() page.CursorKey { return page.NewCursorKey(u.ID, u.CreatedAt) } -func (u *User) scan(r pgx.Row) error { - return r.Scan( - &u.ID, - &u.EmailAddress, - &u.HashedPassword, - &u.FullName, - &u.OrganizationID, - &u.CreatedAt, - &u.UpdatedAt, - ) -} - func (u *User) LoadByEmail( ctx context.Context, conn pg.Conn, @@ -75,14 +64,17 @@ LIMIT 1; args := pgx.NamedArgs{"user_email": email} - r := conn.QueryRow(ctx, q, args) - - u2 := User{} - if err := u2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query user: %w", err) } - *u = u2 + user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User]) + if err != nil { + return fmt.Errorf("cannot collect user: %w", err) + } + + *u = user return nil } @@ -110,14 +102,17 @@ LIMIT 1; args := pgx.NamedArgs{"user_id": userID} - r := conn.QueryRow(ctx, q, args) - - u2 := User{} - if err := u2.scan(r); err != nil { - return err + rows, err := conn.Query(ctx, q, args) + if err != nil { + return fmt.Errorf("cannot query user: %w", err) } - *u = u2 + user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User]) + if err != nil { + return fmt.Errorf("cannot collect user: %w", err) + } + + *u = user return nil }