391 lines
12 KiB
Go
391 lines
12 KiB
Go
package app
|
|
|
|
import (
|
|
"database/sql"
|
|
"fmt"
|
|
"net/http"
|
|
"os"
|
|
"os/exec"
|
|
"path/filepath"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
// ---------------- Pull request API ----------------
|
|
|
|
func (s *Server) handlePRCreate(w http.ResponseWriter, r *http.Request, targetOwner, targetName string) {
|
|
user, ok := s.requireBearerUser(w, r)
|
|
if !ok {
|
|
return
|
|
}
|
|
target, err := s.loadRepo(targetOwner, targetName)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "target repository not found")
|
|
return
|
|
}
|
|
var in struct {
|
|
SourceOwner string `json:"source_owner"`
|
|
SourceRepo string `json:"source_repo"`
|
|
SourceBranch string `json:"source_branch"`
|
|
TargetBranch string `json:"target_branch"`
|
|
Title string `json:"title"`
|
|
Description string `json:"description"`
|
|
}
|
|
if !decodeJSON(w, r, &in) {
|
|
return
|
|
}
|
|
in.SourceOwner = strings.ToLower(strings.TrimSpace(in.SourceOwner))
|
|
in.SourceRepo = strings.ToLower(strings.TrimSpace(in.SourceRepo))
|
|
in.SourceBranch = strings.TrimSpace(in.SourceBranch)
|
|
in.TargetBranch = strings.TrimSpace(in.TargetBranch)
|
|
in.Title = strings.TrimSpace(in.Title)
|
|
if in.SourceOwner == "" {
|
|
in.SourceOwner = target.Owner
|
|
}
|
|
if in.SourceRepo == "" {
|
|
in.SourceRepo = target.Name
|
|
}
|
|
if in.Title == "" || !branchRE.MatchString(in.SourceBranch) || !branchRE.MatchString(in.TargetBranch) {
|
|
writeError(w, http.StatusBadRequest, "title and valid source/target branches are required")
|
|
return
|
|
}
|
|
source, err := s.loadRepo(in.SourceOwner, in.SourceRepo)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "source repository not found")
|
|
return
|
|
}
|
|
if source.ID == target.ID {
|
|
if target.OwnerUserID != user.ID {
|
|
writeError(w, http.StatusForbidden, "same-repository PRs require repository ownership")
|
|
return
|
|
}
|
|
} else {
|
|
if source.OwnerUserID != user.ID {
|
|
writeError(w, http.StatusForbidden, "source repository must be owned by you")
|
|
return
|
|
}
|
|
if target.Visibility != "public" && target.OwnerUserID != user.ID {
|
|
writeError(w, http.StatusForbidden, "target repository is private")
|
|
return
|
|
}
|
|
}
|
|
if !gitBranchExists(s.repoPath(source.Owner, source.Name), in.SourceBranch) {
|
|
writeError(w, http.StatusBadRequest, "source branch does not exist")
|
|
return
|
|
}
|
|
if !gitBranchExists(s.repoPath(target.Owner, target.Name), in.TargetBranch) {
|
|
writeError(w, http.StatusBadRequest, "target branch does not exist")
|
|
return
|
|
}
|
|
var number int
|
|
_ = s.db.QueryRow(`SELECT COALESCE(MAX(number), 0) + 1 FROM pull_requests WHERE target_repository_id = ?`, target.ID).Scan(&number)
|
|
res, err := s.db.Exec(`INSERT INTO pull_requests (target_repository_id, number, author_user_id, source_repository_id, source_branch, target_branch, title, description)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, target.ID, number, user.ID, source.ID, in.SourceBranch, in.TargetBranch, in.Title, in.Description)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
id, _ := res.LastInsertId()
|
|
pr, _ := s.loadPR(target.ID, number)
|
|
pr.ID = id
|
|
writeJSON(w, http.StatusCreated, pr)
|
|
}
|
|
|
|
func (s *Server) handlePRList(w http.ResponseWriter, r *http.Request, owner, name string) {
|
|
repo, err := s.loadRepo(owner, name)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "repository not found")
|
|
return
|
|
}
|
|
user, authed := s.optionalBearerUser(r)
|
|
if !s.canReadRepo(repo, user, authed) {
|
|
writeError(w, http.StatusNotFound, "repository not found")
|
|
return
|
|
}
|
|
rows, err := s.db.Query(prSelectSQL()+` WHERE pr.target_repository_id = ? ORDER BY pr.number DESC`, repo.ID)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
defer rows.Close()
|
|
prs, err := scanPRs(rows)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
writeJSON(w, http.StatusOK, prs)
|
|
}
|
|
|
|
func (s *Server) handlePRView(w http.ResponseWriter, r *http.Request, owner, name string, number int) {
|
|
repo, err := s.loadRepo(owner, name)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "repository not found")
|
|
return
|
|
}
|
|
user, authed := s.optionalBearerUser(r)
|
|
if !s.canReadRepo(repo, user, authed) {
|
|
writeError(w, http.StatusNotFound, "repository not found")
|
|
return
|
|
}
|
|
pr, err := s.loadPR(repo.ID, number)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "pull request not found")
|
|
return
|
|
}
|
|
writeJSON(w, http.StatusOK, pr)
|
|
}
|
|
|
|
func (s *Server) handlePRClose(w http.ResponseWriter, r *http.Request, owner, name string, number int) {
|
|
user, ok := s.requireBearerUser(w, r)
|
|
if !ok {
|
|
return
|
|
}
|
|
repo, err := s.loadRepo(owner, name)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "repository not found")
|
|
return
|
|
}
|
|
if repo.OwnerUserID != user.ID {
|
|
writeError(w, http.StatusForbidden, "only target owner can close pull requests")
|
|
return
|
|
}
|
|
res, err := s.db.Exec(`UPDATE pull_requests SET status = 'closed', closed_at = UTC_TIMESTAMP() WHERE target_repository_id = ? AND number = ? AND status = 'open'`, repo.ID, number)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
affected, _ := res.RowsAffected()
|
|
if affected == 0 {
|
|
writeError(w, http.StatusConflict, "pull request is not open")
|
|
return
|
|
}
|
|
pr, _ := s.loadPR(repo.ID, number)
|
|
writeJSON(w, http.StatusOK, pr)
|
|
}
|
|
|
|
func (s *Server) handlePRMerge(w http.ResponseWriter, r *http.Request, owner, name string, number int) {
|
|
user, ok := s.requireBearerUser(w, r)
|
|
if !ok {
|
|
return
|
|
}
|
|
target, err := s.loadRepo(owner, name)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "repository not found")
|
|
return
|
|
}
|
|
if target.OwnerUserID != user.ID {
|
|
writeError(w, http.StatusForbidden, "only target owner can merge pull requests")
|
|
return
|
|
}
|
|
pr, err := s.loadPR(target.ID, number)
|
|
if err != nil || pr.Status != "open" {
|
|
writeError(w, http.StatusConflict, "pull request is not open")
|
|
return
|
|
}
|
|
if err := s.mergePR(pr); err != nil {
|
|
writeError(w, http.StatusConflict, err.Error())
|
|
return
|
|
}
|
|
_, err = s.db.Exec(`UPDATE pull_requests SET status = 'merged', merged_at = UTC_TIMESTAMP() WHERE target_repository_id = ? AND number = ?`, target.ID, number)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
pr, _ = s.loadPR(target.ID, number)
|
|
writeJSON(w, http.StatusOK, pr)
|
|
}
|
|
|
|
func prSelectSQL() string {
|
|
return `SELECT pr.id, pr.number, pr.target_repository_id, pr.source_repository_id, pr.author_user_id,
|
|
au.username, su.username, sr.name, pr.source_branch, tu.username, tr.name, pr.target_branch,
|
|
pr.title, pr.description, pr.status, pr.created_at, pr.updated_at, pr.closed_at, pr.merged_at
|
|
FROM pull_requests pr
|
|
JOIN users au ON au.id = pr.author_user_id
|
|
JOIN repositories sr ON sr.id = pr.source_repository_id
|
|
JOIN users su ON su.id = sr.owner_user_id
|
|
JOIN repositories tr ON tr.id = pr.target_repository_id
|
|
JOIN users tu ON tu.id = tr.owner_user_id`
|
|
}
|
|
|
|
func (s *Server) loadPR(targetRepoID int64, number int) (PullRequest, error) {
|
|
row := s.db.QueryRow(prSelectSQL()+` WHERE pr.target_repository_id = ? AND pr.number = ?`, targetRepoID, number)
|
|
prs, err := scanOnePR(row)
|
|
return prs, err
|
|
}
|
|
|
|
type scanner interface{ Scan(dest ...any) error }
|
|
|
|
func scanOnePR(row scanner) (PullRequest, error) {
|
|
var pr PullRequest
|
|
var closedAt, mergedAt sql.NullTime
|
|
err := row.Scan(&pr.ID, &pr.Number, &pr.TargetRepositoryID, &pr.SourceRepositoryID, &pr.AuthorUserID,
|
|
&pr.Author, &pr.SourceOwner, &pr.SourceRepo, &pr.SourceBranch, &pr.TargetOwner, &pr.TargetRepo, &pr.TargetBranch,
|
|
&pr.Title, &pr.Description, &pr.Status, &pr.CreatedAt, &pr.UpdatedAt, &closedAt, &mergedAt)
|
|
if closedAt.Valid {
|
|
pr.ClosedAt = &closedAt.Time
|
|
}
|
|
if mergedAt.Valid {
|
|
pr.MergedAt = &mergedAt.Time
|
|
}
|
|
return pr, err
|
|
}
|
|
|
|
func scanPRs(rows *sql.Rows) ([]PullRequest, error) {
|
|
var prs []PullRequest
|
|
for rows.Next() {
|
|
pr, err := scanOnePR(rows)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
prs = append(prs, pr)
|
|
}
|
|
return prs, rows.Err()
|
|
}
|
|
|
|
func (s *Server) mergePR(pr PullRequest) error {
|
|
targetPath := s.repoPath(pr.TargetOwner, pr.TargetRepo)
|
|
sourcePath := s.repoPath(pr.SourceOwner, pr.SourceRepo)
|
|
tmp, err := os.MkdirTemp("", "gitocean-merge-*")
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer os.RemoveAll(tmp)
|
|
work := filepath.Join(tmp, "work")
|
|
if err := runGit("", "clone", targetPath, work); err != nil {
|
|
return err
|
|
}
|
|
if err := runGit(work, "config", "user.name", "gitocean"); err != nil {
|
|
return err
|
|
}
|
|
if err := runGit(work, "config", "user.email", "gitocean@localhost"); err != nil {
|
|
return err
|
|
}
|
|
if err := runGit(work, "checkout", pr.TargetBranch); err != nil {
|
|
return err
|
|
}
|
|
if err := runGit(work, "remote", "add", "source", sourcePath); err != nil {
|
|
return err
|
|
}
|
|
if err := runGit(work, "fetch", "source", pr.SourceBranch); err != nil {
|
|
return err
|
|
}
|
|
msg := fmt.Sprintf("Merge pull request #%d from %s/%s:%s", pr.Number, pr.SourceOwner, pr.SourceRepo, pr.SourceBranch)
|
|
if err := runGit(work, "merge", "--no-ff", "FETCH_HEAD", "-m", msg); err != nil {
|
|
return fmt.Errorf("merge failed, likely due to conflicts: %w", err)
|
|
}
|
|
if err := runGit(work, "push", "origin", "HEAD:"+pr.TargetBranch); err != nil {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (s *Server) handlePRDiff(w http.ResponseWriter, r *http.Request, owner, name string, number int) {
|
|
repo, ok := s.requireReadableRepo(w, r, owner, name)
|
|
if !ok {
|
|
return
|
|
}
|
|
pr, err := s.loadPR(repo.ID, number)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "pull request not found")
|
|
return
|
|
}
|
|
diff, err := s.prDiff(pr)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
|
_, _ = w.Write([]byte(diff))
|
|
}
|
|
|
|
func (s *Server) prDiff(pr PullRequest) (string, error) {
|
|
tmp, err := os.MkdirTemp("", "gitocean-diff-*")
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
defer os.RemoveAll(tmp)
|
|
work := filepath.Join(tmp, "work")
|
|
if err := runGit("", "clone", "--no-checkout", s.repoPath(pr.TargetOwner, pr.TargetRepo), work); err != nil {
|
|
return "", err
|
|
}
|
|
if err := runGit(work, "remote", "add", "source", s.repoPath(pr.SourceOwner, pr.SourceRepo)); err != nil {
|
|
return "", err
|
|
}
|
|
if err := runGit(work, "fetch", "origin", pr.TargetBranch); err != nil {
|
|
return "", err
|
|
}
|
|
if err := runGit(work, "fetch", "source", pr.SourceBranch); err != nil {
|
|
return "", err
|
|
}
|
|
cmd := exec.Command("git", "diff", "--patch", "origin/"+pr.TargetBranch+"...FETCH_HEAD")
|
|
cmd.Dir = work
|
|
out, err := cmd.CombinedOutput()
|
|
if err != nil {
|
|
return "", fmt.Errorf("git diff failed: %s", strings.TrimSpace(string(out)))
|
|
}
|
|
return string(out), nil
|
|
}
|
|
|
|
func (s *Server) handlePRCommentsList(w http.ResponseWriter, r *http.Request, owner, name string, number int) {
|
|
repo, ok := s.requireReadableRepo(w, r, owner, name)
|
|
if !ok {
|
|
return
|
|
}
|
|
pr, err := s.loadPR(repo.ID, number)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "pull request not found")
|
|
return
|
|
}
|
|
rows, err := s.db.Query(`SELECT c.id, u.username, c.body, c.created_at, c.updated_at FROM pull_request_comments c JOIN users u ON u.id = c.author_user_id WHERE c.pull_request_id = ? ORDER BY c.created_at`, pr.ID)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
defer rows.Close()
|
|
var out []PRComment
|
|
for rows.Next() {
|
|
var c PRComment
|
|
if err := rows.Scan(&c.ID, &c.Author, &c.Body, &c.CreatedAt, &c.UpdatedAt); err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
out = append(out, c)
|
|
}
|
|
writeJSON(w, http.StatusOK, out)
|
|
}
|
|
|
|
func (s *Server) handlePRCommentCreate(w http.ResponseWriter, r *http.Request, owner, name string, number int) {
|
|
user, ok := s.requireBearerUser(w, r)
|
|
if !ok {
|
|
return
|
|
}
|
|
repo, err := s.loadRepo(owner, name)
|
|
if err != nil || !s.canReadRepo(repo, user, true) {
|
|
writeError(w, http.StatusNotFound, "repository not found")
|
|
return
|
|
}
|
|
pr, err := s.loadPR(repo.ID, number)
|
|
if err != nil {
|
|
writeError(w, http.StatusNotFound, "pull request not found")
|
|
return
|
|
}
|
|
var in struct {
|
|
Body string `json:"body"`
|
|
}
|
|
if !decodeJSON(w, r, &in) {
|
|
return
|
|
}
|
|
body := strings.TrimSpace(in.Body)
|
|
if body == "" {
|
|
writeError(w, http.StatusBadRequest, "comment body is required")
|
|
return
|
|
}
|
|
res, err := s.db.Exec(`INSERT INTO pull_request_comments (pull_request_id, author_user_id, body) VALUES (?, ?, ?)`, pr.ID, user.ID, body)
|
|
if err != nil {
|
|
writeError(w, http.StatusInternalServerError, err.Error())
|
|
return
|
|
}
|
|
id, _ := res.LastInsertId()
|
|
writeJSON(w, http.StatusCreated, PRComment{ID: id, Author: user.Username, Body: body, CreatedAt: time.Now().UTC(), UpdatedAt: time.Now().UTC()})
|
|
}
|