diff --git a/docs/reference.md b/docs/reference.md index 8e08ed3..20df229 100644 --- a/docs/reference.md +++ b/docs/reference.md @@ -163,8 +163,10 @@ The webhook endpoint accepts only signed GitHub `POST` deliveries up to 10 MiB and requires exactly one configured project for the installation. Supported deliveries return HTTP 202 with `{"accepted":true,"queued":true}`. Unsupported event names, conflicted pull requests, and pull requests from fork repositories -return HTTP 202 with `{"accepted":true,"queued":false}`. Pull request deliveries -whose merge ref is still being prepared are queued and resolved asynchronously. +return HTTP 202 with `{"accepted":true,"queued":false}`. Open pull request +workflows use GitHub's test merge revision for the webhook head. Deliveries wait +up to two minutes for that revision; unavailable or superseded revisions +produce a `Failed` delivery. Queued deliveries are processed asynchronously. Invalid or unsupported workflow definitions fail the whole delivery before any `WorkflowRun` resources are diff --git a/internal/github/client.go b/internal/github/client.go index 7aa9866..2e88bd4 100644 --- a/internal/github/client.go +++ b/internal/github/client.go @@ -45,6 +45,13 @@ type Content struct { Type string `json:"type"` } +type repositoryCommit struct { + SHA string `json:"sha"` + Parents []struct { + SHA string `json:"sha"` + } `json:"parents"` +} + // APIError describes a non-success response from the GitHub API. type APIError struct { StatusCode int @@ -169,23 +176,45 @@ func (c *InstallationClient) GetFile(ctx context.Context, owner, repository, fil // ResolveRevision resolves a branch, tag, or commit expression to a full commit // SHA. func (c *InstallationClient) ResolveRevision(ctx context.Context, owner, repository, revision string) (string, error) { + commit, err := c.resolveRevision(ctx, owner, repository, revision) + if err != nil { + return "", err + } + return commit.SHA, nil +} + +func (c *InstallationClient) resolveRevision(ctx context.Context, owner, repository, revision string) (repositoryCommit, error) { identity := fmt.Sprintf("resolve repository revision %q from %s/%s", revision, owner, repository) requestPath := "repos/" + owner + "/" + repository + "/commits" - commits := []struct { - SHA string `json:"sha"` - }{} + commits := []repositoryCommit{} if err := c.client.doJSONWithQuery(ctx, http.MethodGet, requestPath, url.Values{"sha": []string{revision}, "per_page": []string{"1"}}, c.token, &commits); err != nil { - return "", fmt.Errorf("%s: %w", identity, err) + return repositoryCommit{}, fmt.Errorf("%s: %w", identity, err) } if len(commits) == 0 { - return "", fmt.Errorf("%s: GitHub returned no commits", identity) + return repositoryCommit{}, fmt.Errorf("%s: GitHub returned no commits", identity) } - sha := commits[0].SHA + commit := commits[0] + sha := commit.SHA decoded, err := hex.DecodeString(sha) if err != nil || len(decoded) != gitSHA1Bytes || sha != strings.ToLower(sha) { - return "", fmt.Errorf("%s: GitHub returned invalid commit SHA %q", identity, sha) + return repositoryCommit{}, fmt.Errorf("%s: GitHub returned invalid commit SHA %q", identity, sha) + } + return commit, nil +} + +// ResolvePullRequestRevision resolves a pull request merge ref and reports +// whether its merge commit includes the expected head commit. +func (c *InstallationClient) ResolvePullRequestRevision(ctx context.Context, owner, repository, revision, headSHA string) (string, bool, error) { + commit, err := c.resolveRevision(ctx, owner, repository, revision) + if err != nil { + return "", false, err + } + for _, parent := range commit.Parents { + if parent.SHA == headSHA { + return commit.SHA, true, nil + } } - return sha, nil + return commit.SHA, false, nil } func (c *Client) doJSON(ctx context.Context, method, requestPath, token string, destination any) error { diff --git a/internal/github/client_test.go b/internal/github/client_test.go index 52b1c62..51bbec4 100644 --- a/internal/github/client_test.go +++ b/internal/github/client_test.go @@ -126,6 +126,43 @@ func TestResolveRevisionErrorIncludesRepositoryIdentity(t *testing.T) { } } +func TestResolvePullRequestRevisionRequiresExpectedHeadParent(t *testing.T) { + mergeSHA := strings.Repeat("c", 40) + expectedHeadSHA := strings.Repeat("b", 40) + for _, tt := range []struct { + name string + parentSHA string + wantReady bool + }{ + {name: "current", parentSHA: expectedHeadSHA, wantReady: true}, + {name: "stale", parentSHA: strings.Repeat("a", 40)}, + } { + t.Run(tt.name, func(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/repos/acme/example/commits" || request.URL.Query().Get("sha") != "refs/pull/9/merge" { + http.NotFound(writer, request) + return + } + fmt.Fprintf(writer, `[{"sha":%q,"parents":[{"sha":%q}]}]`, mergeSHA, tt.parentSHA) + })) + defer server.Close() + client, err := NewClient(server.URL, server.Client()) + if err != nil { + t.Fatal(err) + } + installation := &InstallationClient{client: client, token: "token"} + + resolved, ready, err := installation.ResolvePullRequestRevision(context.Background(), "acme", "example", "refs/pull/9/merge", expectedHeadSHA) + if err != nil { + t.Fatal(err) + } + if resolved != mergeSHA || ready != tt.wantReady { + t.Errorf("resolved = %q, ready = %v", resolved, ready) + } + }) + } +} + func TestAPIErrorPreservesStatusAndMessage(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) { writer.WriteHeader(http.StatusNotFound) diff --git a/internal/webhook/delivery.go b/internal/webhook/delivery.go index de49cfc..85ba5b2 100644 --- a/internal/webhook/delivery.go +++ b/internal/webhook/delivery.go @@ -33,6 +33,7 @@ var digestEncoding = base32.StdEncoding.WithPadding(base32.NoPadding) const ( deliveryLabel = "actions.kelos.dev/webhook-delivery" deliveryDataKey = "delivery.json" + deliveryRevisionKey = "resolvedRevision" deliveryStateKey = "state" deliveryMessageKey = "message" deliveryRunCountKey = "workflowRuns" @@ -42,6 +43,7 @@ const ( maxDeliveryBytes = 900_000 maxWorkflowFiles = 100 maxWorkflowJobs = 1000 + mergeRefWaitTimeout = 2 * time.Minute deliveryRetention = 24 * time.Hour ) @@ -164,6 +166,34 @@ func missingWorkflowDirectory(err error) bool { return errors.As(err, &apiError) && apiError.StatusCode == http.StatusNotFound && apiError.Message == "Not Found" } +func missingPullRequestMergeRef(err error) bool { + apiError := &githubclient.APIError{} + return errors.As(err, &apiError) && apiError.StatusCode == http.StatusNotFound +} + +func mergeRefRetryInterval(age time.Duration) time.Duration { + switch { + case age >= 30*time.Second: + return 15 * time.Second + case age >= 10*time.Second: + return 5 * time.Second + default: + return 2 * time.Second + } +} + +func resolveDeliveryRevision(ctx context.Context, installation *githubclient.InstallationClient, owner, repository string, event normalizedEvent) (string, bool, error) { + if event.HeadSHA == "" { + revision, err := installation.ResolveRevision(ctx, owner, repository, event.ResolveRef) + return revision, err == nil, err + } + revision, ready, err := installation.ResolvePullRequestRevision(ctx, owner, repository, event.ResolveRef, event.HeadSHA) + if err != nil && missingPullRequestMergeRef(err) { + return "", false, nil + } + return revision, ready, err +} + func (r *DeliveryReconciler) Reconcile(ctx context.Context, request ctrl.Request) (ctrl.Result, error) { object := &corev1.ConfigMap{} if err := r.Get(ctx, request.NamespacedName, object); err != nil { @@ -179,6 +209,12 @@ func (r *DeliveryReconciler) Reconcile(ctx context.Context, request ctrl.Request if err := json.Unmarshal([]byte(object.Data[deliveryDataKey]), &delivery); err != nil { return ctrl.Result{}, r.finish(ctx, object, deliveryStateFailed, 0, fmt.Sprintf("decode delivery: %v", err)) } + if revision := object.Data[deliveryRevisionKey]; revision != "" { + if !validGitSHA(revision) { + return ctrl.Result{}, r.finish(ctx, object, deliveryStateFailed, 0, "delivery contains an invalid resolved revision") + } + delivery.Event.SHA = revision + } reader := r.APIReader project := &actionsv1alpha1.Project{} if err := reader.Get(ctx, client.ObjectKey{Namespace: object.Namespace, Name: delivery.ProjectName}, project); err != nil { @@ -199,11 +235,25 @@ func (r *DeliveryReconciler) Reconcile(ctx context.Context, request ctrl.Request if err != nil { return ctrl.Result{}, err } - if delivery.Event.ResolveRef != "" { - delivery.Event.SHA, err = installation.ResolveRevision(ctx, delivery.Payload.Repository.Owner.Login, delivery.Payload.Repository.Name, delivery.Event.ResolveRef) + if delivery.Event.ResolveRef != "" && delivery.Event.SHA == "" { + revision, ready, err := resolveDeliveryRevision(ctx, installation, delivery.Payload.Repository.Owner.Login, delivery.Payload.Repository.Name, delivery.Event) if err != nil { return ctrl.Result{}, err } + if !ready { + age := r.deliveryAge(object) + if age >= mergeRefWaitTimeout { + message := fmt.Sprintf("GitHub pull request merge revision did not become ready for head %s within %s", delivery.Event.HeadSHA, mergeRefWaitTimeout) + return ctrl.Result{}, r.finish(ctx, object, deliveryStateFailed, 0, message) + } + retryAfter := mergeRefRetryInterval(age) + r.Logger.Debug("waiting for GitHub pull request merge revision", "delivery_id", delivery.DeliveryID, "head_sha", delivery.Event.HeadSHA, "retry_after", retryAfter) + return ctrl.Result{RequeueAfter: retryAfter}, nil + } + delivery.Event.SHA = revision + if err := r.persistResolvedRevision(ctx, object, delivery.Event.SHA); err != nil { + return ctrl.Result{}, err + } } contents, err := installation.ListDirectory(ctx, delivery.Payload.Repository.Owner.Login, delivery.Payload.Repository.Name, project.Spec.WorkflowDirectory, delivery.Event.SHA) if err != nil { @@ -322,6 +372,15 @@ func (r *DeliveryReconciler) finish(ctx context.Context, object *corev1.ConfigMa return r.Patch(ctx, object, client.MergeFrom(before)) } +func (r *DeliveryReconciler) persistResolvedRevision(ctx context.Context, object *corev1.ConfigMap, revision string) error { + before := object.DeepCopy() + if object.Data == nil { + object.Data = map[string]string{} + } + object.Data[deliveryRevisionKey] = revision + return r.Patch(ctx, object, client.MergeFrom(before)) +} + func (r *DeliveryReconciler) retain(ctx context.Context, object *corev1.ConfigMap) (ctrl.Result, error) { finishedAt, err := time.Parse(time.RFC3339, object.Data[deliveryFinishedKey]) if err != nil { @@ -344,6 +403,17 @@ func (r *DeliveryReconciler) now() time.Time { return time.Now() } +func (r *DeliveryReconciler) deliveryAge(object *corev1.ConfigMap) time.Duration { + if object.CreationTimestamp.IsZero() { + return 0 + } + age := r.now().Sub(object.CreationTimestamp.Time) + if age < 0 { + return 0 + } + return age +} + func (r *DeliveryReconciler) SetupWithManager(manager ctrl.Manager) error { return ctrl.NewControllerManagedBy(manager). For(&corev1.ConfigMap{}, builder.WithPredicates(predicate.NewPredicateFuncs(isWebhookDelivery))). diff --git a/internal/webhook/delivery_test.go b/internal/webhook/delivery_test.go index e728393..3fa49df 100644 --- a/internal/webhook/delivery_test.go +++ b/internal/webhook/delivery_test.go @@ -2,10 +2,19 @@ package webhook import ( "context" + "crypto/rand" + "crypto/rsa" "crypto/sha256" + "crypto/x509" + "encoding/base64" "encoding/json" + "encoding/pem" "errors" "fmt" + "io" + "log/slog" + "net/http" + "net/http/httptest" "strings" "testing" "time" @@ -39,6 +48,15 @@ func TestWebhookReplayIDUsesSignedBody(t *testing.T) { if err := handler.enqueueDelivery(context.Background(), project, event, normalized, "original-delivery", body); err != nil { t.Fatal(err) } + stored := &corev1.ConfigMap{} + key := client.ObjectKey{Namespace: project.Namespace, Name: webhookDeliveryName(body)} + if err := clusterClient.Get(context.Background(), key, stored); err != nil { + t.Fatal(err) + } + stored.Data[deliveryRevisionKey] = strings.Repeat("b", 40) + if err := clusterClient.Update(context.Background(), stored); err != nil { + t.Fatal(err) + } if err := handler.enqueueDelivery(context.Background(), project, event, normalized, "replay-delivery", body); err != nil { t.Fatalf("signed-body replay was not idempotent: %v", err) } @@ -143,6 +161,200 @@ func TestCreateWorkflowRunReplayUsesLiveReader(t *testing.T) { } } +func TestDeliveryPinsCurrentPullRequestMergeRevision(t *testing.T) { + now := time.Date(2026, 8, 9, 23, 0, 0, 0, time.UTC) + headSHA := strings.Repeat("b", 40) + mergeSHA := strings.Repeat("c", 40) + movedMergeSHA := strings.Repeat("d", 40) + parentSHA := strings.Repeat("a", 40) + resolvedSHA := mergeSHA + resolveCalls := 0 + failDiscovery := false + workflowData := []byte("name: CI\non: pull_request\njobs:\n build:\n runs-on: ubuntu-latest\n steps:\n - run: make test\n") + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/app/installations/2/access_tokens": + fmt.Fprint(writer, `{"token":"installation-token"}`) + case "/repos/acme/example/commits": + resolveCalls++ + fmt.Fprintf(writer, `[{"sha":%q,"parents":[{"sha":%q}]}]`, resolvedSHA, parentSHA) + case "/repos/acme/example/contents/.open-actions/workflows": + if failDiscovery { + failDiscovery = false + http.Error(writer, "temporarily unavailable", http.StatusServiceUnavailable) + return + } + fmt.Fprint(writer, `[{"name":"ci.yaml","path":".open-actions/workflows/ci.yaml","type":"file"}]`) + case "/repos/acme/example/contents/.open-actions/workflows/ci.yaml": + fmt.Fprintf(writer, `{"encoding":"base64","content":%q}`, base64.StdEncoding.EncodeToString(workflowData)) + default: + http.NotFound(writer, request) + } + })) + defer server.Close() + clusterClient, reconciler, handler, project := newPullRequestDeliveryTest(t, server, now) + deliveryKey := enqueuePullRequestDelivery(t, handler, clusterClient, project, now, headSHA, []byte(`{"delivery":"current-merge-ref"}`)) + + result, err := reconciler.Reconcile(context.Background(), ctrl.Request{NamespacedName: deliveryKey}) + if err != nil { + t.Fatal(err) + } + if result.RequeueAfter != mergeRefRetryInterval(0) { + t.Fatalf("requeue after = %v, want %v", result.RequeueAfter, mergeRefRetryInterval(0)) + } + + parentSHA = headSHA + failDiscovery = true + if _, err := reconciler.Reconcile(context.Background(), ctrl.Request{NamespacedName: deliveryKey}); err == nil { + t.Fatal("reconcile succeeded during a transient discovery failure") + } + stored := &corev1.ConfigMap{} + if err := clusterClient.Get(context.Background(), deliveryKey, stored); err != nil { + t.Fatal(err) + } + if stored.Data[deliveryRevisionKey] != mergeSHA { + t.Fatalf("resolved revision = %q, want %q", stored.Data[deliveryRevisionKey], mergeSHA) + } + + resolvedSHA = movedMergeSHA + result, err = reconciler.Reconcile(context.Background(), ctrl.Request{NamespacedName: deliveryKey}) + if err != nil { + t.Fatal(err) + } + if result.RequeueAfter != 0 { + t.Fatalf("requeue after = %v, want 0", result.RequeueAfter) + } + if resolveCalls != 2 { + t.Fatalf("merge ref resolutions = %d, want 2", resolveCalls) + } + if err := clusterClient.Get(context.Background(), deliveryKey, stored); err != nil { + t.Fatal(err) + } + if stored.Data[deliveryStateKey] != deliveryStateCompleted { + t.Fatalf("delivery state = %q, want %q", stored.Data[deliveryStateKey], deliveryStateCompleted) + } + runs := &actionsv1alpha1.WorkflowRunList{} + if err := clusterClient.List(context.Background(), runs); err != nil { + t.Fatal(err) + } + if len(runs.Items) != 1 { + t.Fatalf("WorkflowRuns = %d, want 1", len(runs.Items)) + } + if got := runs.Items[0].Spec.Source.GitHub.Revision.SHA; got != mergeSHA { + t.Fatalf("WorkflowRun revision = %q, want pinned revision %q", got, mergeSHA) + } +} + +func TestDeliveryTimesOutWhenPullRequestMergeRefIsUnavailable(t *testing.T) { + now := time.Date(2026, 8, 9, 23, 0, 0, 0, time.UTC) + headSHA := strings.Repeat("b", 40) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + switch request.URL.Path { + case "/app/installations/2/access_tokens": + fmt.Fprint(writer, `{"token":"installation-token"}`) + case "/repos/acme/example/commits": + http.NotFound(writer, request) + default: + http.NotFound(writer, request) + } + })) + defer server.Close() + clusterClient, reconciler, handler, project := newPullRequestDeliveryTest(t, server, now) + deliveryKey := enqueuePullRequestDelivery(t, handler, clusterClient, project, now, headSHA, []byte(`{"delivery":"unavailable"}`)) + + result, err := reconciler.Reconcile(context.Background(), ctrl.Request{NamespacedName: deliveryKey}) + if err != nil { + t.Fatal(err) + } + if result.RequeueAfter != mergeRefRetryInterval(0) { + t.Fatalf("requeue after = %v, want %v", result.RequeueAfter, mergeRefRetryInterval(0)) + } + stored := &corev1.ConfigMap{} + if err := clusterClient.Get(context.Background(), deliveryKey, stored); err != nil { + t.Fatal(err) + } + if stored.Data[deliveryStateKey] != "" { + t.Fatalf("delivery state = %q, want pending", stored.Data[deliveryStateKey]) + } + + reconciler.Now = func() time.Time { return now.Add(mergeRefWaitTimeout) } + result, err = reconciler.Reconcile(context.Background(), ctrl.Request{NamespacedName: deliveryKey}) + if err != nil { + t.Fatal(err) + } + if result.RequeueAfter != 0 { + t.Fatalf("requeue after = %v, want 0", result.RequeueAfter) + } + if err := clusterClient.Get(context.Background(), deliveryKey, stored); err != nil { + t.Fatal(err) + } + if stored.Data[deliveryStateKey] != deliveryStateFailed || !strings.Contains(stored.Data[deliveryMessageKey], headSHA) { + t.Fatalf("terminal delivery data = %#v", stored.Data) + } +} + +func newPullRequestDeliveryTest(t *testing.T, server *httptest.Server, now time.Time) (client.Client, *DeliveryReconciler, *GitHubHandler, *actionsv1alpha1.Project) { + t.Helper() + githubAPI, err := githubclient.NewClient(server.URL, server.Client()) + if err != nil { + t.Fatal(err) + } + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatal(err) + } + privateKeyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(privateKey)}) + project := &actionsv1alpha1.Project{ + ObjectMeta: metav1.ObjectMeta{Name: "default", Namespace: "default", UID: "project-uid"}, + Spec: actionsv1alpha1.ProjectSpec{ + WorkflowDirectory: ".open-actions/workflows", + Source: actionsv1alpha1.ProjectSource{ + Type: actionsv1alpha1.SourceTypeGitHub, + GitHub: &actionsv1alpha1.GitHubAppConfiguration{ + AppID: 1, InstallationID: 2, + PrivateKeySecretRef: corev1.SecretKeySelector{LocalObjectReference: corev1.LocalObjectReference{Name: "github"}, Key: "private-key"}, + WebhookSecretRef: corev1.SecretKeySelector{LocalObjectReference: corev1.LocalObjectReference{Name: "github"}, Key: "webhook-secret"}, + }, + }, + }, + } + secret := &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{Name: "github", Namespace: project.Namespace}, + Data: map[string][]byte{"private-key": privateKeyPEM, "webhook-secret": []byte("secret")}, + } + scheme := deliveryTestScheme(t) + clusterClient := fake.NewClientBuilder().WithScheme(scheme).WithObjects(project, secret).Build() + logger := slog.New(slog.NewTextHandler(io.Discard, nil)) + reconciler := &DeliveryReconciler{Client: clusterClient, APIReader: clusterClient, GitHub: githubAPI, Logger: logger, Now: func() time.Time { return now }} + handler := &GitHubHandler{Client: clusterClient, APIReader: clusterClient} + return clusterClient, reconciler, handler, project +} + +func enqueuePullRequestDelivery(t *testing.T, handler *GitHubHandler, clusterClient client.Client, project *actionsv1alpha1.Project, createdAt time.Time, headSHA string, body []byte) client.ObjectKey { + t.Helper() + event := &payload{} + event.Repository.ID = 1 + event.Repository.Name = "example" + event.Repository.Owner.Login = "acme" + normalized := normalizedEvent{ + Name: "pull_request", Action: "synchronize", Ref: "refs/pull/9/merge", + ResolveRef: "refs/pull/9/merge", HeadRef: "feature", BaseRef: "main", HeadSHA: headSHA, + } + if err := handler.enqueueDelivery(context.Background(), project, event, normalized, "delivery", body); err != nil { + t.Fatal(err) + } + key := client.ObjectKey{Namespace: project.Namespace, Name: webhookDeliveryName(body)} + object := &corev1.ConfigMap{} + if err := clusterClient.Get(context.Background(), key, object); err != nil { + t.Fatal(err) + } + object.CreationTimestamp = metav1.NewTime(createdAt) + if err := clusterClient.Update(context.Background(), object); err != nil { + t.Fatal(err) + } + return key +} + type workflowRunAlreadyExistsClient struct { client.Client } @@ -217,6 +429,21 @@ func TestDeliveryFanOutLimits(t *testing.T) { } } +func TestMergeRefRetryIntervalGrowsWithDeliveryAge(t *testing.T) { + for _, tt := range []struct { + age time.Duration + want time.Duration + }{ + {age: 0, want: 2 * time.Second}, + {age: 10 * time.Second, want: 5 * time.Second}, + {age: 30 * time.Second, want: 15 * time.Second}, + } { + if got := mergeRefRetryInterval(tt.age); got != tt.want { + t.Errorf("mergeRefRetryInterval(%v) = %v, want %v", tt.age, got, tt.want) + } + } +} + func TestTerminalDeliveryRetention(t *testing.T) { now := time.Date(2026, 8, 8, 12, 0, 0, 0, time.UTC) for _, tt := range []struct { diff --git a/internal/webhook/github.go b/internal/webhook/github.go index bfb5aab..c769fb3 100644 --- a/internal/webhook/github.go +++ b/internal/webhook/github.go @@ -57,6 +57,7 @@ type payload struct { Mergeable *bool `json:"mergeable"` Head struct { Ref string `json:"ref"` + SHA string `json:"sha"` Repository struct { ID int64 `json:"id"` } `json:"repo"` @@ -80,6 +81,7 @@ type normalizedEvent struct { HeadRef string `json:"headRef,omitempty"` BaseRef string `json:"baseRef,omitempty"` ResolveRef string `json:"resolveRef,omitempty"` + HeadSHA string `json:"headSHA,omitempty"` } func (h *GitHubHandler) ServeHTTP(writer http.ResponseWriter, request *http.Request) { @@ -201,11 +203,14 @@ func normalize(eventName string, event *payload) (normalizedEvent, bool, error) result.Ref = "refs/heads/" + pullRequest.Base.Ref } else { result.Ref = "refs/pull/" + strconv.FormatInt(pullRequest.Number, 10) + "/merge" - if pullRequest.MergeCommitSHA == "" { - if pullRequest.State != "open" { - return result, false, nil + if pullRequest.State == "open" { + if !validGitSHA(pullRequest.Head.SHA) { + return normalizedEvent{}, false, errors.New("GitHub pull request event contains an invalid head revision") } result.ResolveRef = result.Ref + result.HeadSHA = pullRequest.Head.SHA + } else if pullRequest.MergeCommitSHA == "" { + return result, false, nil } else { result.SHA = pullRequest.MergeCommitSHA } diff --git a/internal/webhook/github_test.go b/internal/webhook/github_test.go index 79a437b..59afaab 100644 --- a/internal/webhook/github_test.go +++ b/internal/webhook/github_test.go @@ -74,26 +74,29 @@ func TestProjectForInstallationUsesConfiguredOwner(t *testing.T) { func TestNormalizePullRequestRevisions(t *testing.T) { mergeable := true conflicted := false + headSHA := strings.Repeat("f", 40) for _, tt := range []struct { name string action string state string merged bool mergeable *bool - sha string + mergeSHA string wantSupport bool wantRef string wantRefName string + wantSHA string wantResolve string + wantHeadSHA string wantError bool }{ - {name: "open mergeable", action: "synchronize", state: "open", mergeable: &mergeable, sha: strings.Repeat("a", 40), wantSupport: true, wantRef: "refs/pull/42/merge", wantRefName: "42/merge"}, - {name: "open conflicted", action: "synchronize", state: "open", mergeable: &conflicted, sha: strings.Repeat("b", 40)}, - {name: "open merge result unavailable", action: "opened", state: "open", mergeable: &mergeable, wantSupport: true, wantRef: "refs/pull/42/merge", wantRefName: "42/merge", wantResolve: "refs/pull/42/merge"}, - {name: "closed unmerged", action: "closed", state: "closed", mergeable: &mergeable, sha: strings.Repeat("c", 40), wantSupport: true, wantRef: "refs/pull/42/merge", wantRefName: "42/merge"}, + {name: "open mergeable resolves test merge revision", action: "synchronize", state: "open", mergeable: &mergeable, mergeSHA: strings.Repeat("a", 40), wantSupport: true, wantRef: "refs/pull/42/merge", wantRefName: "42/merge", wantResolve: "refs/pull/42/merge", wantHeadSHA: headSHA}, + {name: "open conflicted", action: "synchronize", state: "open", mergeable: &conflicted, mergeSHA: strings.Repeat("b", 40)}, + {name: "open merge result unavailable", action: "opened", state: "open", mergeable: &mergeable, wantSupport: true, wantRef: "refs/pull/42/merge", wantRefName: "42/merge", wantResolve: "refs/pull/42/merge", wantHeadSHA: headSHA}, + {name: "closed unmerged", action: "closed", state: "closed", mergeable: &mergeable, mergeSHA: strings.Repeat("c", 40), wantSupport: true, wantRef: "refs/pull/42/merge", wantRefName: "42/merge", wantSHA: strings.Repeat("c", 40)}, {name: "closed unmerged without revision", action: "closed", state: "closed", mergeable: &mergeable}, - {name: "closed merged", action: "closed", state: "closed", merged: true, mergeable: &conflicted, sha: strings.Repeat("d", 40), wantSupport: true, wantRef: "refs/heads/main", wantRefName: "main"}, - {name: "merged payload with non-closed activity", action: "synchronize", state: "closed", merged: true, mergeable: &conflicted, sha: strings.Repeat("e", 40), wantError: true}, + {name: "closed merged", action: "closed", state: "closed", merged: true, mergeable: &conflicted, mergeSHA: strings.Repeat("d", 40), wantSupport: true, wantRef: "refs/heads/main", wantRefName: "main", wantSHA: strings.Repeat("d", 40)}, + {name: "merged payload with non-closed activity", action: "synchronize", state: "closed", merged: true, mergeable: &conflicted, mergeSHA: strings.Repeat("e", 40), wantError: true}, } { t.Run(tt.name, func(t *testing.T) { event := &payload{Action: tt.action} @@ -101,9 +104,10 @@ func TestNormalizePullRequestRevisions(t *testing.T) { event.PullRequest.State = tt.state event.PullRequest.Merged = tt.merged event.PullRequest.Mergeable = tt.mergeable - event.PullRequest.MergeCommitSHA = tt.sha + event.PullRequest.MergeCommitSHA = tt.mergeSHA event.Repository.ID = 1 event.PullRequest.Head.Ref = "feature" + event.PullRequest.Head.SHA = headSHA event.PullRequest.Head.Repository.ID = event.Repository.ID event.PullRequest.Base.Ref = "main" normalized, supported, err := normalize("pull_request", event) @@ -122,13 +126,32 @@ func TestNormalizePullRequestRevisions(t *testing.T) { if !supported { return } - if normalized.Ref != tt.wantRef || githubclient.RefName(normalized.Ref) != tt.wantRefName || normalized.HeadRef != "feature" || normalized.BaseRef != "main" || normalized.SHA != tt.sha || normalized.ResolveRef != tt.wantResolve { + if normalized.Ref != tt.wantRef || githubclient.RefName(normalized.Ref) != tt.wantRefName || normalized.HeadRef != "feature" || normalized.BaseRef != "main" || normalized.SHA != tt.wantSHA || normalized.ResolveRef != tt.wantResolve || normalized.HeadSHA != tt.wantHeadSHA { t.Errorf("normalized event = %#v", normalized) } }) } } +func TestNormalizePullRequestRequiresValidHeadRevision(t *testing.T) { + mergeable := true + for _, headSHA := range []string{"", "not-a-sha", strings.Repeat("A", 40), zeroGitSHA} { + event := &payload{Action: "synchronize"} + event.Repository.ID = 1 + event.PullRequest.Number = 42 + event.PullRequest.State = "open" + event.PullRequest.Mergeable = &mergeable + event.PullRequest.Head.Ref = "feature" + event.PullRequest.Head.SHA = headSHA + event.PullRequest.Head.Repository.ID = event.Repository.ID + event.PullRequest.Base.Ref = "main" + + if _, _, err := normalize("pull_request", event); err == nil { + t.Fatalf("normalize() accepted head revision %q", headSHA) + } + } +} + func TestNormalizeMergeGroupActions(t *testing.T) { for _, tt := range []struct { action string