Commit b94ae2638b
Verified · cmc
Layout: unified · split
cmd/gitbayd/main.go +2 −2
| @@ -249,7 +249,7 @@ func serveCmd() *cobra.Command { | |||
| 249 | slog.Info("ssh handled by host sshd (ssh.mode = system)") | 249 | slog.Info("ssh handled by host sshd (ssh.mode = system)") |
| 250 | } | 250 | } |
| 251 | 251 | ||
| 252 | web := httpd.New(cfg, st) | 252 | web := httpd.New(cfg, st, nil) |
| 253 | // Header and idle timeouts bound what an idle or slow client can | 253 | // Header and idle timeouts bound what an idle or slow client can |
| 254 | // hold open. No write timeout: archives and upload-pack stream | 254 | // hold open. No write timeout: archives and upload-pack stream |
| 255 | // for as long as they take (#104). | 255 | // for as long as they take (#104). |
| @@ -338,7 +338,7 @@ func serveCmd() *cobra.Command { | |||
| 338 | } | 338 | } |
| 339 | slog.Info("git-daemon listening", "addr", gln.Addr()) | 339 | slog.Info("git-daemon listening", "addr", gln.Addr()) |
| 340 | gitLn = gln | 340 | gitLn = gln |
| 341 | go func() { errCh <- gitd.New(cfg, st).Serve(gln) }() | 341 | go func() { errCh <- gitd.New(cfg, st, nil).Serve(gln) }() |
| 342 | } | 342 | } |
| 343 | 343 | ||
| 344 | select { | 344 | select { |
internal/gitd/gitd.go +36 −12
| @@ -7,24 +7,26 @@ import ( | |||
| 7 | "fmt" | 7 | "fmt" |
| 8 | "io" | 8 | "io" |
| 9 | "net" | 9 | "net" |
| 10 | "os" | ||
| 11 | "os/exec" | ||
| 12 | "strconv" | 10 | "strconv" |
| 13 | "strings" | 11 | "strings" |
| 14 | "time" | 12 | "time" |
| 15 | 13 | ||
| 16 | "gitbay.org/gitbay/internal/config" | 14 | "gitbay.org/gitbay/internal/config" |
| 17 | "gitbay.org/gitbay/internal/control" | 15 | "gitbay.org/gitbay/internal/control" |
| 16 | "gitbay.org/gitbay/internal/gitutil" | ||
| 17 | "gitbay.org/gitbay/internal/packlimit" | ||
| 18 | "gitbay.org/gitbay/internal/store" | 18 | "gitbay.org/gitbay/internal/store" |
| 19 | "gitbay.org/gitbay/internal/toolpath" | ||
| 20 | ) | 19 | ) |
| 21 | 20 | ||
| 22 | type Server struct { | 21 | type Server struct { |
| 23 | cfg config.Config | 22 | cfg config.Config |
| 24 | st *store.Store | 23 | st *store.Store |
| 24 | packs *packlimit.Limiter | ||
| 25 | } | 25 | } |
| 26 | 26 | ||
| 27 | func New(cfg config.Config, st *store.Store) *Server { return &Server{cfg: cfg, st: st} } | 27 | func New(cfg config.Config, st *store.Store, packs *packlimit.Limiter) *Server { |
| 28 | return &Server{cfg: cfg, st: st, packs: packs} | ||
| 29 | } | ||
| 28 | 30 | ||
| 29 | func (s *Server) Serve(ln net.Listener) error { | 31 | func (s *Server) Serve(ln net.Listener) error { |
| 30 | for { | 32 | for { |
| @@ -68,13 +70,35 @@ func (s *Server) handle(conn net.Conn) { | |||
| 68 | return | 70 | return |
| 69 | } | 71 | } |
| 70 | 72 | ||
| 73 | host, _, _ := net.SplitHostPort(conn.RemoteAddr().String()) | ||
| 74 | // A nil done: a queued client that leaves, or a restart, does not | ||
| 75 | // end the wait; only the limiter's wait does. | ||
| 76 | release, err := s.packs.Acquire(nil, "ip:"+host) | ||
| 77 | if err != nil { | ||
| 78 | writeErr(conn, err.Error()) | ||
| 79 | return | ||
| 80 | } | ||
| 81 | // Deferred before git runs, so it fires after git has exited and | ||
| 82 | // been waited for. | ||
| 83 | defer release() | ||
| 84 | out, stalled, unwatch := s.packs.Watch(conn) | ||
| 85 | defer unwatch() | ||
| 86 | kill := make(chan struct{}) | ||
| 87 | finished := make(chan struct{}) | ||
| 88 | defer close(finished) | ||
| 89 | go func() { | ||
| 90 | select { | ||
| 91 | case <-finished: | ||
| 92 | case <-stalled: | ||
| 93 | close(kill) | ||
| 94 | // A write blocked on a client that stopped reading outlives | ||
| 95 | // git; closing the connection ends it and the stdin copy. | ||
| 96 | conn.Close() | ||
| 97 | } | ||
| 98 | }() | ||
| 99 | |||
| 71 | dir := control.RepoDir(s.cfg.Server.Root, repo.OwnerName, repo.Name) | 100 | dir := control.RepoDir(s.cfg.Server.Root, repo.OwnerName, repo.Name) |
| 72 | cmd := exec.Command(toolpath.Look("git"), "upload-pack", dir) | 101 | gitutil.Transport("git-upload-pack", dir, conn, out, io.Discard, protoEnv, 0, kill) |
| 73 | cmd.Env = append(os.Environ(), protoEnv...) | ||
| 74 | cmd.Stdin = conn | ||
| 75 | cmd.Stdout = conn | ||
| 76 | cmd.Stderr = io.Discard | ||
| 77 | cmd.Run() | ||
| 78 | } | 102 | } |
| 79 | 103 | ||
| 80 | func readPktLine(r io.Reader) (string, error) { | 104 | func readPktLine(r io.Reader) (string, error) { |
internal/gitd/gitd_test.go added +51
| @@ -0,0 +1,51 @@ | |||
| 1 | package gitd | ||
| 2 | |||
| 3 | import ( | ||
| 4 | "fmt" | ||
| 5 | "net" | ||
| 6 | "path/filepath" | ||
| 7 | "strings" | ||
| 8 | "testing" | ||
| 9 | "time" | ||
| 10 | |||
| 11 | "gitbay.org/gitbay/internal/config" | ||
| 12 | "gitbay.org/gitbay/internal/packlimit" | ||
| 13 | "gitbay.org/gitbay/internal/store" | ||
| 14 | ) | ||
| 15 | |||
| 16 | func TestBusyAnswersERR(t *testing.T) { | ||
| 17 | st, err := store.Open(filepath.Join(t.TempDir(), "gitbay.db")) | ||
| 18 | if err != nil { | ||
| 19 | t.Fatal(err) | ||
| 20 | } | ||
| 21 | defer st.Close() | ||
| 22 | if err := st.MigrateUp(); err != nil { | ||
| 23 | t.Fatal(err) | ||
| 24 | } | ||
| 25 | uid, err := st.CreateUser("alice", false) | ||
| 26 | if err != nil { | ||
| 27 | t.Fatal(err) | ||
| 28 | } | ||
| 29 | repoID, err := st.CreateRepo("user", uid, "app", "public") | ||
| 30 | if err != nil { | ||
| 31 | t.Fatal(err) | ||
| 32 | } | ||
| 33 | if _, err := st.UpdateRepoSettings(repoID, func(rs *store.RepoSettings) { rs.GitDaemon = true }); err != nil { | ||
| 34 | t.Fatal(err) | ||
| 35 | } | ||
| 36 | packs := packlimit.New(1, 0, 0, time.Second) | ||
| 37 | hold, _ := packs.Acquire(nil, "ip:elsewhere") | ||
| 38 | defer hold() | ||
| 39 | |||
| 40 | s := New(config.Config{Server: config.Server{Root: t.TempDir()}}, st, packs) | ||
| 41 | client, server := net.Pipe() | ||
| 42 | defer client.Close() | ||
| 43 | go s.handle(server) | ||
| 44 | req := "git-upload-pack /alice/app.git\x00host=x\x00" | ||
| 45 | fmt.Fprintf(client, "%04x%s", len(req)+4, req) | ||
| 46 | client.SetReadDeadline(time.Now().Add(5 * time.Second)) | ||
| 47 | line, err := readPktLine(client) | ||
| 48 | if err != nil || !strings.HasPrefix(line, "ERR ") || !strings.Contains(line, "busy") { | ||
| 49 | t.Fatalf("got %q, %v", line, err) | ||
| 50 | } | ||
| 51 | } | ||
internal/gitutil/gitutil.go +8 −1
| @@ -47,7 +47,7 @@ func Transport(service, repoPath string, stdin io.Reader, stdout, errW io.Writer | |||
| 47 | } | 47 | } |
| 48 | if service == "git-upload-pack" { | 48 | if service == "git-upload-pack" { |
| 49 | // Keepalives while pack-objects is still counting keep a | 49 | // Keepalives while pack-objects is still counting keep a |
| 50 | // healthy clone writing; sshd kills one that goes quiet. | 50 | // healthy clone writing; a limited transport kills one that goes quiet. |
| 51 | args = []string{"-c", "uploadpack.keepAlive=5"} | 51 | args = []string{"-c", "uploadpack.keepAlive=5"} |
| 52 | } | 52 | } |
| 53 | args = append(args, strings.TrimPrefix(service, "git-"), repoPath) | 53 | args = append(args, strings.TrimPrefix(service, "git-"), repoPath) |
| @@ -59,6 +59,13 @@ func Transport(service, repoPath string, stdin io.Reader, stdout, errW io.Writer | |||
| 59 | cmd.Stdin = stdin | 59 | cmd.Stdin = stdin |
| 60 | cmd.Stdout = stdout | 60 | cmd.Stdout = stdout |
| 61 | cmd.Stderr = errW | 61 | cmd.Stderr = errW |
| 62 | return RunUntil(cmd, cancel) | ||
| 63 | } | ||
| 64 | |||
| 65 | // RunUntil runs cmd in its own process group. Closing cancel kills the | ||
| 66 | // group; RunUntil returns only once cmd has been waited for. A nil | ||
| 67 | // cancel never fires. | ||
| 68 | func RunUntil(cmd *exec.Cmd, cancel <-chan struct{}) error { | ||
| 62 | ownProcessGroup(cmd) | 69 | ownProcessGroup(cmd) |
| 63 | if err := cmd.Start(); err != nil { | 70 | if err := cmd.Start(); err != nil { |
| 64 | return err | 71 | return err |
internal/httpd/account_test.go +4 −4
| @@ -34,7 +34,7 @@ func TestAccountPagePushToggleAndDevices(t *testing.T) { | |||
| 34 | t.Fatal(err) | 34 | t.Fatal(err) |
| 35 | } | 35 | } |
| 36 | 36 | ||
| 37 | s := New(config.Default(), st) | 37 | s := New(config.Default(), st, nil) |
| 38 | rr := httptest.NewRecorder() | 38 | rr := httptest.NewRecorder() |
| 39 | req := httptest.NewRequest("GET", "/settings", nil) | 39 | req := httptest.NewRequest("GET", "/settings", nil) |
| 40 | s.accountPage(rr, req, store.User{ID: uid, Username: "alice"}) | 40 | s.accountPage(rr, req, store.User{ID: uid, Username: "alice"}) |
| @@ -78,7 +78,7 @@ func TestAccountSubmitNotifyPush(t *testing.T) { | |||
| 78 | t.Fatal(err) | 78 | t.Fatal(err) |
| 79 | } | 79 | } |
| 80 | u := store.User{ID: uid, Username: "alice"} | 80 | u := store.User{ID: uid, Username: "alice"} |
| 81 | s := New(config.Default(), st) | 81 | s := New(config.Default(), st, nil) |
| 82 | 82 | ||
| 83 | rr := submitAccountForm(t, s, u, url.Values{"field": {"notify-push"}, "push": {"on"}}) | 83 | rr := submitAccountForm(t, s, u, url.Values{"field": {"notify-push"}, "push": {"on"}}) |
| 84 | if rr.Code != http.StatusSeeOther { | 84 | if rr.Code != http.StatusSeeOther { |
| @@ -117,7 +117,7 @@ func TestAccountSubmitDeviceRemove(t *testing.T) { | |||
| 117 | if err != nil { | 117 | if err != nil { |
| 118 | t.Fatal(err) | 118 | t.Fatal(err) |
| 119 | } | 119 | } |
| 120 | s := New(config.Default(), st) | 120 | s := New(config.Default(), st, nil) |
| 121 | 121 | ||
| 122 | idStr := strconv.FormatInt(id, 10) | 122 | idStr := strconv.FormatInt(id, 10) |
| 123 | 123 | ||
| @@ -158,7 +158,7 @@ func TestAccountPageMasksAShortDeviceToken(t *testing.T) { | |||
| 158 | t.Fatal(err) | 158 | t.Fatal(err) |
| 159 | } | 159 | } |
| 160 | 160 | ||
| 161 | s := New(config.Default(), st) | 161 | s := New(config.Default(), st, nil) |
| 162 | rr := httptest.NewRecorder() | 162 | rr := httptest.NewRecorder() |
| 163 | s.accountPage(rr, httptest.NewRequest("GET", "/settings", nil), store.User{ID: uid, Username: "alice"}) | 163 | s.accountPage(rr, httptest.NewRequest("GET", "/settings", nil), store.User{ID: uid, Username: "alice"}) |
| 164 | 164 | ||
internal/httpd/admin_test.go +1 −1
| @@ -48,7 +48,7 @@ func TestAdminPageShowsThePushQueue(t *testing.T) { | |||
| 48 | t.Fatal(err) | 48 | t.Fatal(err) |
| 49 | } | 49 | } |
| 50 | 50 | ||
| 51 | s := New(config.Default(), st) | 51 | s := New(config.Default(), st, nil) |
| 52 | rr := httptest.NewRecorder() | 52 | rr := httptest.NewRecorder() |
| 53 | s.adminPage(rr, httptest.NewRequest("GET", "/admin", nil), store.User{ID: uid, Username: "root", IsAdmin: true}) | 53 | s.adminPage(rr, httptest.NewRequest("GET", "/admin", nil), store.User{ID: uid, Username: "root", IsAdmin: true}) |
| 54 | if rr.Code != http.StatusOK { | 54 | if rr.Code != http.StatusOK { |
internal/httpd/anchors_test.go +2 −2
| @@ -26,7 +26,7 @@ func TestMarkdownHeadingAnchors(t *testing.T) { | |||
| 26 | // The stylesheet carries an ETag and a cache lifetime; a revalidation | 26 | // The stylesheet carries an ETag and a cache lifetime; a revalidation |
| 27 | // with the same tag is a 304 with no body (#132). | 27 | // with the same tag is a 304 with no body (#132). |
| 28 | func TestStylesheetRevalidates(t *testing.T) { | 28 | func TestStylesheetRevalidates(t *testing.T) { |
| 29 | s := New(config.Default(), nil) | 29 | s := New(config.Default(), nil, nil) |
| 30 | first := httptest.NewRecorder() | 30 | first := httptest.NewRecorder() |
| 31 | s.stylesheet(first, httptest.NewRequest("GET", "/static/style.css", nil)) | 31 | s.stylesheet(first, httptest.NewRequest("GET", "/static/style.css", nil)) |
| 32 | tag := first.Header().Get("ETag") | 32 | tag := first.Header().Get("ETag") |
| @@ -69,7 +69,7 @@ func TestStylesheetURLCarriesTheBuildHash(t *testing.T) { | |||
| 69 | t.Errorf("the page does not link %s", want) | 69 | t.Errorf("the page does not link %s", want) |
| 70 | } | 70 | } |
| 71 | 71 | ||
| 72 | s := New(config.Default(), nil) | 72 | s := New(config.Default(), nil, nil) |
| 73 | versioned := httptest.NewRecorder() | 73 | versioned := httptest.NewRecorder() |
| 74 | s.stylesheet(versioned, httptest.NewRequest("GET", "/static/style.css?v="+stylesheetHash, nil)) | 74 | s.stylesheet(versioned, httptest.NewRequest("GET", "/static/style.css?v="+stylesheetHash, nil)) |
| 75 | if cc := versioned.Header().Get("Cache-Control"); !strings.Contains(cc, "immutable") { | 75 | if cc := versioned.Header().Get("Cache-Control"); !strings.Contains(cc, "immutable") { |
internal/httpd/checkorigin_test.go +1 −1
| @@ -40,7 +40,7 @@ func TestMutatingRoutesRequireCheckOrigin(t *testing.T) { | |||
| 40 | // there is no second code path in routes.go for looping over both to | 40 | // there is no second code path in routes.go for looping over both to |
| 41 | // reach; "open" alone matches production and is enough. | 41 | // reach; "open" alone matches production and is enough. |
| 42 | cfg.Registration.Mode = "open" | 42 | cfg.Registration.Mode = "open" |
| 43 | s := New(cfg, nil) | 43 | s := New(cfg, nil, nil) |
| 44 | 44 | ||
| 45 | for _, r := range s.Routes() { | 45 | for _, r := range s.Routes() { |
| 46 | if !r.Mutating { | 46 | if !r.Mutating { |
internal/httpd/clientip_test.go +1 −1
| @@ -28,7 +28,7 @@ func TestClientIPBehindProxy(t *testing.T) { | |||
| 28 | for _, tc := range cases { | 28 | for _, tc := range cases { |
| 29 | cfg := config.Default() | 29 | cfg := config.Default() |
| 30 | cfg.HTTP.TrustedProxies = tc.proxies | 30 | cfg.HTTP.TrustedProxies = tc.proxies |
| 31 | s := New(cfg, nil) | 31 | s := New(cfg, nil, nil) |
| 32 | r := httptest.NewRequest("GET", "/api/v1/read", nil) | 32 | r := httptest.NewRequest("GET", "/api/v1/read", nil) |
| 33 | r.RemoteAddr = tc.remote | 33 | r.RemoteAddr = tc.remote |
| 34 | if tc.xff != "" { | 34 | if tc.xff != "" { |
internal/httpd/fonts_test.go +2 −2
| @@ -17,7 +17,7 @@ import ( | |||
| 17 | // and the @font-face URLs were once maintained by hand and drifted, so | 17 | // and the @font-face URLs were once maintained by hand and drifted, so |
| 18 | // gitbay.org served no web font at all (#102). | 18 | // gitbay.org served no web font at all (#102). |
| 19 | func TestStylesheetFontsAreServed(t *testing.T) { | 19 | func TestStylesheetFontsAreServed(t *testing.T) { |
| 20 | s := New(config.Default(), nil) | 20 | s := New(config.Default(), nil, nil) |
| 21 | byPattern := map[string]http.HandlerFunc{} | 21 | byPattern := map[string]http.HandlerFunc{} |
| 22 | for _, r := range s.Routes() { | 22 | for _, r := range s.Routes() { |
| 23 | if r.Method == "GET" { | 23 | if r.Method == "GET" { |
| @@ -49,7 +49,7 @@ func TestStylesheetFontsAreServed(t *testing.T) { | |||
| 49 | // TestLandingImagesAreServed: every file under static/img has a route | 49 | // TestLandingImagesAreServed: every file under static/img has a route |
| 50 | // that answers 200 with an image or video type, and a video answers Range. | 50 | // that answers 200 with an image or video type, and a video answers Range. |
| 51 | func TestLandingImagesAreServed(t *testing.T) { | 51 | func TestLandingImagesAreServed(t *testing.T) { |
| 52 | s := New(config.Default(), nil) | 52 | s := New(config.Default(), nil, nil) |
| 53 | byPattern := map[string]http.HandlerFunc{} | 53 | byPattern := map[string]http.HandlerFunc{} |
| 54 | for _, r := range s.Routes() { | 54 | for _, r := range s.Routes() { |
| 55 | if r.Method == "GET" { | 55 | if r.Method == "GET" { |
internal/httpd/logindisabled_test.go +1 −1
| @@ -40,7 +40,7 @@ func TestLoginRefusesTokenForDisabledAccount(t *testing.T) { | |||
| 40 | t.Fatal(err) | 40 | t.Fatal(err) |
| 41 | } | 41 | } |
| 42 | 42 | ||
| 43 | s := New(config.Default(), st) | 43 | s := New(config.Default(), st, nil) |
| 44 | rr := httptest.NewRecorder() | 44 | rr := httptest.NewRecorder() |
| 45 | req := httptest.NewRequest("GET", "/login?token="+tok, nil) | 45 | req := httptest.NewRequest("GET", "/login?token="+tok, nil) |
| 46 | s.login(rr, req) | 46 | s.login(rr, req) |
internal/httpd/packlimit_test.go added +283
| @@ -0,0 +1,283 @@ | |||
| 1 | package httpd | ||
| 2 | |||
| 3 | import ( | ||
| 4 | "context" | ||
| 5 | "crypto/rand" | ||
| 6 | "fmt" | ||
| 7 | "net" | ||
| 8 | "net/http" | ||
| 9 | "net/http/httptest" | ||
| 10 | "os" | ||
| 11 | "os/exec" | ||
| 12 | "path/filepath" | ||
| 13 | "strconv" | ||
| 14 | "strings" | ||
| 15 | "sync" | ||
| 16 | "testing" | ||
| 17 | "time" | ||
| 18 | |||
| 19 | "gitbay.org/gitbay/internal/config" | ||
| 20 | "gitbay.org/gitbay/internal/control" | ||
| 21 | "gitbay.org/gitbay/internal/gitutil" | ||
| 22 | "gitbay.org/gitbay/internal/packlimit" | ||
| 23 | "gitbay.org/gitbay/internal/store" | ||
| 24 | ) | ||
| 25 | |||
| 26 | func busyServer(t *testing.T) *Server { | ||
| 27 | t.Helper() | ||
| 28 | st, err := store.Open(filepath.Join(t.TempDir(), "gitbay.db")) | ||
| 29 | if err != nil { | ||
| 30 | t.Fatal(err) | ||
| 31 | } | ||
| 32 | t.Cleanup(func() { st.Close() }) | ||
| 33 | if err := st.MigrateUp(); err != nil { | ||
| 34 | t.Fatal(err) | ||
| 35 | } | ||
| 36 | uid, err := st.CreateUser("alice", false) | ||
| 37 | if err != nil { | ||
| 38 | t.Fatal(err) | ||
| 39 | } | ||
| 40 | if _, err := st.CreateRepo("user", uid, "app", "public"); err != nil { | ||
| 41 | t.Fatal(err) | ||
| 42 | } | ||
| 43 | packs := packlimit.New(1, 0, 0, time.Second) | ||
| 44 | hold, err := packs.Acquire(nil, "ip:elsewhere") | ||
| 45 | if err != nil { | ||
| 46 | t.Fatal(err) | ||
| 47 | } | ||
| 48 | t.Cleanup(hold) | ||
| 49 | var cfg config.Config | ||
| 50 | cfg.Server.Root = t.TempDir() | ||
| 51 | return &Server{cfg: cfg, st: st, packs: packs, stopping: make(chan struct{})} | ||
| 52 | } | ||
| 53 | |||
| 54 | func post(s *Server, body string) *httptest.ResponseRecorder { | ||
| 55 | r := httptest.NewRequest("POST", "/alice/app/git-upload-pack", strings.NewReader(body)) | ||
| 56 | r.SetPathValue("owner", "alice") | ||
| 57 | r.SetPathValue("repo", "app") | ||
| 58 | w := httptest.NewRecorder() | ||
| 59 | s.uploadPack(w, r) | ||
| 60 | return w | ||
| 61 | } | ||
| 62 | |||
| 63 | func TestUploadPackBusyIs503(t *testing.T) { | ||
| 64 | w := post(busyServer(t), "0000") | ||
| 65 | if w.Code != http.StatusServiceUnavailable || w.Header().Get("Retry-After") == "" { | ||
| 66 | t.Fatalf("status %d, Retry-After %q", w.Code, w.Header().Get("Retry-After")) | ||
| 67 | } | ||
| 68 | } | ||
| 69 | |||
| 70 | // A protocol v2 ref listing generates no pack and is never queued. | ||
| 71 | func TestLsRefsBypassesTheLimit(t *testing.T) { | ||
| 72 | w := post(busyServer(t), "0014command=ls-refs\n0000") | ||
| 73 | if w.Code == http.StatusServiceUnavailable { | ||
| 74 | t.Fatal("ls-refs was held to the pack limit") | ||
| 75 | } | ||
| 76 | } | ||
| 77 | |||
| 78 | // A request with a valid bearer token counts against the account, the | ||
| 79 | // key SSH uses; anything else against the client address. | ||
| 80 | func TestPackPrincipal(t *testing.T) { | ||
| 81 | s := busyServer(t) | ||
| 82 | alice, err := s.st.UserByUsername("alice") | ||
| 83 | if err != nil { | ||
| 84 | t.Fatal(err) | ||
| 85 | } | ||
| 86 | if err := s.st.CreateAPIToken(alice.ID, "t", store.HashToken("secret"), "read", nil, 0); err != nil { | ||
| 87 | t.Fatal(err) | ||
| 88 | } | ||
| 89 | r := httptest.NewRequest("POST", "/alice/app/git-upload-pack", nil) | ||
| 90 | r.RemoteAddr = "192.0.2.7:4000" | ||
| 91 | if got := s.packPrincipal(r); got != "ip:192.0.2.7" { | ||
| 92 | t.Fatalf("anonymous: %q", got) | ||
| 93 | } | ||
| 94 | r.Header.Set("Authorization", "Bearer wrong") | ||
| 95 | if got := s.packPrincipal(r); got != "ip:192.0.2.7" { | ||
| 96 | t.Fatalf("bad token: %q", got) | ||
| 97 | } | ||
| 98 | want := "user:" + strconv.FormatInt(alice.ID, 10) | ||
| 99 | r.Header.Set("Authorization", "Bearer secret") | ||
| 100 | if got := s.packPrincipal(r); got != want { | ||
| 101 | t.Fatalf("token: %q, want %q", got, want) | ||
| 102 | } | ||
| 103 | if err := s.st.CreateWebSession(store.HashToken("sess"), alice.ID, time.Hour); err != nil { | ||
| 104 | t.Fatal(err) | ||
| 105 | } | ||
| 106 | r = httptest.NewRequest("POST", "/alice/app/git-upload-pack", nil) | ||
| 107 | r.RemoteAddr = "192.0.2.7:4000" | ||
| 108 | r.AddCookie(&http.Cookie{Name: sessionCookie, Value: "sess"}) | ||
| 109 | if got := s.packPrincipal(r); got != want { | ||
| 110 | t.Fatalf("session: %q, want %q", got, want) | ||
| 111 | } | ||
| 112 | } | ||
| 113 | |||
| 114 | // stuckClient is a connection whose client sent the start of a request | ||
| 115 | // body and then stopped sending. Reads block until a read deadline is | ||
| 116 | // set; the response is discarded. | ||
| 117 | type stuckClient struct { | ||
| 118 | header http.Header | ||
| 119 | head string // the part of the body that was sent | ||
| 120 | cut chan struct{} | ||
| 121 | once sync.Once | ||
| 122 | } | ||
| 123 | |||
| 124 | func (c *stuckClient) Header() http.Header { return c.header } | ||
| 125 | func (c *stuckClient) WriteHeader(int) {} | ||
| 126 | func (c *stuckClient) Write(b []byte) (int, error) { return len(b), nil } | ||
| 127 | func (c *stuckClient) Read(b []byte) (int, error) { | ||
| 128 | if c.head != "" { | ||
| 129 | n := copy(b, c.head) | ||
| 130 | c.head = c.head[n:] | ||
| 131 | return n, nil | ||
| 132 | } | ||
| 133 | <-c.cut | ||
| 134 | return 0, os.ErrDeadlineExceeded | ||
| 135 | } | ||
| 136 | func (c *stuckClient) Close() error { return nil } | ||
| 137 | func (c *stuckClient) SetReadDeadline(time.Time) error { | ||
| 138 | c.once.Do(func() { close(c.cut) }) | ||
| 139 | return nil | ||
| 140 | } | ||
| 141 | |||
| 142 | // limitedServer is a server with a pack limit and an empty alice/app. | ||
| 143 | func limitedServer(t *testing.T) *Server { | ||
| 144 | t.Helper() | ||
| 145 | s := busyServer(t) | ||
| 146 | s.packs = packlimit.New(1, 0, 0, time.Second) | ||
| 147 | if err := gitutil.InitBare(control.RepoDir(s.cfg.Server.Root, "alice", "app"), "main", t.TempDir()); err != nil { | ||
| 148 | t.Fatal(err) | ||
| 149 | } | ||
| 150 | return s | ||
| 151 | } | ||
| 152 | |||
| 153 | // seed commits files of the given sizes, random and so incompressible, | ||
| 154 | // to alice/app's main and returns the commit. | ||
| 155 | func seed(t *testing.T, s *Server, sizes ...int) string { | ||
| 156 | t.Helper() | ||
| 157 | work := t.TempDir() | ||
| 158 | git := func(args ...string) string { | ||
| 159 | t.Helper() | ||
| 160 | cmd := exec.Command("git", append([]string{"-C", work, "-c", "user.name=t", "-c", "user.email=t@t"}, args...)...) | ||
| 161 | out, err := cmd.CombinedOutput() | ||
| 162 | if err != nil { | ||
| 163 | t.Fatalf("git %v: %v\n%s", args, err, out) | ||
| 164 | } | ||
| 165 | return strings.TrimSpace(string(out)) | ||
| 166 | } | ||
| 167 | git("init", "-q", "-b", "main") | ||
| 168 | for i, n := range sizes { | ||
| 169 | b := make([]byte, n) | ||
| 170 | rand.Read(b) | ||
| 171 | if err := os.WriteFile(filepath.Join(work, strconv.Itoa(i)), b, 0o644); err != nil { | ||
| 172 | t.Fatal(err) | ||
| 173 | } | ||
| 174 | } | ||
| 175 | git("add", ".") | ||
| 176 | git("commit", "-q", "-m", "seed") | ||
| 177 | git("push", "-q", control.RepoDir(s.cfg.Server.Root, "alice", "app"), "main") | ||
| 178 | return git("rev-parse", "HEAD") | ||
| 179 | } | ||
| 180 | |||
| 181 | // fetchBody is a protocol v0 request for sha's whole history. | ||
| 182 | func fetchBody(sha string) string { | ||
| 183 | pkt := func(s string) string { return fmt.Sprintf("%04x%s", len(s)+4, s) } | ||
| 184 | return pkt("want "+sha+" side-band-64k ofs-delta\n") + "0000" + pkt("done\n") | ||
| 185 | } | ||
| 186 | |||
| 187 | // ended requires the handler to finish within limit and its slot to be | ||
| 188 | // free. | ||
| 189 | func ended(t *testing.T, s *Server, finished <-chan struct{}, limit time.Duration) { | ||
| 190 | t.Helper() | ||
| 191 | select { | ||
| 192 | case <-finished: | ||
| 193 | case <-time.After(limit): | ||
| 194 | t.Fatal("fetch still running") | ||
| 195 | } | ||
| 196 | hold, err := s.packs.Acquire(nil, "ip:elsewhere") | ||
| 197 | if err != nil { | ||
| 198 | t.Fatalf("slot not released after the kill: %v", err) | ||
| 199 | } | ||
| 200 | hold() | ||
| 201 | } | ||
| 202 | |||
| 203 | func stallAfter(t *testing.T, d time.Duration) { | ||
| 204 | old := packlimit.StallDeadline | ||
| 205 | packlimit.StallDeadline = d | ||
| 206 | t.Cleanup(func() { packlimit.StallDeadline = old }) | ||
| 207 | } | ||
| 208 | |||
| 209 | // A client that stops sending its body is cut after StallDeadline. | ||
| 210 | func TestFetchKilledWhenClientStopsSending(t *testing.T) { | ||
| 211 | stallAfter(t, 200*time.Millisecond) | ||
| 212 | s := limitedServer(t) | ||
| 213 | // Half a pkt-line: git waits for the rest. | ||
| 214 | c := &stuckClient{header: http.Header{}, head: "0032want 0123456789abcdef", cut: make(chan struct{})} | ||
| 215 | r := httptest.NewRequest("POST", "/alice/app/git-upload-pack", c) | ||
| 216 | r.SetPathValue("owner", "alice") | ||
| 217 | r.SetPathValue("repo", "app") | ||
| 218 | finished := make(chan struct{}) | ||
| 219 | go func() { | ||
| 220 | s.uploadPack(c, r) | ||
| 221 | close(finished) | ||
| 222 | }() | ||
| 223 | ended(t, s, finished, 5*time.Second) | ||
| 224 | } | ||
| 225 | |||
| 226 | // A client that leaves ends git at once, even while git is busy with | ||
| 227 | // neither its input nor its output: here a pack-objects hook that | ||
| 228 | // sleeps, with the whole request read and the response discarded. | ||
| 229 | func TestFetchKilledWhenClientLeaves(t *testing.T) { | ||
| 230 | s := limitedServer(t) | ||
| 231 | sha := seed(t, s, 10) | ||
| 232 | hook := filepath.Join(t.TempDir(), "hook") | ||
| 233 | if err := os.WriteFile(hook, []byte("#!/bin/sh\nsleep 60\n"), 0o755); err != nil { | ||
| 234 | t.Fatal(err) | ||
| 235 | } | ||
| 236 | t.Setenv("GIT_CONFIG_COUNT", "1") | ||
| 237 | t.Setenv("GIT_CONFIG_KEY_0", "uploadpack.packObjectsHook") | ||
| 238 | t.Setenv("GIT_CONFIG_VALUE_0", hook) | ||
| 239 | ctx, cancel := context.WithCancel(context.Background()) | ||
| 240 | defer cancel() | ||
| 241 | r := httptest.NewRequestWithContext(ctx, "POST", "/alice/app/git-upload-pack", strings.NewReader(fetchBody(sha))) | ||
| 242 | r.SetPathValue("owner", "alice") | ||
| 243 | r.SetPathValue("repo", "app") | ||
| 244 | finished := make(chan struct{}) | ||
| 245 | go func() { | ||
| 246 | s.uploadPack(httptest.NewRecorder(), r) | ||
| 247 | close(finished) | ||
| 248 | }() | ||
| 249 | time.Sleep(500 * time.Millisecond) | ||
| 250 | cancel() | ||
| 251 | // Unkilled, git runs until the hook's sleep ends. | ||
| 252 | ended(t, s, finished, 5*time.Second) | ||
| 253 | } | ||
| 254 | |||
| 255 | // A client that stops reading the response is cut after StallDeadline: | ||
| 256 | // the write blocked on its full socket fails at the deadline set then. | ||
| 257 | func TestFetchKilledWhenClientStopsReading(t *testing.T) { | ||
| 258 | stallAfter(t, 500*time.Millisecond) | ||
| 259 | s := limitedServer(t) | ||
| 260 | sizes := make([]int, 16) | ||
| 261 | for i := range sizes { | ||
| 262 | sizes[i] = 2 << 20 | ||
| 263 | } | ||
| 264 | sha := seed(t, s, sizes...) | ||
| 265 | finished := make(chan struct{}) | ||
| 266 | srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { | ||
| 267 | r.SetPathValue("owner", "alice") | ||
| 268 | r.SetPathValue("repo", "app") | ||
| 269 | s.uploadPack(w, r) | ||
| 270 | close(finished) | ||
| 271 | })) | ||
| 272 | defer srv.Close() | ||
| 273 | conn, err := net.Dial("tcp", srv.Listener.Addr().String()) | ||
| 274 | if err != nil { | ||
| 275 | t.Fatal(err) | ||
| 276 | } | ||
| 277 | defer conn.Close() | ||
| 278 | conn.(*net.TCPConn).SetReadBuffer(4096) | ||
| 279 | body := fetchBody(sha) | ||
| 280 | fmt.Fprintf(conn, "POST /alice/app/git-upload-pack HTTP/1.1\r\nHost: x\r\nContent-Length: %d\r\n\r\n%s", len(body), body) | ||
| 281 | // The 32MB pack cannot fit in the socket buffers; nothing is read. | ||
| 282 | ended(t, s, finished, 20*time.Second) | ||
| 283 | } | ||
internal/httpd/routes_test.go +4 −4
| @@ -14,7 +14,7 @@ import ( | |||
| 14 | func TestViewOnlyHasNoMutatingRoutes(t *testing.T) { | 14 | func TestViewOnlyHasNoMutatingRoutes(t *testing.T) { |
| 15 | cfg := config.Default() | 15 | cfg := config.Default() |
| 16 | cfg.Web.Mode = "view_only" | 16 | cfg.Web.Mode = "view_only" |
| 17 | s := New(cfg, nil) | 17 | s := New(cfg, nil, nil) |
| 18 | 18 | ||
| 19 | for _, r := range s.Routes() { | 19 | for _, r := range s.Routes() { |
| 20 | if r.Mutating { | 20 | if r.Mutating { |
| @@ -38,7 +38,7 @@ func TestViewOnlyHasNoMutatingRoutes(t *testing.T) { | |||
| 38 | // TestAPIRouteGating: the API route exists only when [api] enabled = true. | 38 | // TestAPIRouteGating: the API route exists only when [api] enabled = true. |
| 39 | func TestAPIRouteGating(t *testing.T) { | 39 | func TestAPIRouteGating(t *testing.T) { |
| 40 | has := func(cfg config.Config) bool { | 40 | has := func(cfg config.Config) bool { |
| 41 | for _, r := range New(cfg, nil).Routes() { | 41 | for _, r := range New(cfg, nil, nil).Routes() { |
| 42 | if r.Pattern == "/api/v1/cmd" { | 42 | if r.Pattern == "/api/v1/cmd" { |
| 43 | return true | 43 | return true |
| 44 | } | 44 | } |
| @@ -60,7 +60,7 @@ func TestAPIRouteGating(t *testing.T) { | |||
| 60 | func TestAccountsModeHasLoginRoute(t *testing.T) { | 60 | func TestAccountsModeHasLoginRoute(t *testing.T) { |
| 61 | cfg := config.Default() | 61 | cfg := config.Default() |
| 62 | cfg.Web.Mode = "accounts" | 62 | cfg.Web.Mode = "accounts" |
| 63 | s := New(cfg, nil) | 63 | s := New(cfg, nil, nil) |
| 64 | found := false | 64 | found := false |
| 65 | for _, r := range s.Routes() { | 65 | for _, r := range s.Routes() { |
| 66 | if r.Pattern == "/login" { | 66 | if r.Pattern == "/login" { |
| @@ -82,7 +82,7 @@ func TestTopLevelRouteWordsAreReserved(t *testing.T) { | |||
| 82 | // (config.Default() leaves it "closed"); open it so this walk actually | 82 | // (config.Default() leaves it "closed"); open it so this walk actually |
| 83 | // reaches the route the production instance runs with. | 83 | // reaches the route the production instance runs with. |
| 84 | cfg.Registration.Mode = "open" | 84 | cfg.Registration.Mode = "open" |
| 85 | s := New(cfg, nil) | 85 | s := New(cfg, nil, nil) |
| 86 | for _, r := range s.Routes() { | 86 | for _, r := range s.Routes() { |
| 87 | seg := strings.TrimPrefix(r.Pattern, "/") | 87 | seg := strings.TrimPrefix(r.Pattern, "/") |
| 88 | seg, _, _ = strings.Cut(seg, "/") | 88 | seg, _, _ = strings.Cut(seg, "/") |
internal/httpd/smart.go +88 −6
| @@ -6,18 +6,24 @@ | |||
| 6 | package httpd | 6 | package httpd |
| 7 | 7 | ||
| 8 | import ( | 8 | import ( |
| 9 | "bufio" | ||
| 9 | "compress/gzip" | 10 | "compress/gzip" |
| 11 | "errors" | ||
| 10 | "fmt" | 12 | "fmt" |
| 11 | "io" | 13 | "io" |
| 12 | "net" | 14 | "net" |
| 13 | "net/http" | 15 | "net/http" |
| 14 | "os" | 16 | "os" |
| 15 | "os/exec" | 17 | "os/exec" |
| 18 | "strconv" | ||
| 16 | "strings" | 19 | "strings" |
| 17 | "sync" | 20 | "sync" |
| 21 | "time" | ||
| 18 | 22 | ||
| 19 | "gitbay.org/gitbay/internal/config" | 23 | "gitbay.org/gitbay/internal/config" |
| 20 | "gitbay.org/gitbay/internal/control" | 24 | "gitbay.org/gitbay/internal/control" |
| 25 | "gitbay.org/gitbay/internal/gitutil" | ||
| 26 | "gitbay.org/gitbay/internal/packlimit" | ||
| 21 | "gitbay.org/gitbay/internal/store" | 27 | "gitbay.org/gitbay/internal/store" |
| 22 | "gitbay.org/gitbay/internal/toolpath" | 28 | "gitbay.org/gitbay/internal/toolpath" |
| 23 | ) | 29 | ) |
| @@ -25,15 +31,16 @@ import ( | |||
| 25 | type Server struct { | 31 | type Server struct { |
| 26 | cfg config.Config | 32 | cfg config.Config |
| 27 | st *store.Store | 33 | st *store.Store |
| 34 | packs *packlimit.Limiter | ||
| 28 | apiLimit *apiLimiter | 35 | apiLimit *apiLimiter |
| 29 | proxies []*net.IPNet // http.trusted_proxies, parsed once | 36 | proxies []*net.IPNet // http.trusted_proxies, parsed once |
| 30 | stopping chan struct{} // closed by Stop | 37 | stopping chan struct{} // closed by Stop |
| 31 | stopOnce sync.Once | 38 | stopOnce sync.Once |
| 32 | } | 39 | } |
| 33 | 40 | ||
| 34 | func New(cfg config.Config, st *store.Store) *Server { | 41 | func New(cfg config.Config, st *store.Store, packs *packlimit.Limiter) *Server { |
| 35 | proxies, _ := cfg.HTTP.TrustedProxyNets() // validated at config load | 42 | proxies, _ := cfg.HTTP.TrustedProxyNets() // validated at config load |
| 36 | return &Server{cfg: cfg, st: st, apiLimit: newAPILimiter(cfg.Limits.APIRate), proxies: proxies, | 43 | return &Server{cfg: cfg, st: st, packs: packs, apiLimit: newAPILimiter(cfg.Limits.APIRate), proxies: proxies, |
| 37 | stopping: make(chan struct{})} | 44 | stopping: make(chan struct{})} |
| 38 | } | 45 | } |
| 39 | 46 | ||
| @@ -135,14 +142,89 @@ func (s *Server) uploadPack(w http.ResponseWriter, r *http.Request) { | |||
| 135 | defer gz.Close() | 142 | defer gz.Close() |
| 136 | body = gz | 143 | body = gz |
| 137 | } | 144 | } |
| 145 | br := bufio.NewReader(body) | ||
| 146 | cancel := r.Context().Done() | ||
| 147 | out := io.Writer(w) | ||
| 148 | if !lsRefs(br) { | ||
| 149 | // A queued clone waits at most the limiter's wait, and Stop ends | ||
| 150 | // the wait so it does not hold up a restart's drain. net/http | ||
| 151 | // notices a departed client only after the body is read, so | ||
| 152 | // that rarely ends it. | ||
| 153 | release, err := s.packs.Acquire(s.until(r), s.packPrincipal(r)) | ||
| 154 | if err != nil { | ||
| 155 | msg := "the server is restarting; try again in a minute" | ||
| 156 | if errors.Is(err, packlimit.ErrBusy) { | ||
| 157 | msg = "the server is busy: it is at its limit of concurrent clones and fetches; try again in a minute" | ||
| 158 | } | ||
| 159 | w.Header().Set("Retry-After", "30") | ||
| 160 | http.Error(w, msg, http.StatusServiceUnavailable) | ||
| 161 | return | ||
| 162 | } | ||
| 163 | // Deferred before git runs, so it fires after git has exited | ||
| 164 | // and been waited for. | ||
| 165 | defer release() | ||
| 166 | var stalled <-chan struct{} | ||
| 167 | var unwatch func() | ||
| 168 | out, stalled, unwatch = s.packs.Watch(w) | ||
| 169 | defer unwatch() | ||
| 170 | kill := make(chan struct{}) | ||
| 171 | finished := make(chan struct{}) | ||
| 172 | exited := make(chan struct{}) | ||
| 173 | // The watcher must not touch w once the handler has returned. | ||
| 174 | defer func() { | ||
| 175 | close(finished) | ||
| 176 | <-exited | ||
| 177 | }() | ||
| 178 | go func() { | ||
| 179 | defer close(exited) | ||
| 180 | select { | ||
| 181 | case <-finished: | ||
| 182 | return | ||
| 183 | case <-r.Context().Done(): | ||
| 184 | case <-stalled: | ||
| 185 | // A write blocked on a client that stopped reading, or a | ||
| 186 | // read of a body it stopped sending, outlives git; | ||
| 187 | // expired deadlines end both copies, so Wait returns. | ||
| 188 | rc := http.NewResponseController(w) | ||
| 189 | rc.SetReadDeadline(time.Now()) | ||
| 190 | rc.SetWriteDeadline(time.Now()) | ||
| 191 | } | ||
| 192 | close(kill) | ||
| 193 | }() | ||
| 194 | cancel = kill | ||
| 195 | } | ||
| 138 | w.Header().Set("Content-Type", "application/x-git-upload-pack-result") | 196 | w.Header().Set("Content-Type", "application/x-git-upload-pack-result") |
| 139 | w.Header().Set("Cache-Control", "no-cache") | 197 | w.Header().Set("Cache-Control", "no-cache") |
| 140 | dir := control.RepoDir(s.cfg.Server.Root, repo.OwnerName, repo.Name) | 198 | dir := control.RepoDir(s.cfg.Server.Root, repo.OwnerName, repo.Name) |
| 141 | cmd := exec.CommandContext(r.Context(), toolpath.Look("git"), "upload-pack", "--stateless-rpc", dir) | 199 | cmd := exec.Command(toolpath.Look("git"), "-c", "uploadpack.keepAlive=5", "upload-pack", "--stateless-rpc", dir) |
| 142 | cmd.Env = append(os.Environ(), gitProtocolEnv(r)...) | 200 | cmd.Env = append(os.Environ(), gitProtocolEnv(r)...) |
| 143 | cmd.Stdin = body | 201 | cmd.Stdin = br |
| 144 | cmd.Stdout = w | 202 | cmd.Stdout = out |
| 145 | cmd.Run() | 203 | gitutil.RunUntil(cmd, cancel) |
| 204 | } | ||
| 205 | |||
| 206 | // lsRefs reports whether a protocol v2 request is a ref listing, which | ||
| 207 | // generates no pack. Its first pkt-line is "command=ls-refs". | ||
| 208 | func lsRefs(br *bufio.Reader) bool { | ||
| 209 | const want = "command=ls-refs" | ||
| 210 | head, err := br.Peek(4 + len(want)) | ||
| 211 | return err == nil && string(head[4:]) == want | ||
| 212 | } | ||
| 213 | |||
| 214 | // packPrincipal is who a fetch is counted against: the account when the | ||
| 215 | // request carries a valid bearer token or web session, the same key SSH | ||
| 216 | // uses, so switching transport buys no extra slots; otherwise the | ||
| 217 | // client address. | ||
| 218 | func (s *Server) packPrincipal(r *http.Request) string { | ||
| 219 | if tok, ok := strings.CutPrefix(r.Header.Get("Authorization"), "Bearer "); ok && strings.TrimSpace(tok) != "" { | ||
| 220 | if u, _, err := s.st.APITokenUser(store.HashToken(strings.TrimSpace(tok))); err == nil { | ||
| 221 | return "user:" + strconv.FormatInt(u.ID, 10) | ||
| 222 | } | ||
| 223 | } | ||
| 224 | if u := s.viewer(r); u.ID != 0 { | ||
| 225 | return "user:" + strconv.FormatInt(u.ID, 10) | ||
| 226 | } | ||
| 227 | return "ip:" + s.clientIP(r) | ||
| 146 | } | 228 | } |
| 147 | 229 | ||
| 148 | // gitProtocolEnv forwards the client's protocol negotiation header so | 230 | // gitProtocolEnv forwards the client's protocol negotiation header so |
internal/packlimit/watch.go added +61
| @@ -0,0 +1,61 @@ | |||
| 1 | package packlimit | ||
| 2 | |||
| 3 | import ( | ||
| 4 | "io" | ||
| 5 | "sync" | ||
| 6 | "sync/atomic" | ||
| 7 | "time" | ||
| 8 | ) | ||
| 9 | |||
| 10 | // StallDeadline is how long a limited transport may go without | ||
| 11 | // completing a write to its client before it is killed. upload-pack | ||
| 12 | // sends a keepalive every five seconds while it prepares a pack. | ||
| 13 | var StallDeadline = 2 * time.Minute | ||
| 14 | |||
| 15 | // Watch wraps w, a transport's writer to its client, so that a client | ||
| 16 | // that stops reading does not hold its slot for as long as its | ||
| 17 | // connection stays open. stalled closes once no write to w has | ||
| 18 | // completed for StallDeadline; stop ends the watch. With no limit in | ||
| 19 | // force (a nil Limiter) nothing is watched: w comes back as is and | ||
| 20 | // stalled never closes. | ||
| 21 | func (l *Limiter) Watch(w io.Writer) (out io.Writer, stalled <-chan struct{}, stop func()) { | ||
| 22 | if l == nil { | ||
| 23 | return w, nil, func() {} | ||
| 24 | } | ||
| 25 | deadline := StallDeadline | ||
| 26 | pw := &progressWriter{w: w} | ||
| 27 | pw.last.Store(time.Now().UnixNano()) | ||
| 28 | st := make(chan struct{}) | ||
| 29 | quit := make(chan struct{}) | ||
| 30 | go func() { | ||
| 31 | t := time.NewTicker(deadline / 4) | ||
| 32 | defer t.Stop() | ||
| 33 | for { | ||
| 34 | select { | ||
| 35 | case <-quit: | ||
| 36 | return | ||
| 37 | case <-t.C: | ||
| 38 | if time.Since(time.Unix(0, pw.last.Load())) >= deadline { | ||
| 39 | close(st) | ||
| 40 | return | ||
| 41 | } | ||
| 42 | } | ||
| 43 | } | ||
| 44 | }() | ||
| 45 | var once sync.Once | ||
| 46 | return pw, st, func() { once.Do(func() { close(quit) }) } | ||
| 47 | } | ||
| 48 | |||
| 49 | // progressWriter records when a write to the client last completed. | ||
| 50 | type progressWriter struct { | ||
| 51 | w io.Writer | ||
| 52 | last atomic.Int64 // unix nanoseconds | ||
| 53 | } | ||
| 54 | |||
| 55 | func (p *progressWriter) Write(b []byte) (int, error) { | ||
| 56 | n, err := p.w.Write(b) | ||
| 57 | if n > 0 { | ||
| 58 | p.last.Store(time.Now().UnixNano()) | ||
| 59 | } | ||
| 60 | return n, err | ||
| 61 | } | ||
internal/packlimit/watch_test.go added +40
| @@ -0,0 +1,40 @@ | |||
| 1 | package packlimit | ||
| 2 | |||
| 3 | import ( | ||
| 4 | "io" | ||
| 5 | "testing" | ||
| 6 | "time" | ||
| 7 | ) | ||
| 8 | |||
| 9 | func TestWatchStalls(t *testing.T) { | ||
| 10 | old := StallDeadline | ||
| 11 | StallDeadline = 100 * time.Millisecond | ||
| 12 | t.Cleanup(func() { StallDeadline = old }) | ||
| 13 | l := New(1, 0, 0, time.Second) | ||
| 14 | w, stalled, stop := l.Watch(io.Discard) | ||
| 15 | defer stop() | ||
| 16 | // Writes keep it alive past the deadline. | ||
| 17 | for range 6 { | ||
| 18 | time.Sleep(40 * time.Millisecond) | ||
| 19 | w.Write([]byte("x")) | ||
| 20 | select { | ||
| 21 | case <-stalled: | ||
| 22 | t.Fatal("stalled while writing") | ||
| 23 | default: | ||
| 24 | } | ||
| 25 | } | ||
| 26 | select { | ||
| 27 | case <-stalled: | ||
| 28 | case <-time.After(2 * time.Second): | ||
| 29 | t.Fatal("no stall after writes stopped") | ||
| 30 | } | ||
| 31 | } | ||
| 32 | |||
| 33 | func TestWatchWithoutLimit(t *testing.T) { | ||
| 34 | var l *Limiter | ||
| 35 | w, stalled, stop := l.Watch(io.Discard) | ||
| 36 | defer stop() | ||
| 37 | if w != io.Discard || stalled != nil { | ||
| 38 | t.Fatal("a nil limiter watched") | ||
| 39 | } | ||
| 40 | } | ||
internal/sshd/refusal_test.go +4 −4
| @@ -192,11 +192,11 @@ func TestCloneKilledWhenKeyRevoked(t *testing.T) { | |||
| 192 | killedClone(t, io.Discard, nil, closed(), closed()) | 192 | killedClone(t, io.Discard, nil, closed(), closed()) |
| 193 | } | 193 | } |
| 194 | 194 | ||
| 195 | // A client that stops reading is cut after stallDeadline. | 195 | // A client that stops reading is cut after packlimit.StallDeadline. |
| 196 | func TestCloneKilledWhenClientStopsReading(t *testing.T) { | 196 | func TestCloneKilledWhenClientStopsReading(t *testing.T) { |
| 197 | old := stallDeadline | 197 | old := packlimit.StallDeadline |
| 198 | stallDeadline = 200 * time.Millisecond | 198 | packlimit.StallDeadline = 200 * time.Millisecond |
| 199 | t.Cleanup(func() { stallDeadline = old }) | 199 | t.Cleanup(func() { packlimit.StallDeadline = old }) |
| 200 | r, w := io.Pipe() | 200 | r, w := io.Pipe() |
| 201 | t.Cleanup(func() { r.Close() }) | 201 | t.Cleanup(func() { r.Close() }) |
| 202 | killedClone(t, w, nil, nil, nil) | 202 | killedClone(t, w, nil, nil, nil) |
internal/sshd/sshd.go +8 −36
| @@ -516,25 +516,6 @@ func Exec(cfg config.Config, st *store.Store, packs *packlimit.Limiter, user sto | |||
| 516 | return control.Dispatch(ctx, argv) | 516 | return control.Dispatch(ctx, argv) |
| 517 | } | 517 | } |
| 518 | 518 | ||
| 519 | // stallDeadline is how long a limited transport may go without writing | ||
| 520 | // to its client before it is killed. upload-pack sends a keepalive | ||
| 521 | // every five seconds while it prepares a pack. | ||
| 522 | var stallDeadline = 2 * time.Minute | ||
| 523 | |||
| 524 | // progressWriter records when a write to the client last completed. | ||
| 525 | type progressWriter struct { | ||
| 526 | w io.Writer | ||
| 527 | last atomic.Int64 // unix nanoseconds | ||
| 528 | } | ||
| 529 | |||
| 530 | func (p *progressWriter) Write(b []byte) (int, error) { | ||
| 531 | n, err := p.w.Write(b) | ||
| 532 | if n > 0 { | ||
| 533 | p.last.Store(time.Now().UnixNano()) | ||
| 534 | } | ||
| 535 | return n, err | ||
| 536 | } | ||
| 537 | |||
| 538 | // runGit streams a git transport service after access checks. | 519 | // runGit streams a git transport service after access checks. |
| 539 | func runGit(cfg config.Config, st *store.Store, packs *packlimit.Limiter, user store.User, scope string, argv []string, | 520 | func runGit(cfg config.Config, st *store.Store, packs *packlimit.Limiter, user store.User, scope string, argv []string, |
| 540 | stdin io.Reader, stdout, stderr io.Writer, done, stopping, revoked <-chan struct{}) int { | 521 | stdin io.Reader, stdout, stderr io.Writer, done, stopping, revoked <-chan struct{}) int { |
| @@ -641,18 +622,12 @@ func runGit(cfg config.Config, st *store.Store, packs *packlimit.Limiter, user s | |||
| 641 | // exited and been waited for. | 622 | // exited and been waited for. |
| 642 | defer release() | 623 | defer release() |
| 643 | // A client that stops reading would hold its slot for as long | 624 | // A client that stops reading would hold its slot for as long |
| 644 | // as its channel stays open. With a limit in force, a transport | 625 | // as its channel stays open. |
| 645 | // that writes nothing for stallDeadline is killed. | 626 | client := stdout |
| 646 | var pw *progressWriter | 627 | var stalled <-chan struct{} |
| 647 | var tick <-chan time.Time | 628 | var unwatch func() |
| 648 | if packs != nil { | 629 | stdout, stalled, unwatch = packs.Watch(client) |
| 649 | pw = &progressWriter{w: stdout} | 630 | defer unwatch() |
| 650 | pw.last.Store(time.Now().UnixNano()) | ||
| 651 | stdout = pw | ||
| 652 | t := time.NewTicker(stallDeadline / 4) | ||
| 653 | defer t.Stop() | ||
| 654 | tick = t.C | ||
| 655 | } | ||
| 656 | kill := make(chan struct{}) | 631 | kill := make(chan struct{}) |
| 657 | finished := make(chan struct{}) | 632 | finished := make(chan struct{}) |
| 658 | defer close(finished) | 633 | defer close(finished) |
| @@ -673,15 +648,12 @@ func runGit(cfg config.Config, st *store.Store, packs *packlimit.Limiter, user s | |||
| 673 | continue | 648 | continue |
| 674 | default: | 649 | default: |
| 675 | } | 650 | } |
| 676 | case <-tick: | 651 | case <-stalled: |
| 677 | if time.Since(time.Unix(0, pw.last.Load())) < stallDeadline { | ||
| 678 | continue | ||
| 679 | } | ||
| 680 | close(kill) | 652 | close(kill) |
| 681 | // A write blocked on the client's window outlives | 653 | // A write blocked on the client's window outlives |
| 682 | // git; closing the channel ends it and the stdin copy, | 654 | // git; closing the channel ends it and the stdin copy, |
| 683 | // so Transport's Wait returns. | 655 | // so Transport's Wait returns. |
| 684 | if c, ok := pw.w.(io.Closer); ok { | 656 | if c, ok := client.(io.Closer); ok { |
| 685 | c.Close() | 657 | c.Close() |
| 686 | } | 658 | } |
| 687 | return | 659 | return |