bb40f89b34de83368d63742a9244c69554106cfb

Author
Ayman Bagabas <ayman.bagabas@gmail.com>
Committer
Ayman Bagabas <ayman.bagabas@gmail.com>
Date

Message

fix(server): bound the current context to the underlying operation

Use the connection context when running external commands.

Signed-off-by: Ayman Bagabas <ayman.bagabas@gmail.com>

Diff

  1diff --git a/go.mod b/go.mod
  2index 7c1284dc6f3e7906634774fdd4b3a3d5b6b7fbe6..fdba6eca0e06a1610de36770bdafb18f4fd1f671 100644
  3--- a/go.mod
  4+++ b/go.mod
  5@@ -93,3 +93,5 @@ require (
  6 	modernc.org/strutil v1.1.3 // indirect
  7 	modernc.org/token v1.0.1 // indirect
  8 )
  9+
 10+replace github.com/gogs/git-module => github.com/aymanbagabas/git-module v1.4.1-0.20230509180555-975c24cdb79a
 11diff --git a/go.sum b/go.sum
 12index 1e1a15a0a449e6b67a9bd7446a47978657adf1ec..6aaef95e079c6829dcc96c4cf33e0b9a1eaec7e0 100644
 13--- a/go.sum
 14+++ b/go.sum
 15@@ -51,6 +51,8 @@ github.com/anmitsu/go-shlex v0.0.0-20200514113438-38f4b401e2be/go.mod h1:ySMOLuW
 16 github.com/armon/go-socks5 v0.0.0-20160902184237-e75332964ef5/go.mod h1:wHh0iHkYZB8zMSxRWpUBQtwG5a7fFgvEO+odwuTv2gs=
 17 github.com/atotto/clipboard v0.1.4 h1:EH0zSVneZPSuFR11BlR9YppQTVDbh5+16AmcJi4g1z4=
 18 github.com/atotto/clipboard v0.1.4/go.mod h1:ZY9tmq7sm5xIbd9bOK4onWV4S6X0u6GY7Vn0Yu86PYI=
 19+github.com/aymanbagabas/git-module v1.4.1-0.20230509180555-975c24cdb79a h1:rY724fIR0NNR/UTXJufwKwYz+sNYVve/ZdzWX39xMqM=
 20+github.com/aymanbagabas/git-module v1.4.1-0.20230509180555-975c24cdb79a/go.mod h1:GUSSUH+RM7fZOtjhS6Obh4B9aAvs3EeROpazfMNMF8g=
 21 github.com/aymanbagabas/go-osc52 v1.0.3/go.mod h1:zT8H+Rk4VSabYN90pWyugflM3ZhpTZNC7cASDfUCdT4=
 22 github.com/aymanbagabas/go-osc52 v1.2.1 h1:q2sWUyDcozPLcLabEMd+a+7Ea2DitxZVN9hTxab9L4E=
 23 github.com/aymanbagabas/go-osc52 v1.2.1/go.mod h1:zT8H+Rk4VSabYN90pWyugflM3ZhpTZNC7cASDfUCdT4=
 24@@ -148,8 +150,6 @@ github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/me
 25 github.com/gobwas/glob v0.2.3 h1:A4xDbljILXROh+kObIiy5kIaPYD8e96x1tgBhUI5J+Y=
 26 github.com/gobwas/glob v0.2.3/go.mod h1:d3Ez4x06l9bZtSvzIay5+Yzi0fmZzPgnTbPcKjJAkT8=
 27 github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ=
 28-github.com/gogs/git-module v1.8.1 h1:yC5BZ3unJOXC8N6/FgGQ8EtJXpOd217lgDcd2aPOxkc=
 29-github.com/gogs/git-module v1.8.1/go.mod h1:Y3rsSqtFZEbn7lp+3gWf42GKIY1eNTtLt7JrmOy0yAQ=
 30 github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q=
 31 github.com/golang/groupcache v0.0.0-20190702054246-869f871628b6/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
 32 github.com/golang/groupcache v0.0.0-20191227052852-215e87163ea7/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
 33@@ -392,7 +392,6 @@ github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81P
 34 github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
 35 github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
 36 github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
 37-github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
 38 github.com/stretchr/testify v1.8.2 h1:+h33VjcLVPDHtOdpUCuF+7gSuG3yGIftsP1YvFihtJ8=
 39 github.com/stretchr/testify v1.8.2/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
 40 github.com/xanzy/ssh-agent v0.3.3/go.mod h1:6dzNDKs0J9rVPHPhaGCukekBHKqfl+L3KghI1Bc68Uw=
 41diff --git a/server/backend/backend.go b/server/backend/backend.go
 42index 308f342794b03cc0d416c0530d5529b915b16695..1aa408b46a5bf92cf54f4d4dc676eceffabf8e6a 100644
 43--- a/server/backend/backend.go
 44+++ b/server/backend/backend.go
 45@@ -2,6 +2,7 @@ package backend
 46 
 47 import (
 48 	"bytes"
 49+	"context"
 50 
 51 	"github.com/charmbracelet/ssh"
 52 	gossh "golang.org/x/crypto/ssh"
 53@@ -17,6 +18,9 @@ type Backend interface {
 54 	UserStore
 55 	UserAccess
 56 	Hooks
 57+
 58+	// WithContext returns a copy Backend with the given context.
 59+	WithContext(ctx context.Context) Backend
 60 }
 61 
 62 // ParseAuthorizedKey parses an authorized key string into a public key.
 63diff --git a/server/backend/context.go b/server/backend/context.go
 64new file mode 100644
 65index 0000000000000000000000000000000000000000..1af19057d57ddf2c61bba91460d996af66884ec0
 66--- /dev/null
 67+++ b/server/backend/context.go
 68@@ -0,0 +1,19 @@
 69+package backend
 70+
 71+import "context"
 72+
 73+var contextKey = &struct{ string }{"backend"}
 74+
 75+// FromContext returns the backend from a context.
 76+func FromContext(ctx context.Context) Backend {
 77+	if b, ok := ctx.Value(contextKey).(Backend); ok {
 78+		return b
 79+	}
 80+
 81+	return nil
 82+}
 83+
 84+// WithContext returns a new context with the backend attached.
 85+func WithContext(ctx context.Context, b Backend) context.Context {
 86+	return context.WithValue(ctx, contextKey, b)
 87+}
 88diff --git a/server/backend/sqlite/sqlite.go b/server/backend/sqlite/sqlite.go
 89index 1581d5939d1c72d315bf4ad3a411f8cf55ff9b23..25d396a352f1911949db410f96182acf5351888b 100644
 90--- a/server/backend/sqlite/sqlite.go
 91+++ b/server/backend/sqlite/sqlite.go
 92@@ -2,6 +2,7 @@ package sqlite
 93 
 94 import (
 95 	"context"
 96+	"errors"
 97 	"fmt"
 98 	"os"
 99 	"path/filepath"
100@@ -67,6 +68,12 @@ func NewSqliteBackend(ctx context.Context) (*SqliteBackend, error) {
101 	return d, d.initRepos()
102 }
103 
104+// WithContext returns a copy of SqliteBackend with the given context.
105+func (d SqliteBackend) WithContext(ctx context.Context) backend.Backend {
106+	d.ctx = ctx
107+	return &d
108+}
109+
110 // AllowKeyless returns whether or not keyless access is allowed.
111 //
112 // It implements backend.Backend.
113@@ -183,6 +190,8 @@ func (d *SqliteBackend) ImportRepository(name string, remote string, opts backen
114 		Quiet:   true,
115 		Timeout: 15 * time.Minute,
116 		CommandOptions: git.CommandOptions{
117+			Timeout: -1,
118+			Context: d.ctx,
119 			Envs: []string{
120 				fmt.Sprintf(`GIT_SSH_COMMAND=ssh -o UserKnownHostsFile="%s" -o StrictHostKeyChecking=no -i "%s"`,
121 					filepath.Join(d.cfg.DataPath, "ssh", "known_hosts"),
122@@ -194,6 +203,9 @@ func (d *SqliteBackend) ImportRepository(name string, remote string, opts backen
123 
124 	if err := git.Clone(remote, rp, copts); err != nil {
125 		d.logger.Error("failed to clone repository", "err", err, "mirror", opts.Mirror, "remote", remote, "path", rp)
126+		if rerr := os.RemoveAll(rp); rerr != nil {
127+			err = errors.Join(err, rerr)
128+		}
129 		return nil, err
130 	}
131 
132diff --git a/server/cmd/cmd.go b/server/cmd/cmd.go
133index a1d36d893bdf8b9f9e3a00ad6f77561c5328db36..c30dd8b29c69a3cc6a284b4655cb1a78699a39ac 100644
134--- a/server/cmd/cmd.go
135+++ b/server/cmd/cmd.go
136@@ -18,25 +18,9 @@ import (
137 	"github.com/spf13/cobra"
138 )
139 
140-// ContextKey is a type that can be used as a key in a context.
141-type ContextKey string
142-
143-// String returns the string representation of the ContextKey.
144-func (c ContextKey) String() string {
145-	return string(c) + "ContextKey"
146-}
147-
148-var (
149-	// ConfigCtxKey is the key for the config in the context.
150-	ConfigCtxKey = ContextKey("config")
151-	// SessionCtxKey is the key for the session in the context.
152-	SessionCtxKey = ContextKey("session")
153-	// HooksCtxKey is the key for the git hooks in the context.
154-	HooksCtxKey = ContextKey("hooks")
155-)
156-
157 var (
158-	logger = log.WithPrefix("server.cmd")
159+	// sessionCtxKey is the key for the session in the context.
160+	sessionCtxKey = &struct{ string }{"session"}
161 )
162 
163 var templateFuncs = template.FuncMap{
164@@ -152,8 +136,8 @@ func rootCommand(cfg *config.Config, s ssh.Session) *cobra.Command {
165 
166 func fromContext(cmd *cobra.Command) (*config.Config, ssh.Session) {
167 	ctx := cmd.Context()
168-	cfg := ctx.Value(ConfigCtxKey).(*config.Config)
169-	s := ctx.Value(SessionCtxKey).(ssh.Session)
170+	cfg := config.FromContext(ctx)
171+	s := ctx.Value(sessionCtxKey).(ssh.Session)
172 	return cfg, s
173 }
174 
175@@ -213,7 +197,7 @@ func checkIfCollab(cmd *cobra.Command, args []string) error {
176 }
177 
178 // Middleware is the Soft Serve middleware that handles SSH commands.
179-func Middleware(cfg *config.Config) wish.Middleware {
180+func Middleware(cfg *config.Config, logger *log.Logger) wish.Middleware {
181 	return func(sh ssh.Handler) ssh.Handler {
182 		return func(s ssh.Session) {
183 			func() {
184@@ -232,8 +216,16 @@ func Middleware(cfg *config.Config) wish.Middleware {
185 					}
186 				}
187 
188-				ctx := context.WithValue(s.Context(), ConfigCtxKey, cfg)
189-				ctx = context.WithValue(ctx, SessionCtxKey, s)
190+				// Here we copy the server's config and replace the backend
191+				// with a new one that uses the session's context.
192+				var ctx context.Context = s.Context()
193+				scfg := *cfg
194+				cfg = &scfg
195+				be := cfg.Backend.WithContext(ctx)
196+				cfg.Backend = be
197+				ctx = config.WithContext(ctx, cfg)
198+				ctx = backend.WithContext(ctx, be)
199+				ctx = context.WithValue(ctx, sessionCtxKey, s)
200 
201 				rootCmd := rootCommand(cfg, s)
202 				rootCmd.SetArgs(args)
203diff --git a/server/daemon/daemon.go b/server/daemon/daemon.go
204index cf71ae6730ac445fa8122160c86c401edd79cba7..944f0eedb9c955dfd1914f3236541b28a5ac38a2 100644
205--- a/server/daemon/daemon.go
206+++ b/server/daemon/daemon.go
207@@ -83,6 +83,7 @@ type GitDaemon struct {
208 	finished chan struct{}
209 	conns    connections
210 	cfg      *config.Config
211+	be       backend.Backend
212 	wg       sync.WaitGroup
213 	once     sync.Once
214 	logger   *log.Logger
215@@ -97,6 +98,7 @@ func NewGitDaemon(ctx context.Context) (*GitDaemon, error) {
216 		addr:     addr,
217 		finished: make(chan struct{}, 1),
218 		cfg:      cfg,
219+		be:       backend.FromContext(ctx),
220 		conns:    connections{m: make(map[net.Conn]struct{})},
221 		logger:   log.FromContext(ctx).WithPrefix("gitdaemon"),
222 	}
223@@ -231,7 +233,8 @@ func (d *GitDaemon) handleClient(conn net.Conn) {
224 			return
225 		}
226 
227-		if !d.cfg.Backend.AllowKeyless() {
228+		be := d.be.WithContext(ctx)
229+		if !be.AllowKeyless() {
230 			d.fatal(c, git.ErrNotAuthed)
231 			return
232 		}
233@@ -248,7 +251,7 @@ func (d *GitDaemon) handleClient(conn net.Conn) {
234 			return
235 		}
236 
237-		auth := d.cfg.Backend.AccessLevel(name, "")
238+		auth := be.AccessLevel(name, "")
239 		if auth < backend.ReadOnlyAccess {
240 			d.fatal(c, git.ErrNotAuthed)
241 			return
242diff --git a/server/daemon/daemon_test.go b/server/daemon/daemon_test.go
243index b06decbf4724ad3ab3a62e6632bddbf32a0078bd..a28fb68d4404d5e3eaa9e50535a984874275de7e 100644
244--- a/server/daemon/daemon_test.go
245+++ b/server/daemon/daemon_test.go
246@@ -12,6 +12,7 @@ import (
247 	"strings"
248 	"testing"
249 
250+	"github.com/charmbracelet/soft-serve/server/backend"
251 	"github.com/charmbracelet/soft-serve/server/backend/sqlite"
252 	"github.com/charmbracelet/soft-serve/server/config"
253 	"github.com/charmbracelet/soft-serve/server/git"
254@@ -35,15 +36,16 @@ func TestMain(m *testing.M) {
255 	ctx := context.TODO()
256 	cfg := config.DefaultConfig()
257 	ctx = config.WithContext(ctx, cfg)
258-	d, err := NewGitDaemon(ctx)
259+	fb, err := sqlite.NewSqliteBackend(ctx)
260 	if err != nil {
261 		log.Fatal(err)
262 	}
263-	fb, err := sqlite.NewSqliteBackend(ctx)
264+	cfg = cfg.WithBackend(fb)
265+	ctx = backend.WithContext(ctx, fb)
266+	d, err := NewGitDaemon(ctx)
267 	if err != nil {
268 		log.Fatal(err)
269 	}
270-	cfg = cfg.WithBackend(fb)
271 	testDaemon = d
272 	go func() {
273 		if err := d.Start(); err != ErrServerClosed {
274diff --git a/server/server.go b/server/server.go
275index 66a5488fb0c971c44f1eaec00b9276499733485d..3b9554b2bd83069aa890348fd67df8fea2d232c3 100644
276--- a/server/server.go
277+++ b/server/server.go
278@@ -50,6 +50,7 @@ func NewServer(ctx context.Context) (*Server, error) {
279 		}
280 
281 		cfg = cfg.WithBackend(sb)
282+		ctx = backend.WithContext(ctx, sb)
283 	}
284 
285 	srv := &Server{
286diff --git a/server/ssh/ssh.go b/server/ssh/ssh.go
287index 42e51222602770fdc056004cab503986b5dbc76e..a962721d67d19636d31477aa7fb009097c001b45 100644
288--- a/server/ssh/ssh.go
289+++ b/server/ssh/ssh.go
290@@ -77,6 +77,7 @@ var (
291 type SSHServer struct {
292 	srv    *ssh.Server
293 	cfg    *config.Config
294+	be     backend.Backend
295 	ctx    context.Context
296 	logger *log.Logger
297 }
298@@ -84,26 +85,28 @@ type SSHServer struct {
299 // NewSSHServer returns a new SSHServer.
300 func NewSSHServer(ctx context.Context) (*SSHServer, error) {
301 	cfg := config.FromContext(ctx)
302+	logger := log.FromContext(ctx).WithPrefix("ssh")
303 
304 	var err error
305 	s := &SSHServer{
306 		cfg:    cfg,
307 		ctx:    ctx,
308-		logger: log.FromContext(ctx).WithPrefix("ssh"),
309+		be:     backend.FromContext(ctx),
310+		logger: logger,
311 	}
312 
313-	logger := s.logger.StandardLog(log.StandardLogOptions{ForceLevel: log.DebugLevel})
314 	mw := []wish.Middleware{
315 		rm.MiddlewareWithLogger(
316 			logger,
317 			// BubbleTea middleware.
318 			bm.MiddlewareWithProgramHandler(SessionHandler(cfg), termenv.ANSI256),
319 			// CLI middleware.
320-			cm.Middleware(cfg),
321+			cm.Middleware(cfg, logger),
322 			// Git middleware.
323 			s.Middleware(cfg),
324 			// Logging middleware.
325-			lm.MiddlewareWithLogger(logger),
326+			lm.MiddlewareWithLogger(logger.
327+				StandardLog(log.StandardLogOptions{ForceLevel: log.DebugLevel})),
328 		),
329 	}
330 
331@@ -191,6 +194,8 @@ func (ss *SSHServer) Middleware(cfg *config.Config) wish.Middleware {
332 		return func(s ssh.Session) {
333 			func() {
334 				cmd := s.Command()
335+				ctx := s.Context()
336+				be := ss.be.WithContext(ctx)
337 				if len(cmd) >= 2 && strings.HasPrefix(cmd[0], "git") {
338 					gc := cmd[0]
339 					// repo should be in the form of "repo.git"
340@@ -222,8 +227,8 @@ func (ss *SSHServer) Middleware(cfg *config.Config) wish.Middleware {
341 							sshFatal(s, git.ErrNotAuthed)
342 							return
343 						}
344-						if _, err := cfg.Backend.Repository(name); err != nil {
345-							if _, err := cfg.Backend.CreateRepository(name, backend.RepositoryOptions{Private: false}); err != nil {
346+						if _, err := be.Repository(name); err != nil {
347+							if _, err := be.CreateRepository(name, backend.RepositoryOptions{Private: false}); err != nil {
348 								log.Errorf("failed to create repo: %s", err)
349 								sshFatal(s, err)
350 								return
351@@ -248,7 +253,7 @@ func (ss *SSHServer) Middleware(cfg *config.Config) wish.Middleware {
352 							counter = uploadArchiveCounter
353 						}
354 
355-						err := gitPack(s.Context(), s, s, s.Stderr(), repoDir, envs...)
356+						err := gitPack(ctx, s, s, s.Stderr(), repoDir, envs...)
357 						if errors.Is(err, git.ErrInvalidRepo) {
358 							sshFatal(s, git.ErrInvalidRepo)
359 						} else if err != nil {
360diff --git a/server/web/http.go b/server/web/http.go
361index c3e33e0f9960eacccccf6500cc5751e64cd9f0a5..b932f94162bf6c2c6119e451b682d2a132c7e233 100644
362--- a/server/web/http.go
363+++ b/server/web/http.go
364@@ -81,6 +81,7 @@ func (s *HTTPServer) loggingMiddleware(next http.Handler) http.Handler {
365 type HTTPServer struct {
366 	ctx        context.Context
367 	cfg        *config.Config
368+	be         backend.Backend
369 	server     *http.Server
370 	dirHandler http.Handler
371 	logger     *log.Logger
372@@ -92,6 +93,7 @@ func NewHTTPServer(ctx context.Context) (*HTTPServer, error) {
373 	s := &HTTPServer{
374 		ctx:        ctx,
375 		cfg:        cfg,
376+		be:         backend.FromContext(ctx),
377 		logger:     log.FromContext(ctx).WithPrefix("http"),
378 		dirHandler: http.FileServer(http.Dir(filepath.Join(cfg.DataPath, "repos"))),
379 		server: &http.Server{
380@@ -254,6 +256,7 @@ Redirecting to docs at <a href="https://godoc.org/{{ .ImportRoot }}/{{ .Repo }}"
381 func (s *HTTPServer) handleIndex(w http.ResponseWriter, r *http.Request) {
382 	repo := pattern.Path(r.Context())
383 	repo = utils.SanitizeRepo(repo)
384+	be := s.be.WithContext(r.Context())
385 
386 	// Handle go get requests.
387 	//
388@@ -271,7 +274,7 @@ func (s *HTTPServer) handleIndex(w http.ResponseWriter, r *http.Request) {
389 
390 		// find the repo
391 		for {
392-			if _, err := s.cfg.Backend.Repository(repo); err == nil {
393+			if _, err := be.Repository(repo); err == nil {
394 				break
395 			}
396 
397@@ -305,7 +308,8 @@ func (s *HTTPServer) handleIndex(w http.ResponseWriter, r *http.Request) {
398 func (s *HTTPServer) handleGit(w http.ResponseWriter, r *http.Request) {
399 	repo := pat.Param(r, "repo")
400 	repo = utils.SanitizeRepo(repo) + ".git"
401-	if _, err := s.cfg.Backend.Repository(repo); err != nil {
402+	be := s.be.WithContext(r.Context())
403+	if _, err := be.Repository(repo); err != nil {
404 		s.logger.Debug("repository not found", "repo", repo, "err", err)
405 		http.NotFound(w, r)
406 		return