fix(cfsetup): upsert the Access policy — never swallow already-exists over a broader rule set
CreateAccessPolicy treated a Cloudflare "policy_already_exists" as idempotent success and kept whatever policy was there. On a re-run with a changed identity — or against a hand-made broader policy — op.console would stay guarded by something weaker than the fail-closed body this package builds and guards, while Setup reported success. The fail-closed validation only ever ran on the policy we built, never on the one that stayed live. Now it upserts by name: lookup, PUT the guarded body over the existing policy, POST only when absent (a racing POST re-looks up and PUTs). apiPost/apiPut share one apiWrite; three httptest cases pin update-over-existing, create-when- absent, and the race fallback.
This commit is contained in:
2 files changed
+164
-8
No files matched your search
@@ -225,18 +225,53 @@ func (r *ExecRunner) CreateAccessApplication(ctx context.Context, app AccessAppl
|
|||||||
return resp.Result.ID, resp.Result.AUD, nil
|
return resp.Result.ID, resp.Result.AUD, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// CreateAccessPolicy POSTs the policy onto the Access app.
|
// CreateAccessPolicy makes the Access app carry exactly the recommended policy.
|
||||||
|
// It upserts by name instead of POST-then-swallow: when a policy of this name
|
||||||
|
// already exists — a re-run, or a previous hand-made setup — the swallow looked
|
||||||
|
// idempotent but left the OLD rule set in place. If that old policy is broader
|
||||||
|
// than the fail-closed one just built (a changed identity, a hand-made
|
||||||
|
// allow-everyone rule), op.console ends up guarded by something weaker while
|
||||||
|
// Setup reports success. So: find it and PUT our body over it; POST only when
|
||||||
|
// absent (a POST that races into "already exists" falls back to the PUT).
|
||||||
func (r *ExecRunner) CreateAccessPolicy(ctx context.Context, appID string, policy AccessPolicy) error {
|
func (r *ExecRunner) CreateAccessPolicy(ctx context.Context, appID string, policy AccessPolicy) error {
|
||||||
|
pid, err := r.lookupAccessPolicy(ctx, appID, policy.Name)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if pid != "" {
|
||||||
|
return r.apiPut(ctx, fmt.Sprintf("/accounts/%s/access/apps/%s/policies/%s", r.AccountID, appID, pid), policy, nil)
|
||||||
|
}
|
||||||
if err := r.apiPost(ctx, fmt.Sprintf("/accounts/%s/access/apps/%s/policies", r.AccountID, appID), policy, nil); err != nil {
|
if err := r.apiPost(ctx, fmt.Sprintf("/accounts/%s/access/apps/%s/policies", r.AccountID, appID), policy, nil); err != nil {
|
||||||
// If policy already exists, treat it as idempotent success
|
if strings.Contains(err.Error(), "already_exists") || strings.Contains(err.Error(), "11015") {
|
||||||
if strings.Contains(err.Error(), "policy_already_exists") || strings.Contains(err.Error(), "11015") || strings.Contains(err.Error(), "already_exists") {
|
if pid, lerr := r.lookupAccessPolicy(ctx, appID, policy.Name); lerr == nil && pid != "" {
|
||||||
return nil
|
return r.apiPut(ctx, fmt.Sprintf("/accounts/%s/access/apps/%s/policies/%s", r.AccountID, appID, pid), policy, nil)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// lookupAccessPolicy finds the id of the policy named name on an Access app,
|
||||||
|
// returning "" when absent.
|
||||||
|
func (r *ExecRunner) lookupAccessPolicy(ctx context.Context, appID, name string) (string, error) {
|
||||||
|
var resp struct {
|
||||||
|
Result []struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
} `json:"result"`
|
||||||
|
}
|
||||||
|
if err := r.apiGet(ctx, fmt.Sprintf("/accounts/%s/access/apps/%s/policies?per_page=100", r.AccountID, appID), &resp); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
for _, p := range resp.Result {
|
||||||
|
if p.Name == name {
|
||||||
|
return p.ID, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return "", nil
|
||||||
|
}
|
||||||
|
|
||||||
// runCloudflared executes the cloudflared binary with the given args, returning
|
// runCloudflared executes the cloudflared binary with the given args, returning
|
||||||
// combined output. The interactive `tunnel login` browser consent is NOT done
|
// combined output. The interactive `tunnel login` browser consent is NOT done
|
||||||
// here — it is a separate, operator-driven step the TUI suspends to run.
|
// here — it is a separate, operator-driven step the TUI suspends to run.
|
||||||
@@ -255,15 +290,23 @@ func (r *ExecRunner) runCloudflared(ctx context.Context, args ...string) (string
|
|||||||
return buf.String(), nil
|
return buf.String(), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// apiPost sends an authenticated JSON POST to the Cloudflare API and, on a
|
// apiPost and apiPut send an authenticated JSON write to the Cloudflare API and,
|
||||||
// non-2xx or success:false body, returns the error. out, when non-nil, receives
|
// on a non-2xx or success:false body, return the error. out, when non-nil,
|
||||||
// the decoded response.
|
// receives the decoded response.
|
||||||
func (r *ExecRunner) apiPost(ctx context.Context, path string, body, out any) error {
|
func (r *ExecRunner) apiPost(ctx context.Context, path string, body, out any) error {
|
||||||
|
return r.apiWrite(ctx, http.MethodPost, path, body, out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *ExecRunner) apiPut(ctx context.Context, path string, body, out any) error {
|
||||||
|
return r.apiWrite(ctx, http.MethodPut, path, body, out)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *ExecRunner) apiWrite(ctx context.Context, method, path string, body, out any) error {
|
||||||
payload, err := json.Marshal(body)
|
payload, err := json.Marshal(body)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, r.apiBase()+path, bytes.NewReader(payload))
|
req, err := http.NewRequestWithContext(ctx, method, r.apiBase()+path, bytes.NewReader(payload))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package cfsetup
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -116,3 +117,115 @@ func TestExecRunnerSatisfiesAPITokenVerifier(t *testing.T) {
|
|||||||
t.Fatal("ExecRunner must implement apiTokenVerifier so Setup verifies the token before side effects")
|
t.Fatal("ExecRunner must implement apiTokenVerifier so Setup verifies the token before side effects")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func recommendedPolicy(t *testing.T) AccessPolicy {
|
||||||
|
t.Helper()
|
||||||
|
p, err := BuildRecommendedPolicy("felis-recommended", AccessIdentity{Emails: []string{"[email protected]"}})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("build policy: %v", err)
|
||||||
|
}
|
||||||
|
return p
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestExecRunnerCreateAccessPolicyUpdatesExisting is the regression for the
|
||||||
|
// fail-open swap: an Access app that already carries a policy of this name must
|
||||||
|
// be OVERWRITTEN with the guarded body. The old behavior swallowed
|
||||||
|
// "already exists" as success, so a re-run with a changed identity (or a
|
||||||
|
// hand-made broader policy) silently kept the old rule set while reporting a
|
||||||
|
// completed setup.
|
||||||
|
func TestExecRunnerCreateAccessPolicyUpdatesExisting(t *testing.T) {
|
||||||
|
var methods []string
|
||||||
|
var putBody string
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
methods = append(methods, r.Method)
|
||||||
|
switch {
|
||||||
|
case r.Method == http.MethodGet:
|
||||||
|
_, _ = w.Write([]byte(`{"success":true,"errors":[],"result":[{"id":"pol-1","name":"felis-recommended"}]}`))
|
||||||
|
case r.Method == http.MethodPut && strings.HasSuffix(r.URL.Path, "/policies/pol-1"):
|
||||||
|
b, _ := io.ReadAll(r.Body)
|
||||||
|
putBody = string(b)
|
||||||
|
_, _ = w.Write([]byte(`{"success":true,"errors":[],"result":{}}`))
|
||||||
|
default:
|
||||||
|
t.Fatalf("unexpected %s %s", r.Method, r.URL.Path)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
r := &ExecRunner{APIToken: "t", AccountID: "acc", APIBase: srv.URL}
|
||||||
|
if err := r.CreateAccessPolicy(context.Background(), "app-1", recommendedPolicy(t)); err != nil {
|
||||||
|
t.Fatalf("CreateAccessPolicy = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if len(methods) != 2 || methods[0] != http.MethodGet || methods[1] != http.MethodPut {
|
||||||
|
t.Fatalf("call sequence = %v, want [GET PUT] (lookup then overwrite)", methods)
|
||||||
|
}
|
||||||
|
if !strings.Contains(putBody, `"email"`) || !strings.Contains(putBody, "felis-recommended") {
|
||||||
|
t.Fatalf("PUT body must carry the full guarded policy, got %s", putBody)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestExecRunnerCreateAccessPolicyCreatesWhenAbsent: the happy path still POSTs
|
||||||
|
// when no policy of this name exists.
|
||||||
|
func TestExecRunnerCreateAccessPolicyCreatesWhenAbsent(t *testing.T) {
|
||||||
|
var methods []string
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
methods = append(methods, r.Method)
|
||||||
|
switch {
|
||||||
|
case r.Method == http.MethodGet:
|
||||||
|
_, _ = w.Write([]byte(`{"success":true,"errors":[],"result":[]}`))
|
||||||
|
case r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/policies"):
|
||||||
|
_, _ = w.Write([]byte(`{"success":true,"errors":[],"result":{}}`))
|
||||||
|
default:
|
||||||
|
t.Fatalf("unexpected %s %s", r.Method, r.URL.Path)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
r := &ExecRunner{APIToken: "t", AccountID: "acc", APIBase: srv.URL}
|
||||||
|
if err := r.CreateAccessPolicy(context.Background(), "app-1", recommendedPolicy(t)); err != nil {
|
||||||
|
t.Fatalf("CreateAccessPolicy = %v, want nil", err)
|
||||||
|
}
|
||||||
|
if len(methods) != 2 || methods[0] != http.MethodGet || methods[1] != http.MethodPost {
|
||||||
|
t.Fatalf("call sequence = %v, want [GET POST]", methods)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestExecRunnerCreateAccessPolicyRaceFallsBackToUpdate: a POST losing to a
|
||||||
|
// concurrent creator ("already exists") must re-lookup and PUT, never swallow.
|
||||||
|
func TestExecRunnerCreateAccessPolicyRaceFallsBackToUpdate(t *testing.T) {
|
||||||
|
var methods []string
|
||||||
|
getCount := 0
|
||||||
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
methods = append(methods, r.Method)
|
||||||
|
switch r.Method {
|
||||||
|
case http.MethodGet:
|
||||||
|
getCount++
|
||||||
|
if getCount == 1 {
|
||||||
|
_, _ = w.Write([]byte(`{"success":true,"errors":[],"result":[]}`))
|
||||||
|
} else {
|
||||||
|
_, _ = w.Write([]byte(`{"success":true,"errors":[],"result":[{"id":"pol-9","name":"felis-recommended"}]}`))
|
||||||
|
}
|
||||||
|
case http.MethodPost:
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
_, _ = w.Write([]byte(`{"success":false,"errors":[{"code":11015,"message":"policy_already_exists"}]}`))
|
||||||
|
case http.MethodPut:
|
||||||
|
_, _ = w.Write([]byte(`{"success":true,"errors":[],"result":{}}`))
|
||||||
|
default:
|
||||||
|
t.Fatalf("unexpected %s", r.Method)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer srv.Close()
|
||||||
|
|
||||||
|
r := &ExecRunner{APIToken: "t", AccountID: "acc", APIBase: srv.URL}
|
||||||
|
if err := r.CreateAccessPolicy(context.Background(), "app-1", recommendedPolicy(t)); err != nil {
|
||||||
|
t.Fatalf("CreateAccessPolicy = %v, want nil after race fallback", err)
|
||||||
|
}
|
||||||
|
want := []string{"GET", "POST", "GET", "PUT"}
|
||||||
|
if len(methods) != len(want) {
|
||||||
|
t.Fatalf("call sequence = %v, want %v", methods, want)
|
||||||
|
}
|
||||||
|
for i := range want {
|
||||||
|
if methods[i] != want[i] {
|
||||||
|
t.Fatalf("call sequence = %v, want %v", methods, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in new issue
Block a user