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