package db import ( "context" "database/sql" "fmt" "maps" "slices" "sort" "strings" "time" "github.com/bluesky-social/indigo/atproto/syntax" "tangled.org/core/appview/models" "tangled.org/core/appview/pagination" "tangled.org/core/orm" ) func PutPull(ctx context.Context, tx *sql.Tx, pull *models.Pull, references []syntax.ATURI) error { // ensure sequence exists _, err := tx.ExecContext(ctx, ` insert or ignore into repo_pull_seqs (repo_did, next_pull_id) values (?, 1) `, pull.RepoDid) if err != nil { return err } var exists bool if err := tx.QueryRowContext(ctx, `select exists (select 1 from pulls where at_uri = ?)`, pull.AtUri(), ).Scan(&exists); err != nil { return err } if !exists { // assign new ID for a PR if err := tx.QueryRowContext(ctx, `update repo_pull_seqs set next_pull_id = next_pull_id + 1 where repo_did = ? returning next_pull_id - 1`, pull.RepoDid, ).Scan(&pull.PullId); err != nil { return err } } result, err := tx.ExecContext(ctx, `insert into pulls ( owner_did, rkey, cid, repo_did, pull_id, title, body, target_branch, source_repo_did, source_branch, created, state ) values (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) on conflict(at_uri) do update set cid = excluded.cid, repo_did = excluded.repo_did, title = excluded.title, body = excluded.body, target_branch = excluded.target_branch, source_repo_did = excluded.source_repo_did, source_branch = excluded.source_branch, created = excluded.created, state = excluded.state where pulls.cid is not excluded.cid`, pull.OwnerDid, pull.Rkey, pull.Cid, pull.RepoDid, pull.PullId, pull.Title, pull.Body, pull.TargetBranch, pull.SourceRepo, pull.SourceBranch, pull.Created.Format(time.RFC3339), pull.State, ) if err != nil { return fmt.Errorf("inserting pr: %w", err) } id, err := result.LastInsertId() if err != nil { return err } pull.ID = id // delete all existing versions if _, err := tx.ExecContext(ctx, `delete from pull_versions where pull_at = ?`, pull.AtUri(), ); err != nil { return fmt.Errorf("deleting old pr versions: %w", err) } // re-create all versions if len(pull.Versions) > 0 { pullAt := pull.AtUri() var sb strings.Builder sb.WriteString(`insert into pull_versions (pull_at, id, head, base, created) values `) args := make([]any, 0, len(pull.Versions)*5) for i, v := range pull.Versions { if i > 0 { sb.WriteString(", ") } sb.WriteString("(?, ?, ?, ?, ?)") args = append(args, pullAt, v.ID, v.Head, v.Base, v.Created.Format(time.RFC3339)) } if _, err := tx.ExecContext(ctx, sb.String(), args...); err != nil { return fmt.Errorf("inserting pr versions: %w", err) } } // update references when comment is updated if err := putReferences(tx, pull.AtUri(), references); err != nil { return fmt.Errorf("put reference_links: %w", err) } return nil } func SubmitPullVersion(ctx context.Context, q Execer, pullAt syntax.ATURI, version models.PullVersion) error { _, err := q.ExecContext(ctx, `insert into pull_versions (pull_at, id, head, base, created) values (?, ?, ?, ?, ?)`, pullAt, version.ID, version.Head, version.Base, version.Created.Format(time.RFC3339), ) return err } func GetPull(ctx context.Context, q Execer, filters ...orm.Filter) (*models.Pull, error) { pulls, err := GetPullsPaginated(ctx, q, pagination.Page{Limit: 1}, filters...) if err != nil { return nil, err } if len(pulls) == 0 { return nil, sql.ErrNoRows } return pulls[0], nil } func GetPullsPaginated(ctx context.Context, q Execer, page pagination.Page, filters ...orm.Filter) ([]*models.Pull, error) { pulls := make(map[syntax.ATURI]*models.Pull) var conditions []string var args []any for _, filter := range filters { conditions = append(conditions, filter.Condition()) args = append(args, filter.Arg()...) } whereClause := "" if conditions != nil { whereClause = " where " + strings.Join(conditions, " and ") } pageClause := "" if page.Limit != 0 { pageClause = fmt.Sprintf( " limit %d offset %d ", page.Limit, page.Offset, ) } query := fmt.Sprintf(` select id, owner_did, rkey, cid, repo_did, pull_id, title, body, target_branch, source_repo_did, source_branch, created, state from pulls %s order by created desc %s `, whereClause, pageClause) rows, err := q.QueryContext(ctx, query, args...) if err != nil { return nil, err } defer rows.Close() for rows.Next() { var pull models.Pull var createdAt string var sourceRepo, sourceBranch sql.NullString err := rows.Scan( &pull.ID, &pull.OwnerDid, &pull.Rkey, &pull.Cid, &pull.RepoDid, &pull.PullId, &pull.Title, &pull.Body, &pull.TargetBranch, &sourceRepo, &sourceBranch, &createdAt, &pull.State, ) if err != nil { return nil, fmt.Errorf("scanning row: %w", err) } createdTime, err := time.Parse(time.RFC3339, createdAt) if err != nil { return nil, fmt.Errorf("parsing created: %w", err) } pull.Created = createdTime if sourceRepo.Valid { pull.SourceRepo = syntax.DID(sourceRepo.String) } else { // fallback to pull.target.repo pull.SourceRepo = pull.RepoDid } if sourceBranch.Valid { pull.SourceBranch = &sourceBranch.String } pulls[pull.AtUri()] = &pull } if err := rows.Err(); err != nil { return nil, fmt.Errorf("scanning rows: %w", err) } pullAts := slices.Collect(maps.Keys(pulls)) versionsMap, err := ListVersions(ctx, q, pullAts) if err != nil { return nil, fmt.Errorf("querying versions: %w", err) } for pullAt, p := range pulls { if versions, ok := versionsMap[pullAt]; ok { p.Versions = versions } else { return nil, fmt.Errorf("find 0 versions for PR %s", pullAt) } } // collect reverse repos { repoDids := make([]string, 0, len(pulls)) for _, issue := range pulls { repoDids = append(repoDids, string(issue.RepoDid)) } repos, err := GetRepos(q, orm.FilterIn("repo_did", repoDids)) if err != nil { return nil, fmt.Errorf("failed to build repo mappings: %w", err) } repoMap := make(map[syntax.DID]*models.Repo) for i := range repos { repoMap[syntax.DID(repos[i].RepoDid)] = &repos[i] } for pullAt, p := range pulls { if r, ok := repoMap[p.RepoDid]; ok { p.Repo = r } else { delete(pulls, pullAt) } } } // collect allLabels for each PR { allLabels, err := GetLabels(q, orm.FilterIn("subject", pullAts)) if err != nil { return nil, fmt.Errorf("failed to query labels: %w", err) } for pullAt, labels := range allLabels { if pull, ok := pulls[pullAt]; ok { pull.Labels = labels } } } orderedById := []*models.Pull{} for _, p := range pulls { orderedById = append(orderedById, p) } sort.Slice(orderedById, func(i, j int) bool { return orderedById[i].PullId > orderedById[j].PullId }) return orderedById, nil } // mapping from pull -> pull submissions func ListVersions(ctx context.Context, q Execer, pullAts []syntax.ATURI) (map[syntax.ATURI][]models.PullVersion, error) { filter := orm.FilterIn("pull_at", pullAts) query := fmt.Sprintf(` select pull_at, id, head, base, created from pull_versions where %s order by id asc `, filter.Condition()) rows, err := q.QueryContext(ctx, query, filter.Arg()...) if err != nil { return nil, fmt.Errorf("failed to query: %w", err) } defer rows.Close() versionsMap := make(map[syntax.ATURI][]models.PullVersion) for rows.Next() { var version models.PullVersion var pullAt syntax.ATURI var createdAt string err := rows.Scan( &pullAt, &version.ID, &version.Head, &version.Base, &createdAt, ) if err != nil { return nil, fmt.Errorf("scanning row: %w", err) } createdTime, err := time.Parse(time.RFC3339, createdAt) if err != nil { return nil, fmt.Errorf("parsing created: %w", err) } version.Created = createdTime versionsMap[pullAt] = append(versionsMap[pullAt], version) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("scanning rows: %w", err) } comments, err := GetComments(q, orm.FilterIn("subject_uri", pullAts)) if err != nil { return nil, fmt.Errorf("failed to get pull comments: %w", err) } for _, comment := range comments { if comment.PullRoundIdx == nil { continue } versionIdx := *comment.PullRoundIdx if versions, ok := versionsMap[syntax.ATURI(comment.Subject.Uri)]; ok { if versionIdx >= len(versions) { continue } versions[versionIdx].Comments = append(versions[versionIdx].Comments, comment) } } // TODO: reverse-map version.Comments return versionsMap, nil } // timeframe here is directly passed into the sql query filter, and any // timeframe in the past should be negative; e.g.: "-3 months" func GetPullsByOwnerDid(e Execer, did syntax.DID, timeframe string) ([]models.Pull, error) { var pulls []models.Pull rows, err := e.Query(` select p.owner_did, p.repo_did, p.pull_id, p.created, p.title, p.state, r.did, r.name, r.knot, r.rkey, r.created from pulls p join repos r on p.repo_did = r.repo_did where p.owner_did = ? and p.created >= date ('now', ?) order by p.created desc`, did, timeframe) if err != nil { return nil, err } defer rows.Close() for rows.Next() { var pull models.Pull var repo models.Repo var pullCreatedAt, repoCreatedAt string err := rows.Scan( &pull.OwnerDid, &pull.RepoDid, &pull.PullId, &pullCreatedAt, &pull.Title, &pull.State, &repo.Did, &repo.Name, &repo.Knot, &repo.Rkey, &repoCreatedAt, ) if err != nil { return nil, err } pullCreatedTime, err := time.Parse(time.RFC3339, pullCreatedAt) if err != nil { return nil, err } pull.Created = pullCreatedTime repoCreatedTime, err := time.Parse(time.RFC3339, repoCreatedAt) if err != nil { return nil, err } repo.Created = repoCreatedTime pull.Repo = &repo pulls = append(pulls, pull) } if err := rows.Err(); err != nil { return nil, err } return pulls, nil } // use with transaction func setPullsState(e Execer, pullState models.PullState, filters ...orm.Filter) error { var conditions []string var args []any args = append(args, pullState) for _, filter := range filters { conditions = append(conditions, filter.Condition()) args = append(args, filter.Arg()...) } args = append(args, models.PullAbandoned) // only update state of non-deleted pulls args = append(args, models.PullMerged) // only update state of non-merged pulls whereClause := "" if conditions != nil { whereClause = " where " + strings.Join(conditions, " and ") } query := fmt.Sprintf("update pulls set state = ? %s and state <> ? and state <> ?", whereClause) _, err := e.Exec(query, args...) return err } func ClosePulls(e Execer, filters ...orm.Filter) error { return setPullsState(e, models.PullClosed, filters...) } func ReopenPulls(e Execer, filters ...orm.Filter) error { return setPullsState(e, models.PullOpen, filters...) } func MergePulls(e Execer, filters ...orm.Filter) error { return setPullsState(e, models.PullMerged, filters...) } func AbandonPulls(e Execer, filters ...orm.Filter) error { return setPullsState(e, models.PullAbandoned, filters...) } func GetPullCount(e Execer, repoDid string) (models.PullCount, error) { row := e.QueryRow(` select count(case when state = ? then 1 end) as open_count, count(case when state = ? then 1 end) as merged_count, count(case when state = ? then 1 end) as closed_count, count(case when state = ? then 1 end) as deleted_count from pulls where repo_did = ?`, models.PullOpen, models.PullMerged, models.PullClosed, models.PullAbandoned, repoDid, ) var count models.PullCount if err := row.Scan(&count.Open, &count.Merged, &count.Closed, &count.Deleted); err != nil { return models.PullCount{Open: 0, Merged: 0, Closed: 0, Deleted: 0}, err } return count, nil }