This repository has no description
0

Configure Feed

Select the types of activity you want to include in your feed.

core / appview / oauth / handler.go
18 kB 636 lines
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}