better fallback handling for v2/v1 fallback · Entire

better fallback handling for v2/v1 fallback

3b30f87→main·

Soph·2mo ago·6 files·+82 added/-19 removed

Sessions

ae52adb63cecView transcript

Changes

6

15 unmodified lines

16
17
18
19
20
21
22
23
24
25
22 unmodified lines

48
49
50
47
51
52
49
50
53
54
55
56
57
58
59
60
61
34 unmodified lines

96
97
98
91
99
100
101
102
103
104
97
98
105
106
100
101
102
103
104
107
108
109
110

15 unmodified lines

\t"github.com/go-git/go-git/v6/plumbing/transport";

// GitProtocolV2 is the value passed via the Git-Protocol header / GIT_PROTOCOL
// env var to negotiate protocol v2 with the remote.
const GitProtocolV2 = "version=2"

// RefService encapsulates the result of source ref discovery and the negotiated
// protocol, providing methods for subsequent fetch and pack operations.
type RefService struct {
22 unmodified lines

\t\treturn refs, &RefService{Protocol: "v1", V1Adv: adv, HeadTarget: headTargetFromAdv(adv)}, nil

\t\tcase "auto", "v2":
\t\t\tdata, err := RequestInfoRefs(ctx, conn, transport.UploadPackService, "version=2")
\t\t\tdata, err := RequestInfoRefs(ctx, conn, transport.UploadPackService, GitProtocolV2)
\t\t\tif err != nil {
\t\t\t\tif protocolMode == "auto" && shouldFallbackToV1AfterV2ProbeError(conn, err) {
\t\t\t\t\treturn listSourceRefsAutoV1(ctx, conn)
\t\t\t\tif protocolMode == "auto" && isSSHScheme(conn) {
\t\t\t\t\trefs, svc, v1Err := listSourceRefsAutoV1(ctx, conn)
\t\t\t\t\tif v1Err != nil {
\t\t\t\t\t\treturn nil, nil, errors.Join(err, v1Err)
\t\t\t\t\t}
\t\t\t\t\treturn refs, svc, nil
\t\t\t\t}
\t\t\t}\n\t\t\treturn nil, nil, err
\t\t\t}
34 unmodified lines

\t\treturn refs, &RefService{Protocol: "v1", V1Adv: adv, HeadTarget: headTargetFromAdv(adv)}, nil
}

func shouldFallbackToV1AfterV2ProbeError(conn Conn, err error) bool {
func isSSHScheme(conn Conn) bool {
\tif conn == nil || conn.Endpoint() == nil {
\t\treturn false
\t}
\tswitch conn.Endpoint().Scheme {
\tcase "ssh", "git+ssh":
\tdefault:
\t\treturn false
\t\treturn true
\t}

\tmsg := err.Error()
\treturn strings.Contains(msg, "Invalid command:") &&
\t\tstrings.Contains(msg, "GIT_PROTOCOL='version=2'") &&
\t\tstrings.Contains(msg, "git-upload-pack")
\treturn false
}

// AdvertisedRefsV1 fetches and decodes v1 advertised refs for the given service. 

Minternal/gitproto/refs.go+14/-11

202 unmodified lines

203
204
205
206
206
207
208
209
210
211
212
211
212
213
214
215
38 unmodified lines

254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
4 unmodified lines

321
322
323
269
324
325
326
327
328
329
330
331
332

202 unmodified lines

\t}
}

