forked from
tangled.org/core
Monorepo for Tangled
16 kB
541 lines
1package oauth
2
3import (
4 "bytes"
5 "context"
6 "encoding/json"
7 "errors"
8 "fmt"
9 "io"
10 "log/slog"
11 "net/http"
12 "strings"
13 "time"
14
15 comatproto "github.com/bluesky-social/indigo/api/atproto"
16 "github.com/bluesky-social/indigo/atproto/auth/oauth"
17 "github.com/bluesky-social/indigo/atproto/syntax"
18 xrpc "github.com/bluesky-social/indigo/xrpc"
19 "github.com/go-chi/chi/v5"
20 "github.com/posthog/posthog-go"
21 "tangled.org/core/api/tangled"
22 "tangled.org/core/appview/db"
23 "tangled.org/core/appview/knotcompat"
24 "tangled.org/core/appview/models"
25 "tangled.org/core/consts"
26 "tangled.org/core/idresolver"
27 "tangled.org/core/orm"
28 "tangled.org/core/tid"
29)
30
31const knotAdminTimeout = 30 * time.Second
32
33func (o *OAuth) Router() http.Handler {
34 r := chi.NewRouter()
35
36 r.Get("/oauth/client-metadata.json", o.clientMetadata)
37 r.Get("/oauth/jwks.json", o.jwks)
38 r.Get("/oauth/callback", o.callback)
39 return r
40}
41
42func (o *OAuth) clientMetadata(w http.ResponseWriter, r *http.Request) {
43 doc := o.ClientApp.Config.ClientMetadata()
44 doc.JWKSURI = &o.JwksUri
45 doc.ClientName = &o.ClientName
46 doc.ClientURI = &o.ClientUri
47 doc.Scope = doc.Scope + " identity:handle"
48
49 w.Header().Set("Content-Type", "application/json")
50 if err := json.NewEncoder(w).Encode(doc); err != nil {
51 http.Error(w, err.Error(), http.StatusInternalServerError)
52 return
53 }
54}
55
56func (o *OAuth) jwks(w http.ResponseWriter, r *http.Request) {
57 w.Header().Set("Content-Type", "application/json")
58 body := o.ClientApp.Config.PublicJWKS()
59 if err := json.NewEncoder(w).Encode(body); err != nil {
60 http.Error(w, err.Error(), http.StatusInternalServerError)
61 return
62 }
63}
64
65func (o *OAuth) callback(w http.ResponseWriter, r *http.Request) {
66 ctx := r.Context()
67 l := o.Logger.With("query", r.URL.Query())
68
69 redirectURL := o.GetAuthReturn(r)
70 _ = o.ClearAuthReturn(w, r)
71
72 sessData, err := o.ClientApp.ProcessCallback(ctx, r.URL.Query())
73 if err != nil {
74 var callbackErr *oauth.AuthRequestCallbackError
75 if errors.As(err, &callbackErr) {
76 l.Debug("callback error", "err", callbackErr)
77 http.Redirect(w, r, fmt.Sprintf("/login?error=%s", callbackErr.ErrorCode), http.StatusFound)
78 return
79 }
80 l.Error("failed to process callback", "err", err)
81 http.Redirect(w, r, "/login?error=oauth", http.StatusFound)
82 return
83 }
84
85 if err := o.SaveSession(w, r, sessData); err != nil {
86 l.Error("failed to save session", "data", sessData, "err", err)
87 errorCode := "session"
88 if errors.Is(err, ErrMaxAccountsReached) {
89 errorCode = "max_accounts"
90 }
91 http.Redirect(w, r, fmt.Sprintf("/login?error=%s", errorCode), http.StatusFound)
92 return
93 }
94
95 o.Logger.Debug("session saved successfully")
96
97 did := sessData.AccountDID.String()
98
99 // default to true, so users don't have to onboard again
100 isTangledUser, err := db.IsTangledUser(o.Db, did)
101 if err != nil {
102 isTangledUser = true
103 }
104
105 isNewUser := !isTangledUser
106 if isNewUser {
107 if ob, _ := db.GetOnboarding(o.Db, did); ob == nil {
108 if err := db.UpsertOnboarding(o.Db, &models.Onboarding{
109 Did: did,
110 Step: models.OnboardingStepProfile,
111 Status: models.OnboardingInProgress,
112 }); err != nil {
113 o.Logger.Error("failed to seed onboarding record", "did", did, "err", err)
114 }
115 }
116 }
117
118 go o.addToDefaultKnot(sessData.AccountDID)
119 go o.addToDefaultSpindle(sessData.AccountDID.String())
120 go o.autoClaimTnglShDomain(sessData.AccountDID.String())
121
122 if !o.Config.Core.Dev {
123 err = o.Posthog.Enqueue(posthog.Capture{
124 DistinctId: sessData.AccountDID.String(),
125 Event: "signin",
126 })
127 if err != nil {
128 o.Logger.Error("failed to enqueue posthog event", "err", err)
129 }
130 }
131
132 if redirectURL == "" {
133 redirectURL = "/"
134 }
135
136 if o.isAccountDeactivated(sessData) {
137 redirectURL = "/settings/profile"
138 } else if isNewUser {
139 redirectURL = "/welcome"
140 }
141
142 http.Redirect(w, r, redirectURL, http.StatusFound)
143}
144
145func (o *OAuth) isAccountDeactivated(sessData *oauth.ClientSessionData) bool {
146 pdsClient := &xrpc.Client{
147 Host: sessData.HostURL,
148 Client: &http.Client{Timeout: 5 * time.Second},
149 }
150
151 _, err := comatproto.RepoDescribeRepo(
152 context.Background(),
153 pdsClient,
154 sessData.AccountDID.String(),
155 )
156 if err == nil {
157 return false
158 }
159
160 var xrpcErr *xrpc.Error
161 var xrpcBody *xrpc.XRPCError
162 return errors.As(err, &xrpcErr) &&
163 errors.As(xrpcErr.Wrapped, &xrpcBody) &&
164 xrpcBody.ErrStr == "RepoDeactivated"
165}
166
167func (o *OAuth) addToDefaultSpindle(did string) {
168 l := o.Logger.With("subject", did)
169
170 // use the tangled.sh app password to get an accessJwt
171 // and create an sh.tangled.spindle.member record with that
172 spindleMembers, err := db.GetSpindleMembers(
173 o.Db,
174 orm.FilterEq("instance", "spindle.tangled.sh"),
175 orm.FilterEq("subject", did),
176 )
177 if err != nil {
178 l.Error("failed to get spindle members", "err", err)
179 return
180 }
181
182 if len(spindleMembers) != 0 {
183 l.Warn("already a member of the default spindle")
184 return
185 }
186
187 l.Debug("adding to default spindle")
188 session, err := o.getAppPasswordSession()
189 if err != nil {
190 l.Error("failed to create session", "err", err)
191 return
192 }
193
194 record := tangled.SpindleMember{
195 LexiconTypeID: tangled.SpindleMemberNSID,
196 Subject: did,
197 Instance: consts.DefaultSpindle,
198 CreatedAt: time.Now().Format(time.RFC3339),
199 }
200
201 if err := session.putRecord(record, tangled.SpindleMemberNSID); err != nil {
202 l.Error("failed to add to default spindle", "err", err)
203 return
204 }
205
206 l.Debug("successfully added to default spindle", "did", did)
207}
208
209type onboardAction int
210
211const (
212 onboardViaAdminAPI onboardAction = iota
213 onboardViaRecord
214 onboardBlockedMissingSecret
215 onboardBlockedSecretSet
216)
217
218type defaultKnotState struct {
219 native bool
220 adminSecretSet bool
221}
222
223func onboardActionFor(s defaultKnotState) onboardAction {
224 switch {
225 case s.native && s.adminSecretSet:
226 return onboardViaAdminAPI
227 case s.native:
228 return onboardBlockedMissingSecret
229 case s.adminSecretSet:
230 return onboardBlockedSecretSet
231 default:
232 return onboardViaRecord
233 }
234}
235
236func (o *OAuth) addToDefaultKnot(did syntax.DID) {
237 l := o.Logger.With("subject", did)
238
239 ctx := context.Background()
240
241 if o.Acl.IsKnotMember(ctx, o.Config.Knot.Default, did.String()) {
242 l.Warn("already a member of the default knot")
243 return
244 }
245
246 native := knotcompat.KnotHasCapability(ctx, o.Config.Knot.Default, o.Config.Core.Dev, consts.CapKnotACL)
247
248 switch onboardActionFor(defaultKnotState{native: native, adminSecretSet: o.Config.Knot.AdminSecret != ""}) {
249 case onboardViaAdminAPI:
250 if err := o.addMemberViaKnotAdmin(ctx, o.Config.Knot.Default, did); err != nil {
251 l.Error("failed to add to default knot via admin api", "err", err)
252 return
253 }
254 o.Acl.InvalidateMembers(o.Config.Knot.Default)
255 l.Debug("successfully added to default knot via admin api")
256
257 case onboardBlockedMissingSecret:
258 l.Error("cannot add to default knot: knot admin secret not configured")
259
260 case onboardBlockedSecretSet:
261 l.Warn("default knot probe failed, skipping legacy fallback because an admin secret is configured")
262
263 case onboardViaRecord:
264 l.Debug("adding to default knot")
265 session, err := o.getAppPasswordSession()
266 if err != nil {
267 l.Error("failed to create session", "err", err)
268 return
269 }
270
271 record := tangled.KnotMember{
272 LexiconTypeID: tangled.KnotMemberNSID,
273 Subject: did.String(),
274 Domain: o.Config.Knot.Default,
275 CreatedAt: time.Now().Format(time.RFC3339),
276 }
277
278 if err := session.putRecord(record, tangled.KnotMemberNSID); err != nil {
279 l.Error("failed to add to default knot", "err", err)
280 return
281 }
282
283 if err := o.Enforcer.AddKnotMember(o.Config.Knot.Default, did.String()); err != nil {
284 l.Error("failed to set up enforcer rules", "err", err)
285 return
286 }
287
288 l.Debug("successfully added to default knot")
289 }
290}
291
292func (o *OAuth) addMemberViaKnotAdmin(ctx context.Context, knotHost string, subject syntax.DID) error {
293 ctx, cancel := context.WithTimeout(ctx, knotAdminTimeout)
294 defer cancel()
295
296 scheme := "https://"
297 if o.Config.Core.Dev {
298 scheme = "http://"
299 }
300 endpoint := fmt.Sprintf("%s%s/admin/addMember", scheme, knotHost)
301
302 body, err := json.Marshal(tangled.KnotAddMember_Input{Subject: subject.String()})
303 if err != nil {
304 return err
305 }
306
307 req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body))
308 if err != nil {
309 return err
310 }
311 req.Header.Set("Content-Type", "application/json")
312 req.SetBasicAuth("admin", o.Config.Knot.AdminSecret)
313
314 resp, err := http.DefaultClient.Do(req)
315 if err != nil {
316 return err
317 }
318 defer resp.Body.Close()
319
320 if resp.StatusCode != http.StatusOK {
321 msg, _ := io.ReadAll(resp.Body)
322 return fmt.Errorf("knot admin addMember returned status %d: %s", resp.StatusCode, bytes.TrimSpace(msg))
323 }
324
325 return nil
326}
327
328// create a AppPasswordSession using apppasswords
329type AppPasswordSession struct {
330 AccessJwt string `json:"accessJwt"`
331 RefreshJwt string `json:"refreshJwt"`
332 PdsEndpoint string
333 Did string
334 Logger *slog.Logger
335 ExpiresAt time.Time
336}
337
338func CreateAppPasswordSession(res *idresolver.Resolver, appPassword, did string, logger *slog.Logger) (*AppPasswordSession, error) {
339 if appPassword == "" {
340 return nil, fmt.Errorf("no app password configured")
341 }
342
343 resolved, err := res.ResolveIdent(context.Background(), did)
344 if err != nil {
345 return nil, fmt.Errorf("failed to resolve tangled.sh DID %s: %v", did, err)
346 }
347
348 pdsEndpoint := resolved.PDSEndpoint()
349 if pdsEndpoint == "" {
350 return nil, fmt.Errorf("no PDS endpoint found for tangled.sh DID %s", did)
351 }
352
353 sessionPayload := map[string]string{
354 "identifier": did,
355 "password": appPassword,
356 }
357 sessionBytes, err := json.Marshal(sessionPayload)
358 if err != nil {
359 return nil, fmt.Errorf("failed to marshal session payload: %v", err)
360 }
361
362 sessionURL := pdsEndpoint + "/xrpc/com.atproto.server.createSession"
363 sessionReq, err := http.NewRequestWithContext(context.Background(), "POST", sessionURL, bytes.NewBuffer(sessionBytes))
364 if err != nil {
365 return nil, fmt.Errorf("failed to create session request: %v", err)
366 }
367 sessionReq.Header.Set("Content-Type", "application/json")
368
369 logger.Debug("creating app password session", "url", sessionURL, "headers", sessionReq.Header)
370
371 client := &http.Client{Timeout: 30 * time.Second}
372 sessionResp, err := client.Do(sessionReq)
373 if err != nil {
374 return nil, fmt.Errorf("failed to create session: %v", err)
375 }
376 defer sessionResp.Body.Close()
377
378 if sessionResp.StatusCode != http.StatusOK {
379 return nil, fmt.Errorf("failed to create session: HTTP %d", sessionResp.StatusCode)
380 }
381
382 var session AppPasswordSession
383 if err := json.NewDecoder(sessionResp.Body).Decode(&session); err != nil {
384 return nil, fmt.Errorf("failed to decode session response: %v", err)
385 }
386
387 session.PdsEndpoint = pdsEndpoint
388 session.Did = did
389 session.Logger = logger
390 session.ExpiresAt = time.Now().Add(115 * time.Minute)
391
392 return &session, nil
393}
394
395func (s *AppPasswordSession) RefreshSession() error {
396 refreshURL := s.PdsEndpoint + "/xrpc/com.atproto.server.refreshSession"
397 req, err := http.NewRequestWithContext(context.Background(), "POST", refreshURL, nil)
398 if err != nil {
399 return fmt.Errorf("failed to create refresh request: %w", err)
400 }
401
402 req.Header.Set("Authorization", "Bearer "+s.RefreshJwt)
403
404 s.Logger.Debug("refreshing app password session", "url", refreshURL)
405
406 client := &http.Client{Timeout: 30 * time.Second}
407 resp, err := client.Do(req)
408 if err != nil {
409 return fmt.Errorf("failed to refresh session: %w", err)
410 }
411 defer resp.Body.Close()
412
413 if resp.StatusCode != http.StatusOK {
414 var errorResponse map[string]any
415 if err := json.NewDecoder(resp.Body).Decode(&errorResponse); err != nil {
416 return fmt.Errorf("failed to refresh session: HTTP %d (failed to decode error response: %w)", resp.StatusCode, err)
417 }
418 errorBytes, _ := json.Marshal(errorResponse)
419 return fmt.Errorf("failed to refresh session: HTTP %d, response: %s", resp.StatusCode, string(errorBytes))
420 }
421
422 var refreshResponse struct {
423 AccessJwt string `json:"accessJwt"`
424 RefreshJwt string `json:"refreshJwt"`
425 }
426 if err := json.NewDecoder(resp.Body).Decode(&refreshResponse); err != nil {
427 return fmt.Errorf("failed to decode refresh response: %w", err)
428 }
429
430 s.AccessJwt = refreshResponse.AccessJwt
431 s.RefreshJwt = refreshResponse.RefreshJwt
432 // Set new expiry time with 5 minute buffer
433 s.ExpiresAt = time.Now().Add(115 * time.Minute)
434
435 s.Logger.Debug("successfully refreshed app password session")
436 return nil
437}
438
439func (s *AppPasswordSession) IsValid() bool {
440 return time.Now().Before(s.ExpiresAt)
441}
442
443func (s *AppPasswordSession) putRecord(record any, collection string) error {
444 if !s.IsValid() {
445 s.Logger.Debug("access token expired, refreshing session")
446 if err := s.RefreshSession(); err != nil {
447 return fmt.Errorf("failed to refresh session: %w", err)
448 }
449 s.Logger.Debug("session refreshed")
450 }
451
452 recordBytes, err := json.Marshal(record)
453 if err != nil {
454 return fmt.Errorf("failed to marshal knot member record: %w", err)
455 }
456
457 payload := map[string]any{
458 "repo": s.Did,
459 "collection": collection,
460 "rkey": tid.TID(),
461 "record": json.RawMessage(recordBytes),
462 }
463
464 payloadBytes, err := json.Marshal(payload)
465 if err != nil {
466 return fmt.Errorf("failed to marshal request payload: %w", err)
467 }
468
469 url := s.PdsEndpoint + "/xrpc/com.atproto.repo.putRecord"
470 req, err := http.NewRequestWithContext(context.Background(), "POST", url, bytes.NewBuffer(payloadBytes))
471 if err != nil {
472 return fmt.Errorf("failed to create HTTP request: %w", err)
473 }
474
475 req.Header.Set("Content-Type", "application/json")
476 req.Header.Set("Authorization", "Bearer "+s.AccessJwt)
477
478 s.Logger.Debug("putting record", "url", url, "collection", collection)
479
480 client := &http.Client{Timeout: 30 * time.Second}
481 resp, err := client.Do(req)
482 if err != nil {
483 return fmt.Errorf("failed to add user to default service: %w", err)
484 }
485 defer resp.Body.Close()
486
487 if resp.StatusCode != http.StatusOK {
488 var errorResponse map[string]any
489 if err := json.NewDecoder(resp.Body).Decode(&errorResponse); err != nil {
490 return fmt.Errorf("failed to add user to default service: HTTP %d (failed to decode error response: %w)", resp.StatusCode, err)
491 }
492 return fmt.Errorf("failed to add user to default service: HTTP %d, response: %v", resp.StatusCode, errorResponse)
493 }
494
495 return nil
496}
497
498// autoClaimTnglShDomain checks if the user has a .tngl.sh handle and, if so,
499// ensures their corresponding sites domain is claimed. This is idempotent —
500// ClaimDomain is a no-op if the claim already exists.
501func (o *OAuth) autoClaimTnglShDomain(did string) {
502 l := o.Logger.With("did", did)
503
504 pdsDomain := strings.TrimPrefix(o.Config.Pds.Host, "https://")
505 pdsDomain = strings.TrimPrefix(pdsDomain, "http://")
506
507 resolved, err := o.IdResolver.ResolveIdent(context.Background(), did)
508 if err != nil {
509 l.Error("autoClaimTnglShDomain: failed to resolve ident", "err", err)
510 return
511 }
512
513 handle := resolved.Handle.String()
514 if !strings.HasSuffix(handle, "."+pdsDomain) {
515 return
516 }
517
518 if err := db.ClaimDomain(o.Db, did, handle); err != nil {
519 l.Warn("autoClaimTnglShDomain: failed to claim domain", "domain", handle, "err", err)
520 } else {
521 l.Info("autoClaimTnglShDomain: claimed domain", "domain", handle)
522 }
523}
524
525// getAppPasswordSession returns a cached AppPasswordSession, creating one if needed.
526func (o *OAuth) getAppPasswordSession() (*AppPasswordSession, error) {
527 o.appPasswordSessionMu.Lock()
528 defer o.appPasswordSessionMu.Unlock()
529
530 if o.appPasswordSession != nil {
531 return o.appPasswordSession, nil
532 }
533
534 session, err := CreateAppPasswordSession(o.IdResolver, o.Config.Core.AppPassword, consts.TangledDid, o.Logger)
535 if err != nil {
536 return nil, err
537 }
538
539 o.appPasswordSession = session
540 return session, nil
541}