Skip to content

Commit a6fc8c3

Browse files
committed
Merge branch 'fix/784-correctness-fixes' into 'master'
fix: snapshot exemptions, observation stop hang, start timeout, token reload, panic guards (#784) Closes #29 and #784 See merge request postgres-ai/database-lab!1199
2 parents f03b982 + d116150 commit a6fc8c3

16 files changed

Lines changed: 439 additions & 49 deletions

File tree

‎engine/cmd/cli/commands/clone/actions.go‎

Lines changed: 21 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -132,7 +132,12 @@ func create(cliCtx *cli.Context) error {
132132
cloneRequest.Snapshot = &types.SnapshotCloneFieldRequest{ID: cliCtx.String("snapshot-id")}
133133
}
134134

135-
cloneRequest.ExtraConf = splitFlags(cliCtx.StringSlice("extra-config"))
135+
extraConf, err := splitFlags(cliCtx.StringSlice("extra-config"))
136+
if err != nil {
137+
return fmt.Errorf("invalid --extra-config value: %w", err)
138+
}
139+
140+
cloneRequest.ExtraConf = extraConf
136141

137142
var clone *models.Clone
138143

@@ -373,10 +378,15 @@ func startObservation(cliCtx *cli.Context) error {
373378
MaxDuration: cliCtx.Uint64("max-duration"),
374379
}
375380

381+
tags, err := splitFlags(cliCtx.StringSlice("tags"))
382+
if err != nil {
383+
return fmt.Errorf("invalid --tags value: %w", err)
384+
}
385+
376386
start := types.StartObservationRequest{
377387
CloneID: cloneID,
378388
Config: observationConfig,
379-
Tags: splitFlags(cliCtx.StringSlice("tags")),
389+
Tags: tags,
380390
DBName: cliCtx.String("db-name"),
381391
}
382392

@@ -569,19 +579,18 @@ func retrieveClonePort(cliCtx *cli.Context, wg *sync.WaitGroup, remoteHost *url.
569579
return clone.DB.Port, nil
570580
}
571581

572-
func splitFlags(flags []string) map[string]string {
573-
const maxSplitParts = 2
574-
582+
// splitFlags parses repeated key=value flag entries into a map; an entry without "=" is an error.
583+
func splitFlags(flags []string) (map[string]string, error) {
575584
extraConfig := make(map[string]string, len(flags))
576585

577-
if len(flags) == 0 {
578-
return extraConfig
579-
}
580-
581586
for _, cfg := range flags {
582-
parsed := strings.SplitN(cfg, "=", maxSplitParts)
583-
extraConfig[parsed[0]] = parsed[1]
587+
key, value, found := strings.Cut(cfg, "=")
588+
if !found || key == "" {
589+
return nil, fmt.Errorf("%q is not in the key=value form", cfg)
590+
}
591+
592+
extraConfig[key] = value
584593
}
585594

586-
return extraConfig
595+
return extraConfig, nil
587596
}

‎engine/cmd/cli/commands/clone/actions_test.go‎

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,3 +41,36 @@ func TestUpgradeCommandIsRegistered(t *testing.T) {
4141
assert.True(t, flagNames["async"])
4242
assert.False(t, flagNames["target-version"], "the target follows from the instance configuration, not from a flag")
4343
}
44+
45+
func TestSplitFlags(t *testing.T) {
46+
testCases := []struct {
47+
name string
48+
flags []string
49+
want map[string]string
50+
wantErr bool
51+
}{
52+
{name: "empty", flags: nil, want: map[string]string{}},
53+
{name: "single pair", flags: []string{"shared_buffers=1GB"}, want: map[string]string{"shared_buffers": "1GB"}},
54+
{name: "value keeps extra equals", flags: []string{"search_path=a=b"}, want: map[string]string{"search_path": "a=b"}},
55+
{name: "empty value", flags: []string{"work_mem="}, want: map[string]string{"work_mem": ""}},
56+
{name: "several pairs", flags: []string{"a=1", "b=2"}, want: map[string]string{"a": "1", "b": "2"}},
57+
{name: "missing equals", flags: []string{"foo"}, wantErr: true},
58+
{name: "missing key", flags: []string{"=bar"}, wantErr: true},
59+
{name: "bad entry among good", flags: []string{"a=1", "foo"}, wantErr: true},
60+
}
61+
62+
for _, tc := range testCases {
63+
t.Run(tc.name, func(t *testing.T) {
64+
got, err := splitFlags(tc.flags)
65+
if tc.wantErr {
66+
assert.Error(t, err)
67+
assert.Nil(t, got)
68+
69+
return
70+
}
71+
72+
assert.NoError(t, err)
73+
assert.Equal(t, tc.want, got)
74+
})
75+
}
76+
}

‎engine/internal/observer/observer_test.go‎

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
package observer
22

33
import (
4+
"context"
45
"regexp"
56
"sync"
67
"testing"
@@ -9,7 +10,9 @@ import (
910
"github.com/stretchr/testify/assert"
1011
"github.com/stretchr/testify/require"
1112

13+
"gitlab.com/postgres-ai/database-lab/v3/internal/provision/resources"
1214
"gitlab.com/postgres-ai/database-lab/v3/pkg/client/dblabapi/types"
15+
"gitlab.com/postgres-ai/database-lab/v3/pkg/models"
1316
)
1417

1518
func TestMaskingField(t *testing.T) {
@@ -219,3 +222,47 @@ func TestObservingClone_SetOverallError(t *testing.T) {
219222
oc.SetOverallError(false)
220223
assert.False(t, oc.session.state.OverallError)
221224
}
225+
226+
func TestObservingClone_RunSessionErrorSignalsDone(t *testing.T) {
227+
oc := NewObservingClone(types.Config{}, nil)
228+
229+
require.Error(t, oc.RunSession(), "a session that was never initialized cannot run")
230+
231+
select {
232+
case <-oc.done:
233+
case <-time.After(time.Second):
234+
t.Fatal("done must be signalled when RunSession exits with an error")
235+
}
236+
}
237+
238+
func TestObservingClone_StopExpiredContext(t *testing.T) {
239+
oc := NewObservingClone(types.Config{}, nil)
240+
oc.session = &Session{SessionID: 42}
241+
242+
ctx, cancel := context.WithCancel(context.Background())
243+
cancel()
244+
245+
err := oc.Stop(ctx)
246+
247+
require.Error(t, err, "Stop must not block when RunSession never started")
248+
assert.ErrorIs(t, err, context.Canceled)
249+
}
250+
251+
func TestObservingClone_StopUninitializedSession(t *testing.T) {
252+
oc := NewObservingClone(types.Config{}, nil)
253+
254+
require.Error(t, oc.Stop(context.Background()))
255+
}
256+
257+
func TestObservingClone_StopTwiceAfterRunSessionExit(t *testing.T) {
258+
oc := NewObservingClone(types.Config{}, nil)
259+
oc.pool = &resources.Pool{Name: "pool", MountDir: t.TempDir()}
260+
oc.session = &Session{SessionID: 7, Config: types.Config{MaxDuration: 60}, Result: &models.ObservationResult{}}
261+
262+
close(oc.done)
263+
264+
require.NotPanics(t, func() {
265+
require.NoError(t, oc.Stop(context.Background()))
266+
require.NoError(t, oc.Stop(context.Background()))
267+
})
268+
}

‎engine/internal/observer/observing_clone.go‎

Lines changed: 15 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -97,7 +97,7 @@ func NewObservingClone(config types.Config, sudb *pgx.Conn) *ObservingClone {
9797
config: config,
9898
ctx: ctx,
9999
cancel: cancel,
100-
done: make(chan struct{}, 1),
100+
done: make(chan struct{}),
101101
csvFields: csvFields,
102102
registryMu: &sync.Mutex{},
103103
sessionRegistry: make(map[uint64]struct{}),
@@ -202,8 +202,11 @@ func (c *ObservingClone) Init(clone *models.Clone, sessionID uint64, startedAt t
202202
return nil
203203
}
204204

205-
// RunSession runs observing session.
205+
// RunSession runs observing session. It closes the done channel on every return path so that
206+
// Stop never waits for a session that has already ended.
206207
func (c *ObservingClone) RunSession() error {
208+
defer close(c.done)
209+
207210
if c.session == nil || c.db.IsClosed() {
208211
return errors.New("failed to run session because it has not been initialized")
209212
}
@@ -259,8 +262,6 @@ func (c *ObservingClone) RunSession() error {
259262
log.Err("failed to store artifacts: ", err)
260263
}
261264

262-
c.done <- struct{}{}
263-
264265
return nil
265266
}
266267

@@ -436,16 +437,20 @@ where table_name = 'postgres_log'`)
436437
}
437438

438439
// Stop stops an observation session.
439-
func (c *ObservingClone) Stop() error {
440+
func (c *ObservingClone) Stop(ctx context.Context) error {
441+
if c.session == nil {
442+
return errors.New("failed to summarize session because it has not been initialized")
443+
}
444+
440445
log.Msg(fmt.Sprintf("Observation session %v is stopping...", c.session.SessionID))
441446

442447
c.cancel()
443448

444-
// Waiting for the observation process stops.
445-
<-c.done
446-
447-
if c.session == nil {
448-
return errors.New("failed to summarize session because it has not been initialized")
449+
// Wait until the observation process ends or the caller gives up.
450+
select {
451+
case <-c.done:
452+
case <-ctx.Done():
453+
return fmt.Errorf("observation session %v has not stopped: %w", c.session.SessionID, ctx.Err())
449454
}
450455

451456
c.summarize()

‎engine/internal/provision/databases/postgres/postgres.go‎

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -119,7 +119,7 @@ func Start(r runners.Runner, c *resources.AppConfig) error {
119119
log.Err(runnerErr)
120120
}
121121

122-
return errors.Wrap(err, "postgres start timeout")
122+
return startTimeoutError(out, err)
123123
}
124124

125125
time.Sleep(checkPostgresStatusPeriod * time.Millisecond)
@@ -128,6 +128,16 @@ func Start(r runners.Runner, c *resources.AppConfig) error {
128128
return nil
129129
}
130130

131+
// startTimeoutError describes why Postgres was not ready before the deadline: recoveryState is
132+
// the last pg_is_in_recovery() result and queryErr the error of the last status query, if any.
133+
func startTimeoutError(recoveryState string, queryErr error) error {
134+
if queryErr != nil {
135+
return fmt.Errorf("postgres start timeout: last status check failed: %w", queryErr)
136+
}
137+
138+
return fmt.Errorf("postgres start timeout: instance is still in recovery (pg_is_in_recovery: %q)", recoveryState)
139+
}
140+
131141
func collectDiagnostics(c *resources.AppConfig) {
132142
dockerClient, err := client.New(client.FromEnv)
133143
if err != nil {

‎engine/internal/provision/databases/postgres/postgres_test.go‎

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -170,3 +170,32 @@ func TestRemoveContainers(t *testing.T) {
170170
assert.Equal(t, tc.err, errors.Cause(err))
171171
}
172172
}
173+
174+
func TestStartTimeoutError(t *testing.T) {
175+
queryErr := errors.New("connection refused")
176+
177+
testCases := []struct {
178+
name string
179+
recoveryState string
180+
queryErr error
181+
wantContains string
182+
}{
183+
{name: "still in recovery", recoveryState: "t", queryErr: nil, wantContains: `still in recovery (pg_is_in_recovery: "t")`},
184+
{name: "no status at all", recoveryState: "", queryErr: nil, wantContains: `pg_is_in_recovery: ""`},
185+
{name: "status check failed", recoveryState: "", queryErr: queryErr, wantContains: "last status check failed: connection refused"},
186+
}
187+
188+
for _, tc := range testCases {
189+
t.Run(tc.name, func(t *testing.T) {
190+
err := startTimeoutError(tc.recoveryState, tc.queryErr)
191+
192+
require.Error(t, err)
193+
assert.Contains(t, err.Error(), "postgres start timeout")
194+
assert.Contains(t, err.Error(), tc.wantContains)
195+
196+
if tc.queryErr != nil {
197+
assert.ErrorIs(t, err, tc.queryErr)
198+
}
199+
})
200+
}
201+
}

‎engine/internal/retrieval/retrieval.go‎

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -368,9 +368,7 @@ func (r *Retrieval) run(ctx context.Context, fsm pool.FSManager) (err error) {
368368
r.State.cleanAlerts()
369369
}
370370

371-
var existsErr *thinclones.SnapshotExistsError
372-
373-
if err := r.SnapshotData(ctx, poolName); err != nil && (err != errNoJobs || !errors.As(err, &existsErr)) {
371+
if err := r.SnapshotData(ctx, poolName); err != nil && !isSnapshotExempt(err) {
374372
return err
375373
}
376374

@@ -390,6 +388,14 @@ func (r *Retrieval) run(ctx context.Context, fsm pool.FSManager) (err error) {
390388
return nil
391389
}
392390

391+
// isSnapshotExempt reports whether a SnapshotData error must not abort the run: having no
392+
// snapshot jobs or an already existing snapshot still leaves the pool ready to be activated.
393+
func isSnapshotExempt(err error) bool {
394+
var existsErr *thinclones.SnapshotExistsError
395+
396+
return errors.Is(err, errNoJobs) || errors.As(err, &existsErr)
397+
}
398+
393399
// RefreshData runs a group of data refresh jobs.
394400
func (r *Retrieval) RefreshData(ctx context.Context, poolName string) error {
395401
fsm, err := r.poolManager.GetFSManager(poolName)

‎engine/internal/retrieval/retrieval_test.go‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,13 +2,16 @@ package retrieval
22

33
import (
44
"context"
5+
"errors"
6+
"fmt"
57
"os"
68
"path"
79
"testing"
810

911
"github.com/stretchr/testify/assert"
1012
"github.com/stretchr/testify/require"
1113

14+
"gitlab.com/postgres-ai/database-lab/v3/internal/provision/thinclones"
1215
"gitlab.com/postgres-ai/database-lab/v3/internal/retrieval/config"
1316
"gitlab.com/postgres-ai/database-lab/v3/internal/retrieval/engine/postgres/logical"
1417
"gitlab.com/postgres-ai/database-lab/v3/pkg/models"
@@ -306,3 +309,23 @@ func TestSkipRefreshingError(t *testing.T) {
306309
assert.EqualError(t, err, "some error")
307310
})
308311
}
312+
313+
func TestIsSnapshotExempt(t *testing.T) {
314+
testCases := []struct {
315+
name string
316+
err error
317+
exempt bool
318+
}{
319+
{name: "no jobs", err: errNoJobs, exempt: true},
320+
{name: "wrapped no jobs", err: fmt.Errorf("snapshot: %w", errNoJobs), exempt: true},
321+
{name: "snapshot exists", err: thinclones.NewSnapshotExistsError("snap"), exempt: true},
322+
{name: "wrapped snapshot exists", err: fmt.Errorf("snapshot: %w", thinclones.NewSnapshotExistsError("snap")), exempt: true},
323+
{name: "other error", err: errors.New("zfs failed"), exempt: false},
324+
}
325+
326+
for _, tc := range testCases {
327+
t.Run(tc.name, func(t *testing.T) {
328+
assert.Equal(t, tc.exempt, isSnapshotExempt(tc.err))
329+
})
330+
}
331+
}

‎engine/internal/runci/server.go‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,7 @@ func NewServer(cfg *Config, dle *dblabapi.Client, platform *platform.Service, co
5454
func (s *Server) Run() error {
5555
r := mux.NewRouter().StrictSlash(true)
5656

57-
authMW := mw.NewAuth(s.config.App.VerificationToken, s.platform)
57+
authMW := mw.NewAuth(mw.StaticToken(s.config.App.VerificationToken), s.platform)
5858

5959
r.HandleFunc("/migration/run", authMW.Authorized(s.runMigration)).Methods(http.MethodPost)
6060
r.HandleFunc("/artifact/download", authMW.Authorized(s.downloadArtifact)).Methods(http.MethodGet)
@@ -63,7 +63,7 @@ func (s *Server) Run() error {
6363

6464
addr := fmt.Sprintf("%s:%d", s.config.App.Host, s.config.App.Port)
6565

66-
s.httpServer = &http.Server{Addr: addr, Handler: mw.Logging(r)}
66+
s.httpServer = &http.Server{Addr: addr, Handler: mw.Logging(mw.Recover(r))}
6767

6868
log.Msg(fmt.Sprintf("Server started listening on %s...", addr))
6969

0 commit comments

Comments
 (0)