diff --git a/internal/controller/exec_sessions.go b/internal/controller/exec_sessions.go index 2ba41ea..6ca6855 100644 --- a/internal/controller/exec_sessions.go +++ b/internal/controller/exec_sessions.go @@ -73,48 +73,46 @@ func (registry *execSessionRegistry) getOrCreate( key execSessionKey, create func() (*execSession, error), ) (*execSession, bool, error) { - for { - registry.mu.Lock() + registry.mu.Lock() - if session, ok := registry.sessions[key]; ok { - registry.mu.Unlock() - - return session, false, nil - } - - if creation, ok := registry.creating[key]; ok { - registry.mu.Unlock() - - select { - case <-ctx.Done(): - return nil, false, ctx.Err() - case <-creation.done: - if creation.err != nil { - return nil, false, creation.err - } - - return creation.session, false, nil - } - } - - creation := &execSessionCreation{done: make(chan struct{})} - registry.creating[key] = creation + if session, ok := registry.sessions[key]; ok { registry.mu.Unlock() - session, err := create() - - registry.mu.Lock() - delete(registry.creating, key) - if err == nil { - registry.sessions[key] = session - } - creation.session = session - creation.err = err - close(creation.done) - registry.mu.Unlock() - - return session, true, err + return session, false, nil } + + if creation, ok := registry.creating[key]; ok { + registry.mu.Unlock() + + select { + case <-ctx.Done(): + return nil, false, ctx.Err() + case <-creation.done: + if creation.err != nil { + return nil, false, creation.err + } + + return creation.session, false, nil + } + } + + creation := &execSessionCreation{done: make(chan struct{})} + registry.creating[key] = creation + registry.mu.Unlock() + + session, err := create() + + registry.mu.Lock() + delete(registry.creating, key) + if err == nil { + registry.sessions[key] = session + } + creation.session = session + creation.err = err + close(creation.done) + registry.mu.Unlock() + + return session, true, err } func (registry *execSessionRegistry) remove(key execSessionKey, expected *execSession) { diff --git a/internal/tests/exec_test.go b/internal/tests/exec_test.go index af87560..18602aa 100644 --- a/internal/tests/exec_test.go +++ b/internal/tests/exec_test.go @@ -2,6 +2,7 @@ package tests_test import ( "bytes" + "context" "encoding/json" "testing" "time" @@ -152,6 +153,10 @@ func TestVMExecSessionReconnectHistory(t *testing.T) { require.NoError(t, err) _ = wsConn.CloseNow() + // Let the detached process finish so this test verifies partial replay + // without relying on live-output timing. + time.Sleep(2 * time.Second) + wsConn, err = devClient.VMs().ExecSession(t.Context(), vmName, client.ExecSessionOptions{ WaitSeconds: 30, Session: sessionID, @@ -314,14 +319,22 @@ func prepareForExec(t *testing.T) (*client.Client, string) { } func readFrame(t *testing.T, wsConn *websocket.Conn) *execstream.Frame { + t.Helper() + var frame execstream.Frame - messageType, payloadBytes, err := wsConn.Read(t.Context()) + readCtx, readCtxCancel := context.WithTimeout(t.Context(), 30*time.Second) + defer readCtxCancel() + + messageType, payloadBytes, err := wsConn.Read(readCtx) require.NoError(t, err) require.Equal(t, websocket.MessageText, messageType) err = json.Unmarshal(payloadBytes, &frame) require.NoError(t, err) + if frame.Type == execstream.FrameTypeError { + require.FailNowf(t, "exec stream error", "%s", frame.Error) + } return &frame }