ssh: clean up failed info-refs startup · Entire

ssh: clean up failed info-refs startup

9ada339→main· Soph·2mo ago·2 files·+61 added/-1 removed

Sessions

5f82375c7a8aView transcript

[?
yes, start implementing, make a new branch, make meaningful commits, and add tests as you go, when it makes sense create tests firstCodex·GPT-5.4·1 step](/content/gh/entireio/git-sync/session/019e260b-71e8-73a1-9e68-5857379ffab6#timeline-5f82375c7a8a/index.html)

Changes

2

51 unmodified lines

52
53
54
55
56
57
58
59
56
60
61
62
63
131 unmodified lines

195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229

51 unmodified lines

if err != nil {
        return nil, err
    }
    return requestInfoRefsWithCommand(ctx, service, cmd, stderr)
}

func requestInfoRefsWithCommand(ctx context.Context, service string, cmd *sshCommand, stderr *sshCommandError) ([]byte, error) {
    if err := cmd.Stdin.Close(); err != nil {
        return nil, fmt.Errorf("close ssh stdin for %s: %w", service, err)
        return nil, fmt.Errorf("close ssh stdin for %s: %w", service, errors.Join(err, cleanupSSHCommand(cmd)))
    }

data, readErr := io.ReadAll(cmd.Stdout)
    waitErr := cmd.wait()
    131 unmodified lines

Cmd    *exec.Cmd
    Stdin  io.WriteCloser
    Stdout io.ReadCloser
    waitFn func() error
}

func (c *sshCommand) wait() error {
    if c.waitFn != nil {
        return c.waitFn()
    }
    if err := c.Cmd.Wait(); err != nil {
        return fmt.Errorf("wait for ssh command: %w", err)
    }
    return nil
}

func cleanupSSHCommand(cmd *sshCommand) error {
    if cmd == nil {
        return nil
    }
    var errs []error
    if cmd.Stdout != nil {
        if err := cmd.Stdout.Close(); err != nil {
            errs = append(errs, fmt.Errorf("close ssh stdout: %w", err))
        }
    }
    if err := cmd.wait(); err != nil {
        err
    }
    return errors.Join(errs...)
}

func discardSSHAdvertisement(stdout io.ReadCloser) (io.ReadCloser, error) {
    buffered := bufio.NewReader(stdout)
    header, err := buffered.Peek(4)

Minternal/gitproto/ssh.go+25/-1

167 unmodified lines

168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215

167 unmodified lines

}

func TestRequestInfoRefsCleansUpWhenStdinCloseFails(t *testing.T) {

t.Parallel()

stdoutClosed := false
    waitCalled := false
    _, err := requestInfoRefsWithCommand(
    
t.Context(),
        "git-upload-pack",
        &sshCommand{
            Stdin:  closeWriterFunc(func() error { return errors.New("close failed") }),
            Stdout: closeReaderFunc(func() error { stdoutClosed = true; return nil }),
            waitFn: func() error { waitCalled = true; return nil },
        },
        &sshCommandError{},
    )
    if err == nil || !strings.Contains(err.Error(), "close ssh stdin for git-upload-pack") {
        t.Fatalf("requestInfoRefsWithCommand error = %v", err)
    }
    if !stdoutClosed {
        t.Fatal("stdout was not closed on stdin-close failure")
    }
    if !waitCalled {
        t.Fatal("wait was not called on stdin-close failure")
    }
}

type sshShimEnv struct {
    script     string
    logFile    string
    bodyPrefix string
}

type closeWriterFunc func() error

func (f closeWriterFunc) Write(p []byte) (int, error) { return len(p), nil }
func (f closeWriterFunc) Close() error                { return f() }

type closeReaderFunc func() error

func (f closeReaderFunc) Read([]byte) (int, error) { return 0, io.EOF }
func (f closeReaderFunc) Close() error             { return f() }

func newSSHShimEnv(t *testing.T) sshShimEnv {
    t.Helper()
    dir := t.TempDir()