This repository has no description
15 kB
489 lines
1package middleware
2
3import (
4 "context"
5 "database/sql"
6 "errors"
7 "fmt"
8 "log/slog"
9 "net/http"
10 "net/url"
11 "path"
12 "slices"
13 "strconv"
14 "strings"
15
16 "github.com/bluesky-social/indigo/atproto/identity"
17 "github.com/bluesky-social/indigo/atproto/syntax"
18 "github.com/go-chi/chi/v5"
19 "tangled.org/core/appview/cache"
20 "tangled.org/core/appview/db"
21 "tangled.org/core/appview/knotacl"
22 "tangled.org/core/appview/models"
23 "tangled.org/core/appview/oauth"
24 "tangled.org/core/appview/pages"
25 "tangled.org/core/appview/pagination"
26 "tangled.org/core/appview/reporesolver"
27 "tangled.org/core/appview/state/userutil"
28 "tangled.org/core/idresolver"
29 "tangled.org/core/orm"
30 "tangled.org/core/rbac"
31)
32
33type Middleware struct {
34 oauth *oauth.OAuth
35 db *db.DB
36 enforcer *rbac.Enforcer
37 acl *knotacl.Service
38 repoResolver *reporesolver.RepoResolver
39 idResolver *idresolver.Resolver
40 pages *pages.Pages
41 rdb *cache.Cache
42 logger *slog.Logger
43}
44
45func New(oauth *oauth.OAuth, db *db.DB, enforcer *rbac.Enforcer, acl *knotacl.Service, repoResolver *reporesolver.RepoResolver, idResolver *idresolver.Resolver, pages *pages.Pages, rdb *cache.Cache, logger *slog.Logger) Middleware {
46 return Middleware{
47 oauth: oauth,
48 db: db,
49 enforcer: enforcer,
50 acl: acl,
51 repoResolver: repoResolver,
52 idResolver: idResolver,
53 pages: pages,
54 rdb: rdb,
55 logger: logger,
56 }
57}
58
59type middlewareFunc func(http.Handler) http.Handler
60
61func AuthMiddleware(o *oauth.OAuth) middlewareFunc {
62 return func(next http.Handler) http.Handler {
63 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
64 returnURL := "/"
65 if u, err := url.Parse(r.Header.Get("Referer")); err == nil {
66 returnURL = u.RequestURI()
67 }
68
69 loginURL := fmt.Sprintf("/login?return_url=%s", url.QueryEscape(returnURL))
70
71 redirectFunc := func(w http.ResponseWriter, r *http.Request) {
72 http.Redirect(w, r, loginURL, http.StatusTemporaryRedirect)
73 }
74 if r.Header.Get("HX-Request") == "true" {
75 redirectFunc = func(w http.ResponseWriter, _ *http.Request) {
76 w.Header().Set("HX-Redirect", loginURL)
77 w.WriteHeader(http.StatusOK)
78 }
79 }
80
81 sess, err := o.ResumeSession(r)
82 if err != nil {
83 slog.Default().Warn("failed to resume session, redirecting", "err", err, "url", r.URL.String())
84 redirectFunc(w, r)
85 return
86 }
87
88 if sess == nil {
89 slog.Default().Warn("session is nil, redirecting")
90 redirectFunc(w, r)
91 return
92 }
93
94 next.ServeHTTP(w, r)
95 })
96 }
97}
98
99func (m *Middleware) InjectBaseParams(next http.Handler) http.Handler {
100 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
101 user := m.oauth.GetMultiAccountUser(r)
102 bp := pages.BaseParams{
103 LoggedInUser: user,
104 }
105 if user != nil {
106 if focusing, _ := db.GetFocusStatus(m.db, user.Did); focusing {
107 if item, _ := db.GetNextFocusItem(m.db, user.Did); item != nil {
108 count, _ := db.CountFocusNotifs(m.db, user.Did)
109 bp.FocusParams = pages.FocusParams{
110 Focusing: true,
111 FocusLink: item.URL(m.idResolver),
112 FocusNotificationID: item.ID,
113 CurrentPath: r.URL.Path,
114 FocusCount: int(count),
115 }
116 } else {
117 // queue exhausted — auto-exit focus mode
118 _ = db.EndFocus(m.db, user.Did)
119 }
120 }
121 }
122 ctx := pages.BaseParamsIntoContext(r.Context(), bp)
123 next.ServeHTTP(w, r.WithContext(ctx))
124 })
125}
126
127func Paginate(next http.Handler) http.Handler {
128 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
129 page := pagination.FirstPage()
130
131 offsetVal := r.URL.Query().Get("offset")
132 if offsetVal != "" {
133 offset, err := strconv.Atoi(offsetVal)
134 if err != nil {
135 slog.Default().Warn("invalid offset", "value", offsetVal)
136 } else {
137 page.Offset = offset
138 }
139 }
140
141 limitVal := r.URL.Query().Get("limit")
142 if limitVal != "" {
143 limit, err := strconv.Atoi(limitVal)
144 if err != nil {
145 slog.Default().Warn("invalid limit", "value", limitVal)
146 } else {
147 page.Limit = limit
148 }
149 }
150
151 ctx := pagination.IntoContext(r.Context(), page)
152 next.ServeHTTP(w, r.WithContext(ctx))
153 })
154}
155
156func (mw Middleware) knotRoleMiddleware(group string) middlewareFunc {
157 return func(next http.Handler) http.Handler {
158 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
159 l := mw.logger.With("middleware", "knotRoleMiddleware")
160 // requires auth also
161 actor := mw.oauth.GetMultiAccountUser(r)
162 if actor == nil {
163 // we need a logged in user
164 l.Warn("not logged in, redirecting")
165 http.Error(w, "Forbidden", http.StatusUnauthorized)
166 return
167 }
168 domain := chi.URLParam(r, "domain")
169 if domain == "" {
170 http.Error(w, "malformed url", http.StatusBadRequest)
171 return
172 }
173
174 ok, err := mw.enforcer.E.HasGroupingPolicy(actor.Did, group, domain)
175 if err != nil || !ok {
176 l.Warn("permission denied", "did", actor.Did, "group", group, "domain", domain)
177 http.Error(w, "Forbidden", http.StatusUnauthorized)
178 return
179 }
180
181 next.ServeHTTP(w, r)
182 })
183 }
184}
185
186func (mw Middleware) KnotOwner() middlewareFunc {
187 return mw.knotRoleMiddleware("server:owner")
188}
189
190func (mw Middleware) RepoPermissionMiddleware(requiredPerm string) middlewareFunc {
191 return func(next http.Handler) http.Handler {
192 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
193 l := mw.logger.With("middleware", "RepoPermissionMiddleware")
194 // requires auth also
195 actor := mw.oauth.GetMultiAccountUser(r)
196 if actor == nil {
197 // we need a logged in user
198 l.Warn("not logged in, redirecting")
199 http.Error(w, "Forbidden", http.StatusUnauthorized)
200 return
201 }
202 f, err := mw.repoResolver.Resolve(r)
203 if err != nil {
204 http.Error(w, "malformed url", http.StatusBadRequest)
205 return
206 }
207
208 if !mw.acl.HasRepoPermission(r.Context(), f, actor.Did, requiredPerm) {
209 l.Warn("permission denied", "did", actor.Did, "perm", requiredPerm, "repo", f.RepoIdentifier())
210 http.Error(w, "Forbidden", http.StatusUnauthorized)
211 return
212 }
213
214 next.ServeHTTP(w, r)
215 })
216 }
217}
218
219func (mw Middleware) ResolveIdent() middlewareFunc {
220 excluded := []string{"favicon.ico"}
221
222 return func(next http.Handler) http.Handler {
223 return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
224 origSeg := chi.URLParam(req, "user")
225 didOrHandle := strings.TrimPrefix(origSeg, "@")
226 didOrHandle = strings.TrimSuffix(didOrHandle, ".keys")
227
228 if slices.Contains(excluded, didOrHandle) {
229 next.ServeHTTP(w, req)
230 return
231 }
232
233 id, err := mw.idResolver.ResolveAtIdentifier(req.Context(), didOrHandle)
234 if err != nil {
235 if h, parseErr := syntax.ParseHandle(didOrHandle); parseErr == nil {
236 if did := cache.LookupDidByPreferredHandle(req.Context(), mw.rdb, mw.db, h); did != "" {
237 id, err = mw.idResolver.ResolveAtIdentifier(req.Context(), did)
238 }
239 }
240 }
241 if err != nil {
242 mw.logger.Error("failed to resolve did/handle", "didOrHandle", didOrHandle, "err", err)
243 mw.pages.Error404(w)
244 return
245 }
246
247 if req.Method == http.MethodGet && !userutil.IsDid(didOrHandle) {
248 if pref := cache.LookupPreferredHandle(req.Context(), mw.rdb, mw.db, id.DID.String()); pref != "" && didOrHandle != pref {
249 rest := strings.TrimPrefix(req.URL.Path, "/"+origSeg)
250 target := "/" + pref + rest
251 if req.URL.RawQuery != "" {
252 target += "?" + req.URL.RawQuery
253 }
254 http.Redirect(w, req, target, http.StatusFound)
255 return
256 }
257 }
258
259 ctx := context.WithValue(req.Context(), "resolvedId", *id)
260
261 next.ServeHTTP(w, req.WithContext(ctx))
262 })
263 }
264}
265
266func (mw Middleware) ResolveRepo() middlewareFunc {
267 return func(next http.Handler) http.Handler {
268 return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
269 l := mw.logger.With("middleware", "ResolveRepo")
270 repoName := strings.TrimSuffix(chi.URLParam(req, "repo"), ".git")
271 rkey := strings.ToLower(repoName)
272
273 id, ok := req.Context().Value("resolvedId").(identity.Identity)
274 if !ok {
275 l.Error("malformed middleware")
276 w.WriteHeader(http.StatusInternalServerError)
277 return
278 }
279
280 repo, isRename := resolveRepoForOwner(mw.db, id.DID.String(), repoName, rkey, l)
281 if repo == nil {
282 w.WriteHeader(http.StatusNotFound)
283 mw.pages.ErrorKnot404(w)
284 return
285 }
286 if isRename {
287 handle := id.Handle.String()
288 if id.Handle.IsInvalidHandle() || handle == "" {
289 handle = id.DID.String()
290 }
291 canonical := reporesolver.CanonicalRepoPath(handle, repo)
292 if path.Join(chi.URLParam(req, "user"), repoName) != canonical {
293 target := reporesolver.CanonicalRedirectTarget(req, canonical)
294 http.Redirect(w, req, target, http.StatusMovedPermanently)
295 return
296 }
297 }
298
299 ctx := context.WithValue(req.Context(), "repo", repo)
300 next.ServeHTTP(w, req.WithContext(ctx))
301 })
302 }
303}
304
305func resolveRepoForOwner(d db.Execer, ownerDid, repoName, rkey string, l *slog.Logger) (*models.Repo, bool) {
306 repo, err := db.GetRepo(d, orm.FilterEq("did", ownerDid), orm.FilterEq("rkey", rkey))
307 if err == nil {
308 return repo, false
309 }
310 if !errors.Is(err, sql.ErrNoRows) {
311 l.Error("failed to resolve repo by rkey", "err", err)
312 return nil, false
313 }
314
315 hint, hintErr := db.LookupRepoRename(d, ownerDid, rkey)
316 if hintErr != nil && !errors.Is(hintErr, sql.ErrNoRows) {
317 l.Error("failed to lookup repo rename hint", "err", hintErr)
318 }
319 if hint != nil {
320 return hint, true
321 }
322
323 nameRepos, nameErr := db.GetRepos(d, orm.FilterEq("did", ownerDid), orm.FilterEq("name", repoName))
324 if nameErr != nil {
325 l.Error("failed to resolve repo by name", "err", nameErr)
326 return nil, false
327 }
328 if len(nameRepos) == 1 {
329 return &nameRepos[0], false
330 }
331 return nil, false
332}
333
334func (mw Middleware) CanonicalizeRepoURL() middlewareFunc {
335 return func(next http.Handler) http.Handler {
336 return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
337 if req.Method != http.MethodGet && req.Method != http.MethodHead {
338 next.ServeHTTP(w, req)
339 return
340 }
341 id, idOk := req.Context().Value("resolvedId").(identity.Identity)
342 repo, repoOk := req.Context().Value("repo").(*models.Repo)
343 if !idOk || !repoOk || id.Handle.IsInvalidHandle() {
344 next.ServeHTTP(w, req)
345 return
346 }
347 handle := id.Handle.String()
348 if handle == "" {
349 next.ServeHTTP(w, req)
350 return
351 }
352 canonical := reporesolver.CanonicalRepoPath(handle, repo)
353 urlUser := chi.URLParam(req, "user")
354 urlRepo := strings.TrimSuffix(chi.URLParam(req, "repo"), ".git")
355 if urlUser+"/"+urlRepo == canonical {
356 next.ServeHTTP(w, req)
357 return
358 }
359
360 http.Redirect(w, req, reporesolver.CanonicalRedirectTarget(req, canonical), http.StatusFound)
361 })
362 }
363}
364
365// middleware that is tacked on top of /{user}/{repo}/pulls/{pull}
366func (mw Middleware) ResolvePull() middlewareFunc {
367 return func(next http.Handler) http.Handler {
368 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
369 l := mw.logger.With("middleware", "ResolvePull")
370 f, err := mw.repoResolver.Resolve(r)
371 if err != nil {
372 l.Error("failed to fully resolve repo", "err", err)
373 w.WriteHeader(http.StatusNotFound)
374 mw.pages.ErrorKnot404(w)
375 return
376 }
377
378 prId := chi.URLParam(r, "pull")
379 prIdInt, err := strconv.Atoi(prId)
380 if err != nil {
381 l.Error("failed to parse pr id", "err", err)
382 mw.pages.Error404(w)
383 return
384 }
385
386 pr, err := db.GetPull(mw.db, orm.FilterEq("repo_did", f.RepoDid), orm.FilterEq("pull_id", prIdInt))
387 if err != nil {
388 l.Error("failed to get pull and comments", "err", err)
389 mw.pages.Error404(w)
390 return
391 }
392
393 ctx := context.WithValue(r.Context(), "pull", pr)
394
395 stack, err := db.GetStack(mw.db, pr.AtUri())
396 if err != nil {
397 l.Error("failed to get stack", "err", err)
398 mw.pages.Error404(w)
399 return
400 }
401
402 ctx = context.WithValue(ctx, "stack", stack)
403
404 next.ServeHTTP(w, r.WithContext(ctx))
405 })
406 }
407}
408
409// middleware that is tacked on top of /{user}/{repo}/issues/{issue}
410func (mw Middleware) ResolveIssue(next http.Handler) http.Handler {
411 l := mw.logger.With("middleware", "ResolveIssue")
412 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
413 f, err := mw.repoResolver.Resolve(r)
414 if err != nil {
415 l.Error("failed to fully resolve repo", "err", err)
416 w.WriteHeader(http.StatusNotFound)
417 mw.pages.ErrorKnot404(w)
418 return
419 }
420
421 issueIdStr := chi.URLParam(r, "issue")
422 issueId, err := strconv.Atoi(issueIdStr)
423 if err != nil {
424 l.Error("failed to fully resolve issue ID", "err", err)
425 mw.pages.Error404(w)
426 return
427 }
428
429 issue, err := db.GetIssue(mw.db, f.RepoDid, issueId)
430 if err != nil {
431 l.Error("failed to get issues", "err", err)
432 mw.pages.Error404(w)
433 return
434 }
435
436 ctx := context.WithValue(r.Context(), "issue", issue)
437 next.ServeHTTP(w, r.WithContext(ctx))
438 })
439}
440
441// this should serve the go-import meta tag even if the path is technically
442// a 404 like tangled.sh/oppi.li/go-git/v5
443//
444// we're keeping the tangled.sh go-import tag too to maintain backward
445// compatibility for modules that still point there. they will be redirected
446// to fetch source from tangled.org
447func (mw Middleware) GoImport() middlewareFunc {
448 return func(next http.Handler) http.Handler {
449 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
450 l := mw.logger.With("middleware", "GoImport")
451 f, err := mw.repoResolver.Resolve(r)
452 if err != nil {
453 l.Error("failed to fully resolve repo", "err", err)
454 w.WriteHeader(http.StatusNotFound)
455 mw.pages.ErrorKnot404(w)
456 return
457 }
458
459 fullName := reporesolver.GetBaseRepoPath(r, f)
460
461 if r.Header.Get("User-Agent") == "Go-http-client/1.1" {
462 if r.URL.Query().Get("go-get") == "1" {
463 modulePath := userutil.FlattenDid(fullName)
464 if strings.Contains(modulePath, ":") {
465 modulePath = userutil.FlattenDid(f.Did) + "/" + f.Rkey
466 }
467 tags := []string{
468 fmt.Sprintf(`<meta name="go-import" content="tangled.sh/%s git https://tangled.sh/%s"/>`, modulePath, fullName),
469 fmt.Sprintf(`<meta name="go-import" content="tangled.org/%s git https://tangled.org/%s"/>`, modulePath, fullName),
470 }
471 if f.RepoDid != "" {
472 stable := userutil.FlattenDid(f.RepoDid)
473 if stable != modulePath {
474 tags = append(tags,
475 fmt.Sprintf(`<meta name="go-import" content="tangled.sh/%s git https://tangled.sh/%s"/>`, stable, f.RepoDid),
476 fmt.Sprintf(`<meta name="go-import" content="tangled.org/%s git https://tangled.org/%s"/>`, stable, f.RepoDid),
477 )
478 }
479 }
480 w.Header().Set("Content-Type", "text/html")
481 w.Write([]byte(strings.Join(tags, "\n")))
482 return
483 }
484 }
485
486 next.ServeHTTP(w, r)
487 })
488 }
489}