package oauth import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "log/slog" "net/http" "net/url" "slices" "strings" "sync" "time" comatproto "github.com/bluesky-social/indigo/api/atproto" "github.com/bluesky-social/indigo/atproto/atcrypto" "github.com/bluesky-social/indigo/atproto/auth" indigooauth "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/identity" "github.com/bluesky-social/indigo/atproto/syntax" "github.com/go-chi/chi/v5" tangled "tangled.org/core/api/tangled" "tangled.org/core/migrator/config" "tangled.org/core/migrator/db" ) var ( ErrGrantRequired = errors.New("oauth grant required") ErrRepoExists = errors.New("repository with this name already exists on the knot") ) var scopes = []string{ "atproto", "rpc:sh.tangled.repo.create?aud=*", "rpc:sh.tangled.repo.describeRepo?aud=*", "repo:sh.tangled.repo?action=create", } type Client struct { app *indigooauth.ClientApp store *Store cfg *config.Config privateKey atcrypto.PrivateKey http *http.Client logger *slog.Logger ownerLocks sync.Map sessions sync.Map } func NewClient(cfg *config.Config, database *db.DB, directory identity.Directory, logger *slog.Logger) (*Client, error) { privateKey, err := atcrypto.ParsePrivateMultibase(cfg.PrivateKey) if err != nil { return nil, fmt.Errorf("parsing service key: %w", err) } store := NewStore(database, cfg.ParsedMasterKey) clientID := strings.TrimSuffix(cfg.ServiceURL(), "/") + "/oauth-client-metadata.json" callback := strings.TrimSuffix(cfg.ServiceURL(), "/") + "/oauth/callback" oauthConfig := indigooauth.NewPublicConfig(clientID, callback, scopes) app := indigooauth.NewClientApp(&oauthConfig, store) app.Dir = directory return &Client{ app: app, store: store, cfg: cfg, privateKey: privateKey, http: &http.Client{Timeout: 30 * time.Second}, logger: logger, }, nil } func (c *Client) Routes() http.Handler { r := chi.NewRouter() r.Get("/oauth-client-metadata.json", c.handleMetadata) r.Get("/oauth/start", c.handleStart) r.Get("/oauth/callback", c.handleCallback) return r } func (c *Client) HasSession(ctx context.Context, did string) bool { parsed, err := syntax.ParseDID(did) if err != nil { return false } session, err := c.store.GetSession(ctx, parsed, "") return err == nil && hasScopes(session.Scopes, scopes) } func hasScopes(granted, required []string) bool { for _, req := range required { if !slices.Contains(granted, req) { return false } } return true } func (c *Client) handleMetadata(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "application/json") w.Header().Set("Cache-Control", "no-store") _ = json.NewEncoder(w).Encode(c.app.Config.ClientMetadata()) } const ( returnToCookie = "migrator_return_to" returnToMaxAge = 10 * 60 ) func setReturnToCookie(w http.ResponseWriter, value string, maxAge int) { http.SetCookie(w, &http.Cookie{ Name: returnToCookie, Value: url.PathEscape(value), Path: "/oauth", MaxAge: maxAge, HttpOnly: true, Secure: true, SameSite: http.SameSiteLaxMode, // the PDS bounces the browser back with a top-level GET }) } func sanitizeReturnTo(raw string) (string, bool) { if raw == "" || len(raw) > 2048 { return "", false } if !strings.HasPrefix(raw, "/") || strings.HasPrefix(raw, "//") || strings.Contains(raw, "\\") { return "", false } if strings.IndexFunc(raw, func(r rune) bool { return r < ' ' || r == 0x7f }) >= 0 { return "", false } return raw, true } func (c *Client) returnTarget(w http.ResponseWriter, r *http.Request, errSignal string) string { path := config.DefaultReturnPath if cookie, err := r.Cookie(returnToCookie); err == nil { if value, err := url.PathUnescape(cookie.Value); err == nil { if safe, ok := sanitizeReturnTo(value); ok { path = safe } } setReturnToCookie(w, "", -1) } sep := "?" if strings.Contains(path, "?") { sep = "&" } if errSignal != "" { path += sep + "oauth_error=" + errSignal } else { path += sep + "migrator=granted" } return strings.TrimSuffix(c.cfg.AppURL, "/") + path } func (c *Client) handleStart(w http.ResponseWriter, r *http.Request) { did := r.URL.Query().Get("did") if _, err := syntax.ParseDID(did); err != nil { http.Error(w, "a valid did query parameter is required", http.StatusBadRequest) return } if safe, ok := sanitizeReturnTo(r.URL.Query().Get("return_to")); ok { setReturnToCookie(w, safe, returnToMaxAge) } redirect, err := c.app.StartAuthFlow(r.Context(), did) if err != nil { c.logger.Error("starting oauth grant", "did", did) http.Error(w, "could not start oauth grant", http.StatusBadGateway) return } http.Redirect(w, r, redirect, http.StatusFound) } func (c *Client) handleCallback(w http.ResponseWriter, r *http.Request) { sessData, err := c.app.ProcessCallback(r.Context(), r.URL.Query()) errSignal := "" if err != nil { c.logger.Error("completing oauth grant", "err", err) errSignal = "grant_failed" } else if sessData != nil { // fresh grant rotates session id; evict cache so worker cannot replay revoked token c.sessions.Delete(sessData.AccountDID.String()) } http.Redirect(w, r, c.returnTarget(w, r, errSignal), http.StatusSeeOther) } // serializes session access per owner because refresh tokens are single-use; concurrent refreshes kill the token family func (c *Client) withOwnerSession(ctx context.Context, ownerDid string, fn func(*indigooauth.ClientSession) error) error { lock, _ := c.ownerLocks.LoadOrStore(ownerDid, new(sync.Mutex)) mu := lock.(*sync.Mutex) mu.Lock() defer mu.Unlock() session, err := c.session(ctx, ownerDid) if err != nil { return err } if err := grantErr(fn(session)); err != nil { if errors.Is(err, ErrGrantRequired) { c.dropSession(ctx, ownerDid) } return err } return nil } // auth server already revoked the token family, so upstream revoke is omitted func (c *Client) dropSession(ctx context.Context, ownerDid string) { c.sessions.Delete(ownerDid) did, err := syntax.ParseDID(ownerDid) if err != nil { return } if err := c.store.DeleteSession(ctx, did, ""); err != nil { c.logger.Error("deleting refused oauth session", "did", ownerDid, "err", err) } } // refresh happens lazily in DoWithAuth on 401; eager refresh would burn single-use tokens func (c *Client) session(ctx context.Context, ownerDid string) (*indigooauth.ClientSession, error) { if cached, ok := c.sessions.Load(ownerDid); ok { current := cached.(*indigooauth.ClientSession) // avoids a lost eviction race reviving stale tokens if rowID, err := c.store.SessionID(ctx, ownerDid); err == nil && rowID == current.Data.SessionID { return current, nil } c.sessions.Delete(ownerDid) } did, err := syntax.ParseDID(ownerDid) if err != nil { return nil, err } sessionID, err := c.store.SessionID(ctx, ownerDid) if errors.Is(err, db.ErrNotFound) { return nil, ErrGrantRequired } if err != nil { return nil, err } session, err := c.app.ResumeSession(ctx, did, sessionID) if err != nil { return nil, fmt.Errorf("resuming oauth session: %w", err) } if !hasScopes(session.Data.Scopes, scopes) { return nil, ErrGrantRequired } c.sessions.Store(ownerDid, session) return session, nil } func (c *Client) UsableGrant(ctx context.Context, ownerDid string) error { return c.withOwnerSession(ctx, ownerDid, func(session *indigooauth.ClientSession) error { if _, err := comatproto.ServerGetServiceAuth(ctx, session.APIClient(), c.cfg.ServiceDid.String(), time.Now().Add(time.Minute).Unix(), tangled.RepoDescribeRepoNSID); err != nil { return fmt.Errorf("minting user service token: %w", err) } return nil }) } // invalid_grant is terminal because the auth server invalidates the entire token family func grantErr(err error) error { if err == nil { return nil } if strings.Contains(err.Error(), "invalid_grant") { return fmt.Errorf("%w: %s", ErrGrantRequired, err) } return err } func (c *Client) userKnotToken(ctx context.Context, ownerDid, knotDid, lxm string) (string, error) { var token string err := c.withOwnerSession(ctx, ownerDid, func(session *indigooauth.ClientSession) error { response, err := comatproto.ServerGetServiceAuth(ctx, session.APIClient(), knotDid, time.Now().Add(time.Minute).Unix(), lxm) if err != nil { return fmt.Errorf("minting user service token for %s: %w", lxm, err) } token = response.Token return nil }) return token, err } func (c *Client) CreateRepo(ctx context.Context, ownerDid, knotDid, rkey, name, sourceURL string) (string, error) { token, err := c.userKnotToken(ctx, ownerDid, knotDid, tangled.RepoCreateNSID) if err != nil { return "", err } endpoint, err := knotEndpoint(knotDid, tangled.RepoCreateNSID) if err != nil { return "", err } body, err := json.Marshal(&tangled.RepoCreate_Input{ Rkey: rkey, Name: name, Source: &tangled.RepoCreate_Source{ Url: sourceURL, Kind: "import", }, }) if err != nil { return "", err } req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) if err != nil { return "", err } req.Header.Set("Content-Type", "application/json") req.Header.Set("Authorization", "Bearer "+token) resp, err := c.http.Do(req) if err != nil { return "", fmt.Errorf("calling knot create: %w", err) } defer resp.Body.Close() if resp.StatusCode == http.StatusConflict { return "", fmt.Errorf("%w: %s", ErrRepoExists, responseError("knot create", resp)) } if resp.StatusCode < 200 || resp.StatusCode >= 300 { return "", responseError("knot create", resp) } var out tangled.RepoCreate_Output if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(&out); err != nil { return "", fmt.Errorf("decoding knot create response: %w", err) } if out.RepoDid == nil || *out.RepoDid == "" { return "", errors.New("knot create response omitted repoDid") } return *out.RepoDid, nil } func (c *Client) Content(ctx context.Context, knotDid, repoDid string) (string, error) { lxm, _ := syntax.ParseNSID(tangled.RepoDescribeRepoNSID) token, err := auth.SignServiceAuth(c.cfg.ServiceDid, knotDid, time.Minute, &lxm, c.privateKey) if err != nil { return "", fmt.Errorf("minting describeRepo service token: %w", err) } var out tangled.RepoDescribeRepo_Output if err := c.describeRepo(ctx, token, knotDid, repoDid, &out); err != nil { return "", err } if out.Content == nil { return "", nil } return *out.Content, nil } func (c *Client) DescribeRepo(ctx context.Context, ownerDid, knotDid, repoDid string) (string, *string, error) { token, err := c.userKnotToken(ctx, ownerDid, knotDid, tangled.RepoDescribeRepoNSID) if err != nil { return "", nil, fmt.Errorf("minting describeRepo service token: %w", err) } var out tangled.RepoDescribeRepo_Output if err := c.describeRepo(ctx, token, knotDid, repoDid, &out); err != nil { return "", nil, err } if out.Content == nil { return "", out.Reason, nil } return *out.Content, out.Reason, nil } func (c *Client) describeRepo(ctx context.Context, token, knotDid, repoDid string, out *tangled.RepoDescribeRepo_Output) error { endpoint, err := knotEndpoint(knotDid, tangled.RepoDescribeRepoNSID) if err != nil { return err } u, _ := url.Parse(endpoint) query := u.Query() query.Set("repoDid", repoDid) u.RawQuery = query.Encode() req, err := http.NewRequestWithContext(ctx, http.MethodGet, u.String(), nil) if err != nil { return err } req.Header.Set("Authorization", "Bearer "+token) resp, err := c.http.Do(req) if err != nil { return fmt.Errorf("calling knot describeRepo: %w", err) } defer resp.Body.Close() if resp.StatusCode < 200 || resp.StatusCode >= 300 { return responseError("knot describeRepo", resp) } if err := json.NewDecoder(io.LimitReader(resp.Body, 1<<20)).Decode(out); err != nil { return fmt.Errorf("decoding knot describeRepo response: %w", err) } return nil } func (c *Client) PutRepoRecord(ctx context.Context, ownerDid, rkey, name, description, knotDid, repoDid string) error { did, err := syntax.ParseDID(ownerDid) if err != nil { return fmt.Errorf("parsing owner did: %w", err) } return c.withOwnerSession(ctx, ownerDid, func(session *indigooauth.ClientSession) error { return c.putRepoRecord(ctx, session, did, ownerDid, rkey, name, description, knotDid, repoDid) }) } func (c *Client) putRepoRecord(ctx context.Context, session *indigooauth.ClientSession, did syntax.DID, ownerDid, rkey, name, description, knotDid, repoDid string) error { doc, err := c.app.Dir.LookupDID(ctx, did) if err != nil { return fmt.Errorf("resolving owner did %q: %w", ownerDid, err) } pds := doc.PDSEndpoint() if pds == "" { return fmt.Errorf("owner did %q publishes no pds", ownerDid) } host, err := knotHost(knotDid) if err != nil { return err } body, err := json.Marshal(map[string]any{ "repo": ownerDid, "collection": tangled.RepoNSID, "rkey": rkey, "record": repoRecord(rkey, name, description, host, repoDid, time.Now().UTC()), }) if err != nil { return err } req, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimSuffix(pds, "/")+"/xrpc/"+"com.atproto.repo.createRecord", bytes.NewReader(body)) if err != nil { return err } req.Header.Set("Content-Type", "application/json") lxm, err := syntax.ParseNSID("com.atproto.repo.createRecord") if err != nil { return err } resp, err := session.DoWithAuth(c.http, req, lxm) if err != nil { return fmt.Errorf("writing repo record: %w", err) } defer resp.Body.Close() if resp.StatusCode >= 200 && resp.StatusCode < 300 { return nil } if resp.StatusCode == http.StatusBadRequest || resp.StatusCode == http.StatusConflict { // The record already exists; retry swaps lost to concurrent writers. for attempt := 0; attempt < swapConflictAttempts; attempt++ { existing, found := c.existingRepoRecord(ctx, session, pds, ownerDid, rkey) if !found { break } if existing.repoDid == repoDid { c.logger.Info("repo record was already written by an earlier attempt", "rkey", rkey, "repoDid", repoDid) return nil } if existing.repoDid != "" { return fmt.Errorf("a record for %q already names repository %s", rkey, existing.repoDid) } retry, err := c.patchRepoRecord(ctx, session, pds, ownerDid, rkey, repoDid, existing) if err != nil { return err } if !retry { return nil } } } return responseError("repo record write", resp) } // knot field stores the knot's host, not its DID; a DID breaks knot and mirror reads func knotHost(knotDid string) (string, error) { const prefix = "did:web:" if !strings.HasPrefix(knotDid, prefix) { return "", fmt.Errorf("knot %q is not a did:web DID", knotDid) } // ports are percent-escaped in did:web hosts host, err := url.PathUnescape(strings.TrimPrefix(knotDid, prefix)) if err != nil { return "", fmt.Errorf("decoding host from knot %q: %w", knotDid, err) } return host, nil } // Match the label definitions subscribed by the app for new repos. const defaultLabelOwner = "did:plc:wshs7t2adsemcrrd4snkeqli" var defaultLabels = []string{ "at://" + defaultLabelOwner + "/sh.tangled.label.definition/wontfix", "at://" + defaultLabelOwner + "/sh.tangled.label.definition/good-first-issue", "at://" + defaultLabelOwner + "/sh.tangled.label.definition/duplicate", "at://" + defaultLabelOwner + "/sh.tangled.label.definition/documentation", "at://" + defaultLabelOwner + "/sh.tangled.label.definition/assignee", } const swapConflictAttempts = 3 func (c *Client) patchRepoRecord(ctx context.Context, session *indigooauth.ClientSession, pds, ownerDid, rkey, repoDid string, existing *existingRepoRecord) (bool, error) { record := existing.value if record == nil { record = make(map[string]any) } record["repoDid"] = repoDid putBody, err := json.Marshal(map[string]any{ "repo": ownerDid, "collection": tangled.RepoNSID, "rkey": rkey, "record": record, "swapRecord": existing.cid, }) if err != nil { return false, err } putReq, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimSuffix(pds, "/")+"/xrpc/com.atproto.repo.putRecord", bytes.NewReader(putBody)) if err != nil { return false, err } putReq.Header.Set("Content-Type", "application/json") putLxm, err := syntax.ParseNSID("com.atproto.repo.putRecord") if err != nil { return false, err } putResp, err := session.DoWithAuth(c.http, putReq, putLxm) if err != nil { return false, fmt.Errorf("updating repo record: %w", err) } defer putResp.Body.Close() if putResp.StatusCode >= 200 && putResp.StatusCode < 300 { return false, nil } if putResp.StatusCode == http.StatusBadRequest || putResp.StatusCode == http.StatusConflict { return true, nil } return false, responseError("repo record update", putResp) } func repoRecord(rkey, name, description, knotHost, repoDid string, now time.Time) *tangled.Repo { record := &tangled.Repo{ LexiconTypeID: tangled.RepoNSID, CreatedAt: now.UTC().Format(time.RFC3339), Knot: knotHost, RepoDid: &repoDid, Labels: defaultLabels, } // Omit name when identical to rkey to match the app's record shape. if trimmed := strings.TrimSpace(name); trimmed != "" && trimmed != strings.ToLower(strings.TrimSpace(rkey)) { record.Name = &trimmed } if trimmed := strings.TrimSpace(description); trimmed != "" { record.Description = &trimmed } return record } type existingRepoRecord struct { cid string repoDid string value map[string]any } func (c *Client) existingRepoRecord(ctx context.Context, session *indigooauth.ClientSession, pds, ownerDid, rkey string) (*existingRepoRecord, bool) { query := url.Values{} query.Set("repo", ownerDid) query.Set("collection", tangled.RepoNSID) query.Set("rkey", rkey) req, err := http.NewRequestWithContext(ctx, http.MethodGet, strings.TrimSuffix(pds, "/")+"/xrpc/com.atproto.repo.getRecord?"+query.Encode(), nil) if err != nil { return nil, false } lxm, err := syntax.ParseNSID("com.atproto.repo.getRecord") if err != nil { return nil, false } resp, err := session.DoWithAuth(c.http, req, lxm) if err != nil { return nil, false } defer resp.Body.Close() if resp.StatusCode < 200 || resp.StatusCode >= 300 { return nil, false } var out struct { Cid string `json:"cid"` Value map[string]any `json:"value"` } if err := json.NewDecoder(io.LimitReader(resp.Body, 64<<10)).Decode(&out); err != nil { return nil, false } var repoDid string if d, ok := out.Value["repoDid"].(string); ok { repoDid = d } return &existingRepoRecord{ cid: out.Cid, repoDid: repoDid, value: out.Value, }, true } func knotEndpoint(did, nsid string) (string, error) { if !strings.HasPrefix(did, "did:web:") { return "", fmt.Errorf("knot DID must be did:web: %q", did) } parts := strings.Split(strings.TrimPrefix(did, "did:web:"), ":") host, err := url.PathUnescape(parts[0]) if err != nil || host == "" { return "", fmt.Errorf("invalid knot DID %q", did) } base := "https://" + host if len(parts) > 1 { base += "/" + strings.Join(parts[1:], "/") } return strings.TrimSuffix(base, "/") + "/xrpc/" + nsid, nil } func responseError(operation string, resp *http.Response) error { body, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<10)) return fmt.Errorf("%s returned HTTP %d: %s", operation, resp.StatusCode, strings.TrimSpace(string(body))) }