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 625 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 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}