This repository has no description
1package oauth
2
3import (
4 "bytes"
5 "context"
6 "encoding/json"
7 "errors"
8 "fmt"
9 "log/slog"
10 "net/http"
11 "slices"
12 "strings"
13 "time"
14
15 comatproto "github.com/bluesky-social/indigo/api/atproto"
16 atpclient "github.com/bluesky-social/indigo/atproto/atclient"
17 "github.com/bluesky-social/indigo/atproto/auth/oauth"
18 lexutil "github.com/bluesky-social/indigo/lex/util"
19 xrpc "github.com/bluesky-social/indigo/xrpc"
20 "github.com/go-chi/chi/v5"
21 "github.com/posthog/posthog-go"
22 "tangled.org/core/api/tangled"
23 "tangled.org/core/appview/db"
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
31func (o *OAuth) Router() http.Handler {
32 r := chi.NewRouter()
33
34 r.Get("/oauth/client-metadata.json", o.clientMetadata)
35 r.Get("/oauth/jwks.json", o.jwks)
36 r.Get("/oauth/callback", o.callback)
37 return r
38}
39
40func (o *OAuth) clientMetadata(w http.ResponseWriter, r *http.Request) {
41 doc := o.ClientApp.Config.ClientMetadata()
42 doc.JWKSURI = &o.JwksUri
43 doc.ClientName = &o.ClientName
44 doc.ClientURI = &o.ClientUri
45 doc.Scope = doc.Scope + " identity:handle"
46
47 w.Header().Set("Content-Type", "application/json")
48 if err := json.NewEncoder(w).Encode(doc); err != nil {
49 http.Error(w, err.Error(), http.StatusInternalServerError)
50 return
51 }
52}
53
54func (o *OAuth) jwks(w http.ResponseWriter, r *http.Request) {
55 w.Header().Set("Content-Type", "application/json")
56 body := o.ClientApp.Config.PublicJWKS()
57 if err := json.NewEncoder(w).Encode(body); err != nil {
58 http.Error(w, err.Error(), http.StatusInternalServerError)
59 return
60 }
61}
62
63func (o *OAuth) callback(w http.ResponseWriter, r *http.Request) {
64 ctx := r.Context()
65 l := o.Logger.With("query", r.URL.Query())
66
67 authReturn := o.GetAuthReturn(r)
68 _ = o.ClearAuthReturn(w, r)
69
70 sessData, err := o.ClientApp.ProcessCallback(ctx, r.URL.Query())
71 if err != nil {
72 var callbackErr *oauth.AuthRequestCallbackError
73 if errors.As(err, &callbackErr) {
74 l.Debug("callback error", "err", callbackErr)
75 http.Redirect(w, r, fmt.Sprintf("/login?error=%s", callbackErr.ErrorCode), http.StatusFound)
76 return
77 }
78 l.Error("failed to process callback", "err", err)
79 http.Redirect(w, r, "/login?error=oauth", http.StatusFound)
80 return
81 }
82
83 if err := o.SaveSession(w, r, sessData); err != nil {
84 l.Error("failed to save session", "data", sessData, "err", err)
85 errorCode := "session"
86 if errors.Is(err, ErrMaxAccountsReached) {
87 errorCode = "max_accounts"
88 }
89 http.Redirect(w, r, fmt.Sprintf("/login?error=%s", errorCode), http.StatusFound)
90 return
91 }
92
93 o.Logger.Debug("session saved successfully")
94
95 go o.addToDefaultKnot(sessData.AccountDID.String())
96 go o.addToDefaultSpindle(sessData.AccountDID.String())
97 go o.ensureTangledProfile(sessData)
98 go o.autoClaimTnglShDomain(sessData.AccountDID.String())
99 go o.drainPdsRewrites(sessData)
100
101 if !o.Config.Core.Dev {
102 err = o.Posthog.Enqueue(posthog.Capture{
103 DistinctId: sessData.AccountDID.String(),
104 Event: "signin",
105 })
106 if err != nil {
107 o.Logger.Error("failed to enqueue posthog event", "err", err)
108 }
109 }
110
111 redirectURL := "/"
112 if authReturn.ReturnURL != "" {
113 redirectURL = authReturn.ReturnURL
114 }
115
116 if o.isAccountDeactivated(sessData) {
117 redirectURL = "/settings/profile"
118 }
119
120 http.Redirect(w, r, redirectURL, http.StatusFound)
121}
122
123func (o *OAuth) isAccountDeactivated(sessData *oauth.ClientSessionData) bool {
124 pdsClient := &xrpc.Client{
125 Host: sessData.HostURL,
126 Client: &http.Client{Timeout: 5 * time.Second},
127 }
128
129 _, err := comatproto.RepoDescribeRepo(
130 context.Background(),
131 pdsClient,
132 sessData.AccountDID.String(),
133 )
134 if err == nil {
135 return false
136 }
137
138 var xrpcErr *xrpc.Error
139 var xrpcBody *xrpc.XRPCError
140 return errors.As(err, &xrpcErr) &&
141 errors.As(xrpcErr.Wrapped, &xrpcBody) &&
142 xrpcBody.ErrStr == "RepoDeactivated"
143}
144
145func (o *OAuth) addToDefaultSpindle(did string) {
146 l := o.Logger.With("subject", did)
147
148 // use the tangled.sh app password to get an accessJwt
149 // and create an sh.tangled.spindle.member record with that
150 spindleMembers, err := db.GetSpindleMembers(
151 o.Db,
152 orm.FilterEq("instance", "spindle.tangled.sh"),
153 orm.FilterEq("subject", did),
154 )
155 if err != nil {
156 l.Error("failed to get spindle members", "err", err)
157 return
158 }
159
160 if len(spindleMembers) != 0 {
161 l.Warn("already a member of the default spindle")
162 return
163 }
164
165 l.Debug("adding to default spindle")
166 session, err := o.getAppPasswordSession()
167 if err != nil {
168 l.Error("failed to create session", "err", err)
169 return
170 }
171
172 record := tangled.SpindleMember{
173 LexiconTypeID: tangled.SpindleMemberNSID,
174 Subject: did,
175 Instance: consts.DefaultSpindle,
176 CreatedAt: time.Now().Format(time.RFC3339),
177 }
178
179 if err := session.putRecord(record, tangled.SpindleMemberNSID); err != nil {
180 l.Error("failed to add to default spindle", "err", err)
181 return
182 }
183
184 l.Debug("successfully added to default spindle", "did", did)
185}
186
187func (o *OAuth) addToDefaultKnot(did string) {
188 l := o.Logger.With("subject", did)
189
190 // use the tangled.sh app password to get an accessJwt
191 // and create an sh.tangled.spindle.member record with that
192
193 allKnots, err := o.Enforcer.GetKnotsForUser(did)
194 if err != nil {
195 l.Error("failed to get knot members for did", "err", err)
196 return
197 }
198
199 if slices.Contains(allKnots, consts.DefaultKnot) {
200 l.Warn("already a member of the default knot")
201 return
202 }
203
204 l.Debug("adding to default knot")
205 session, err := o.getAppPasswordSession()
206 if err != nil {
207 l.Error("failed to create session", "err", err)
208 return
209 }
210
211 record := tangled.KnotMember{
212 LexiconTypeID: tangled.KnotMemberNSID,
213 Subject: did,
214 Domain: consts.DefaultKnot,
215 CreatedAt: time.Now().Format(time.RFC3339),
216 }
217
218 if err := session.putRecord(record, tangled.KnotMemberNSID); err != nil {
219 l.Error("failed to add to default knot", "err", err)
220 return
221 }
222
223 if err := o.Enforcer.AddKnotMember(consts.DefaultKnot, did); err != nil {
224 l.Error("failed to set up enforcer rules", "err", err)
225 return
226 }
227
228 l.Debug("successfully added to default knot")
229}
230
231func (o *OAuth) ensureTangledProfile(sessData *oauth.ClientSessionData) {
232 ctx := context.Background()
233 did := sessData.AccountDID.String()
234 l := o.Logger.With("did", did)
235
236 profile, _ := db.GetProfile(o.Db, did)
237 if profile != nil {
238 l.Debug("profile already exists in DB")
239 return
240 }
241
242 l.Debug("creating empty Tangled profile")
243
244 sess, err := o.ClientApp.ResumeSession(ctx, sessData.AccountDID, sessData.SessionID)
245 if err != nil {
246 l.Error("failed to resume session for profile creation", "err", err)
247 return
248 }
249 client := sess.APIClient()
250
251 _, err = comatproto.RepoPutRecord(ctx, client, &comatproto.RepoPutRecord_Input{
252 Collection: tangled.ActorProfileNSID,
253 Repo: did,
254 Rkey: "self",
255 Record: &lexutil.LexiconTypeDecoder{Val: &tangled.ActorProfile{}},
256 })
257
258 if err != nil {
259 l.Error("failed to create empty profile on PDS", "err", err)
260 return
261 }
262
263 tx, err := o.Db.BeginTx(ctx, nil)
264 if err != nil {
265 l.Error("failed to start transaction", "err", err)
266 return
267 }
268
269 emptyProfile := &models.Profile{Did: did}
270 if err := db.UpsertProfile(tx, emptyProfile); err != nil {
271 l.Error("failed to create empty profile in DB", "err", err)
272 return
273 }
274
275 l.Debug("successfully created empty Tangled profile on PDS and DB")
276}
277
278func (o *OAuth) drainPdsRewrites(sessData *oauth.ClientSessionData) {
279 ctx := context.Background()
280 did := sessData.AccountDID.String()
281 l := o.Logger.With("did", did, "handler", "drainPdsRewrites")
282
283 rewrites, err := db.GetPendingPdsRewrites(o.Db, did)
284 if err != nil {
285 l.Error("failed to get pending rewrites", "err", err)
286 return
287 }
288 if len(rewrites) == 0 {
289 return
290 }
291
292 l.Info("draining pending PDS rewrites", "count", len(rewrites))
293
294 sess, err := o.ClientApp.ResumeSession(ctx, sessData.AccountDID, sessData.SessionID)
295 if err != nil {
296 l.Error("failed to resume session for PDS rewrites", "err", err)
297 return
298 }
299 client := sess.APIClient()
300
301 for _, rw := range rewrites {
302 if err := o.rewritePdsRecord(ctx, client, did, rw); err != nil {
303 l.Error("failed to rewrite PDS record",
304 "nsid", rw.RecordNsid,
305 "rkey", rw.RecordRkey,
306 "repo_did", rw.RepoDid,
307 "err", err)
308 continue
309 }
310
311 if err := db.CompletePdsRewrite(o.Db, rw.Id); err != nil {
312 l.Error("failed to mark rewrite complete", "id", rw.Id, "err", err)
313 }
314 }
315}
316
317func (o *OAuth) rewritePdsRecord(ctx context.Context, client *atpclient.APIClient, userDid string, rw db.PdsRewrite) error {
318 ex, err := comatproto.RepoGetRecord(ctx, client, "", rw.RecordNsid, userDid, rw.RecordRkey)
319 if err != nil {
320 return fmt.Errorf("get record: %w", err)
321 }
322
323 val := ex.Value.Val
324 repoDid := rw.RepoDid
325
326 switch rw.RecordNsid {
327 case tangled.RepoNSID:
328 rec, ok := val.(*tangled.Repo)
329 if !ok {
330 return fmt.Errorf("unexpected type for repo record")
331 }
332 rec.RepoDid = &repoDid
333
334 case tangled.RepoIssueNSID:
335 rec, ok := val.(*tangled.RepoIssue)
336 if !ok {
337 return fmt.Errorf("unexpected type for issue record")
338 }
339 rec.RepoDid = &repoDid
340
341 case tangled.RepoPullNSID:
342 rec, ok := val.(*tangled.RepoPull)
343 if !ok {
344 return fmt.Errorf("unexpected type for pull record")
345 }
346 if rec.Target != nil {
347 rec.Target.RepoDid = &repoDid
348 }
349 if rec.Source != nil && rec.Source.Repo != nil && *rec.Source.Repo == rw.OldRepoAt {
350 rec.Source.RepoDid = &repoDid
351 }
352
353 case tangled.RepoCollaboratorNSID:
354 rec, ok := val.(*tangled.RepoCollaborator)
355 if !ok {
356 return fmt.Errorf("unexpected type for collaborator record")
357 }
358 rec.RepoDid = &repoDid
359
360 case tangled.RepoArtifactNSID:
361 rec, ok := val.(*tangled.RepoArtifact)
362 if !ok {
363 return fmt.Errorf("unexpected type for artifact record")
364 }
365 rec.RepoDid = &repoDid
366
367 case tangled.FeedStarNSID:
368 rec, ok := val.(*tangled.FeedStar)
369 if !ok {
370 return fmt.Errorf("unexpected type for star record")
371 }
372 rec.SubjectDid = &repoDid
373
374 case tangled.ActorProfileNSID:
375 rec, ok := val.(*tangled.ActorProfile)
376 if !ok {
377 return fmt.Errorf("unexpected type for profile record")
378 }
379 rewritten := make([]string, 0, len(rec.PinnedRepositories))
380 for _, pin := range rec.PinnedRepositories {
381 if strings.HasPrefix(pin, "did:") {
382 rewritten = append(rewritten, pin)
383 continue
384 }
385 repo, repoErr := db.GetRepoByAtUri(o.Db, pin)
386 if repoErr != nil || repo.RepoDid == "" {
387 rewritten = append(rewritten, pin)
388 continue
389 }
390 rewritten = append(rewritten, repo.RepoDid)
391 }
392 rec.PinnedRepositories = rewritten
393
394 default:
395 return fmt.Errorf("unsupported NSID for PDS rewrite: %s", rw.RecordNsid)
396 }
397
398 _, err = comatproto.RepoPutRecord(ctx, client, &comatproto.RepoPutRecord_Input{
399 Collection: rw.RecordNsid,
400 Repo: userDid,
401 Rkey: rw.RecordRkey,
402 SwapRecord: ex.Cid,
403 Record: &lexutil.LexiconTypeDecoder{Val: val},
404 })
405 if err != nil {
406 return fmt.Errorf("put record: %w", err)
407 }
408
409 return nil
410}
411
412// create a AppPasswordSession using apppasswords
413type AppPasswordSession struct {
414 AccessJwt string `json:"accessJwt"`
415 RefreshJwt string `json:"refreshJwt"`
416 PdsEndpoint string
417 Did string
418 Logger *slog.Logger
419 ExpiresAt time.Time
420}
421
422func CreateAppPasswordSession(res *idresolver.Resolver, appPassword, did string, logger *slog.Logger) (*AppPasswordSession, error) {
423 if appPassword == "" {
424 return nil, fmt.Errorf("no app password configured")
425 }
426
427 resolved, err := res.ResolveIdent(context.Background(), did)
428 if err != nil {
429 return nil, fmt.Errorf("failed to resolve tangled.sh DID %s: %v", did, err)
430 }
431
432 pdsEndpoint := resolved.PDSEndpoint()
433 if pdsEndpoint == "" {
434 return nil, fmt.Errorf("no PDS endpoint found for tangled.sh DID %s", did)
435 }
436
437 sessionPayload := map[string]string{
438 "identifier": did,
439 "password": appPassword,
440 }
441 sessionBytes, err := json.Marshal(sessionPayload)
442 if err != nil {
443 return nil, fmt.Errorf("failed to marshal session payload: %v", err)
444 }
445
446 sessionURL := pdsEndpoint + "/xrpc/com.atproto.server.createSession"
447 sessionReq, err := http.NewRequestWithContext(context.Background(), "POST", sessionURL, bytes.NewBuffer(sessionBytes))
448 if err != nil {
449 return nil, fmt.Errorf("failed to create session request: %v", err)
450 }
451 sessionReq.Header.Set("Content-Type", "application/json")
452
453 logger.Debug("creating app password session", "url", sessionURL, "headers", sessionReq.Header)
454
455 client := &http.Client{Timeout: 30 * time.Second}
456 sessionResp, err := client.Do(sessionReq)
457 if err != nil {
458 return nil, fmt.Errorf("failed to create session: %v", err)
459 }
460 defer sessionResp.Body.Close()
461
462 if sessionResp.StatusCode != http.StatusOK {
463 return nil, fmt.Errorf("failed to create session: HTTP %d", sessionResp.StatusCode)
464 }
465
466 var session AppPasswordSession
467 if err := json.NewDecoder(sessionResp.Body).Decode(&session); err != nil {
468 return nil, fmt.Errorf("failed to decode session response: %v", err)
469 }
470
471 session.PdsEndpoint = pdsEndpoint
472 session.Did = did
473 session.Logger = logger
474 session.ExpiresAt = time.Now().Add(115 * time.Minute)
475
476 return &session, nil
477}
478
479func (s *AppPasswordSession) refreshSession() error {
480 refreshURL := s.PdsEndpoint + "/xrpc/com.atproto.server.refreshSession"
481 req, err := http.NewRequestWithContext(context.Background(), "POST", refreshURL, nil)
482 if err != nil {
483 return fmt.Errorf("failed to create refresh request: %w", err)
484 }
485
486 req.Header.Set("Authorization", "Bearer "+s.RefreshJwt)
487
488 s.Logger.Debug("refreshing app password session", "url", refreshURL)
489
490 client := &http.Client{Timeout: 30 * time.Second}
491 resp, err := client.Do(req)
492 if err != nil {
493 return fmt.Errorf("failed to refresh session: %w", err)
494 }
495 defer resp.Body.Close()
496
497 if resp.StatusCode != http.StatusOK {
498 var errorResponse map[string]any
499 if err := json.NewDecoder(resp.Body).Decode(&errorResponse); err != nil {
500 return fmt.Errorf("failed to refresh session: HTTP %d (failed to decode error response: %w)", resp.StatusCode, err)
501 }
502 errorBytes, _ := json.Marshal(errorResponse)
503 return fmt.Errorf("failed to refresh session: HTTP %d, response: %s", resp.StatusCode, string(errorBytes))
504 }
505
506 var refreshResponse struct {
507 AccessJwt string `json:"accessJwt"`
508 RefreshJwt string `json:"refreshJwt"`
509 }
510 if err := json.NewDecoder(resp.Body).Decode(&refreshResponse); err != nil {
511 return fmt.Errorf("failed to decode refresh response: %w", err)
512 }
513
514 s.AccessJwt = refreshResponse.AccessJwt
515 s.RefreshJwt = refreshResponse.RefreshJwt
516 // Set new expiry time with 5 minute buffer
517 s.ExpiresAt = time.Now().Add(115 * time.Minute)
518
519 s.Logger.Debug("successfully refreshed app password session")
520 return nil
521}
522
523func (s *AppPasswordSession) isValid() bool {
524 return time.Now().Before(s.ExpiresAt)
525}
526
527func (s *AppPasswordSession) putRecord(record any, collection string) error {
528 if !s.isValid() {
529 s.Logger.Debug("access token expired, refreshing session")
530 if err := s.refreshSession(); err != nil {
531 return fmt.Errorf("failed to refresh session: %w", err)
532 }
533 s.Logger.Debug("session refreshed")
534 }
535
536 recordBytes, err := json.Marshal(record)
537 if err != nil {
538 return fmt.Errorf("failed to marshal knot member record: %w", err)
539 }
540
541 payload := map[string]any{
542 "repo": s.Did,
543 "collection": collection,
544 "rkey": tid.TID(),
545 "record": json.RawMessage(recordBytes),
546 }
547
548 payloadBytes, err := json.Marshal(payload)
549 if err != nil {
550 return fmt.Errorf("failed to marshal request payload: %w", err)
551 }
552
553 url := s.PdsEndpoint + "/xrpc/com.atproto.repo.putRecord"
554 req, err := http.NewRequestWithContext(context.Background(), "POST", url, bytes.NewBuffer(payloadBytes))
555 if err != nil {
556 return fmt.Errorf("failed to create HTTP request: %w", err)
557 }
558
559 req.Header.Set("Content-Type", "application/json")
560 req.Header.Set("Authorization", "Bearer "+s.AccessJwt)
561
562 s.Logger.Debug("putting record", "url", url, "collection", collection)
563
564 client := &http.Client{Timeout: 30 * time.Second}
565 resp, err := client.Do(req)
566 if err != nil {
567 return fmt.Errorf("failed to add user to default service: %w", err)
568 }
569 defer resp.Body.Close()
570
571 if resp.StatusCode != http.StatusOK {
572 var errorResponse map[string]any
573 if err := json.NewDecoder(resp.Body).Decode(&errorResponse); err != nil {
574 return fmt.Errorf("failed to add user to default service: HTTP %d (failed to decode error response: %w)", resp.StatusCode, err)
575 }
576 return fmt.Errorf("failed to add user to default service: HTTP %d, response: %v", resp.StatusCode, errorResponse)
577 }
578
579 return nil
580}
581
582// autoClaimTnglShDomain checks if the user has a .tngl.sh handle and, if so,
583// ensures their corresponding sites domain is claimed. This is idempotent —
584// ClaimDomain is a no-op if the claim already exists.
585func (o *OAuth) autoClaimTnglShDomain(did string) {
586 l := o.Logger.With("did", did)
587
588 pdsDomain := strings.TrimPrefix(o.Config.Pds.Host, "https://")
589 pdsDomain = strings.TrimPrefix(pdsDomain, "http://")
590
591 resolved, err := o.IdResolver.ResolveIdent(context.Background(), did)
592 if err != nil {
593 l.Error("autoClaimTnglShDomain: failed to resolve ident", "err", err)
594 return
595 }
596
597 handle := resolved.Handle.String()
598 if !strings.HasSuffix(handle, "."+pdsDomain) {
599 return
600 }
601
602 if err := db.ClaimDomain(o.Db, did, handle); err != nil {
603 l.Warn("autoClaimTnglShDomain: failed to claim domain", "domain", handle, "err", err)
604 } else {
605 l.Info("autoClaimTnglShDomain: claimed domain", "domain", handle)
606 }
607}
608
609// getAppPasswordSession returns a cached AppPasswordSession, creating one if needed.
610func (o *OAuth) getAppPasswordSession() (*AppPasswordSession, error) {
611 o.appPasswordSessionMu.Lock()
612 defer o.appPasswordSessionMu.Unlock()
613
614 if o.appPasswordSession != nil {
615 return o.appPasswordSession, nil
616 }
617
618 session, err := CreateAppPasswordSession(o.IdResolver, o.Config.Core.AppPassword, consts.TangledDid, o.Logger)
619 if err != nil {
620 return nil, err
621 }
622
623 o.appPasswordSession = session
624 return session, nil
625}