func TestListSourceRefsAutoFallsBackToV1AfterSSHV2ProbeCommandRejection(t *testing.T) {
func TestListSourceRefsAutoFallsBackToV1AfterSSHV2ProbeError(t *testing.T) {
\tt.Parallel()

\tconn := &stubConn{
\t\treqInfoRefs: func(_ context.Context, _ string, gitProtocol string) ([]byte, error) {
\t\t\tif gitProtocol == "version=2" {
\t\t\t\treturn nil, errors.New("request info refs: git-upload-pack info-refs: wait for ssh command: exit status 1: Invalid command: GIT_PROTOCOL='version=2' git-upload-pack 'entireio/cli.git'")
\t\t\t}
\t\t\tif gitProtocol == GitProtocolV2 {
\t\t\t\treturn nil, errors.New("ssh server rejected v2 probe")
\t\t\t}

\t\t\tvar body strings.Builder
38 unmodified lines

\t}
}

func TestListSourceRefsAutoJoinsErrorsWhenV1FallbackAlsoFails(t *testing.T) {
\tt.Parallel()

\tv2Err := errors.New("ssh v2 probe failed")
\tv1Err := errors.New("ssh v1 fallback failed")
\tconn := &stubConn{
\t\treqInfoRefs: func(_ context.Context, _ string, gitProtocol string) ([]byte, error) {
\t\t\tif gitProtocol == GitProtocolV2 {
\t\t\t\treturn nil, v2Err
\t\t\t}
\t\t\treturn nil, v1Err
\t\t}
\t}

\t_, _, err := ListSourceRefs(t.Context(), conn, "auto", nil)
\tif err == nil {
\t\tt.Fatal("expected error when both v2 probe and v1 fallback fail")
\t}
\tif !errors.Is(err, v2Err) {
\t\tt.Errorf("err does not wrap v2 error: %v", err)
\t}
\tif !errors.Is(err, v1Err) {
\t\tt.Errorf("err does not wrap v1 error: %v", err)
\t}
}

func TestListSourceRefsAutoDoesNotFallBackForNonSSH(t *testing.T) {
\tt.Parallel()

\tv2Err := errors.New("https v2 probe failed")
\tv1Called := false
\tconn := &stubConn{
\t\treqInfoRefs: func(_ context.Context, _ string, gitProtocol string) ([]byte, error) {
\t\t\tif gitProtocol == GitProtocolV2 {
\t\t\t\treturn nil, v2Err
\t\t\t}
\t\tv1Called = true
\t\treturn nil, errors.New("unexpected v1 call")
\t\t},
\t\tendpoint: &url.URL{Scheme: "https", Host: "github.com"},
\t}

\t_, _, err := ListSourceRefs(t.Context(), conn, "auto", nil)
\tif err == nil {
\t\tt.Fatal("expected error from v2 probe")
\t}
\tif !errors.Is(err, v2Err) {
\t\tt.Errorf("err does not wrap v2 error: %v", err)
\t}
\tif v1Called {
\t\tt.Error("v1 fallback should not be attempted for non-SSH schemes")
\t}
}

type stubConn struct {
\treqInfoRefs func(ctx context.Context, service string, gitProtocol string) ([]byte, error)
\tendpoint    *url.URL
}

func (s *stubConn) RequestInfoRefs(ctx context.Context, service string, gitProtocol string) ([]byte, error) {
4 unmodified lines

\treturn nil, errors.New("unexpected PostRPCStreamBody call")
}

func (s *stubConn) Endpoint() *url.URL { return &url.URL{Scheme: "ssh", Host: "github.com"} }
func (s *stubConn) Endpoint() *url.URL {
\tif s.endpoint != nil {
\t\treturn s.endpoint
\t}
\treturn &url.URL{Scheme: "ssh", Host: "github.com"}
}

func (s *stubConn) ProgressWriter() io.Writer { return nil }

Minternal/gitproto/refs_test.go+64/-4

261 unmodified lines

262
263
264
265
265
266
267
268

261 unmodified lines

\treq.Header.Set("User-Agent", capability.DefaultAgent())
\treq.Header.Set(StatsPhaseHeader, phase)
\tif v2 {
\t\treq.Header.Set("Git-Protocol", "version=2")
\t\treq.Header.Set("Git-Protocol", GitProtocolV2)
\t}
\tApplyAuth(req, c.Auth)

Minternal/gitproto/smarthttp.go+1/-1

150 unmodified lines

151
152
153
154
154
155
156
157

150 unmodified lines

\tdone := make(chan error, 1)
\tgo func() {
\t\t_, err := RequestInfoRefs(ctx, conn, "git-upload-pack", "version=2")
\t\t_, err := RequestInfoRefs(ctx, conn, "git-upload-pack", GitProtocolV2)
\t\tdone <- err
\t}()

Minternal/gitproto/smarthttp_test.go+1/-1

79 unmodified lines

80
81
82
83
83
84
85
86

79 unmodified lines

\t_ = phase // phase labels are HTTP-only today; SSH transport has no per-RPC stats tagging
\tgitProtocol := ""
\tif v2 {
\t\tgitProtocol = "version=2"
\t\tgitProtocol = GitProtocolV2
\t}
\tcmd, stderr, err := c.startRPC(ctx, service, gitProtocol)
\tif err != nil {

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

34 unmodified lines

35
36
37
38
38
39
40
41

34 unmodified lines

\tenv := newSSHShimEnv(t)
\tconn := newSSHTestConn(t, "ssh://example.com/repo.git", env.script)

\tbody, err := conn.RequestInfoRefs(t.Context(), "git-upload-pack", "version=2")
\tbody, err := conn.RequestInfoRefs(t.Context(), "git-upload-pack", GitProtocolV2)
\tif err != nil {
\t\tt.Fatalf("RequestInfoRefs: %v", err)
\t}

Minternal/gitproto/ssh_test.go+1/-1