Repository navigation
Expand file tree
/
Copy pathservice_after_rollback_test.go
More file actions
126 lines (122 loc) · 4.99 KB
/
Copy pathservice_after_rollback_test.go
File metadata and controls
126 lines (122 loc) · 4.99 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
package kata_test
import (
"bytes"
"context"
"database/sql"
"errors"
"net/http"
"net/http/httptest"
"path/filepath"
"strconv"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"go.kenn.io/kata"
"go.kenn.io/kata/internal/testenv"
)
func TestHostDenialFinalizesAfterRollback(t *testing.T) {
for _, backend := range []string{"sqlite", "postgres"} {
t.Run(backend, func(t *testing.T) {
ctx, cancel := context.WithTimeout(t.Context(), 30*time.Second)
defer cancel()
config := kata.Config{DSN: filepath.Join(t.TempDir(), "service.db")}
driver, prefix := "sqlite", ""
if backend == "postgres" {
dsn, cleanup := testenv.NewPostgresContainer(t, ctx)
t.Cleanup(cleanup)
config.DSN = dsn
config.Postgres = kata.PostgresConfig{Schema: "kata", SchemaMode: kata.PostgresSchemaBootstrap}
driver, prefix = "pgx", "kata."
}
controller := &recordingAccessController{}
config.Access = controller
service, err := kata.New(ctx, config)
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, service.Close()) })
project, err := service.EnsureProject(ctx, kata.ProjectSpec{
UID: "01HZNQ7VFPK1XGD8R5MABCD4EX", Name: "example-project",
})
require.NoError(t, err)
inspection, err := sql.Open(driver, config.DSN)
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, inspection.Close()) })
_, err = inspection.ExecContext(ctx, `CREATE TABLE `+prefix+`fence_markers (attempt INTEGER NOT NULL)`)
require.NoError(t, err)
for _, testCase := range []struct {
name string
finishError error
wantStatus int
wantRecords int
}{
{name: "recorded", wantStatus: http.StatusNotFound, wantRecords: 1},
{name: "recording failed", finishError: errors.New("host recording unavailable"), wantStatus: http.StatusServiceUnavailable},
} {
t.Run(testCase.name, func(t *testing.T) {
_, err := inspection.ExecContext(ctx, `DELETE FROM `+prefix+`fence_markers`)
require.NoError(t, err)
calls := 0
controller.transactionFence = func(ctx context.Context, tx kata.Transaction) error {
if _, err := tx.ExecContext(ctx, `INSERT INTO `+prefix+`fence_markers VALUES (1)`); err != nil {
return err
}
return kata.AfterTransactionRollback(kata.ErrAccessDenied, func(ctx context.Context) error {
calls++
var markers int
if err := inspection.QueryRowContext(ctx, `SELECT count(*) FROM `+prefix+`fence_markers`).Scan(&markers); err != nil {
return err
}
assert.Zero(t, markers, "the callback must see the completed rollback")
if testCase.finishError != nil {
return testCase.finishError
}
// This write uses a separate transaction and must survive denial.
_, err := inspection.ExecContext(ctx, `INSERT INTO `+prefix+`fence_markers VALUES (2)`)
return err
})
}
request := httptest.NewRequestWithContext(ctx, http.MethodPost,
"/api/v1/projects/"+strconv.FormatInt(project.Project.ID, 10)+"/issues",
bytes.NewBufferString(`{"actor":"ignored","title":"must not be stored"}`))
request.Header.Set("Content-Type", "application/json")
request = request.WithContext(kata.WithPrincipal(request.Context(), kata.Principal{Subject: "user-a", Actor: "Example User"}))
response := httptest.NewRecorder()
service.Handler().ServeHTTP(response, request)
assert.Equal(t, testCase.wantStatus, response.Code)
if testCase.finishError != nil {
assert.Contains(t, response.Body.String(), `"code":"access_unavailable"`)
assert.NotContains(t, response.Body.String(), testCase.finishError.Error())
}
assert.Equal(t, 1, calls)
var records, issues int
require.NoError(t, inspection.QueryRowContext(ctx, `SELECT count(*) FROM `+prefix+`fence_markers WHERE attempt = 2`).Scan(&records))
assert.Equal(t, testCase.wantRecords, records)
require.NoError(t, inspection.QueryRowContext(ctx, `SELECT count(*) FROM `+prefix+`issues`).Scan(&issues))
assert.Zero(t, issues)
})
}
})
}
}
func TestFederationDenialKeepsRollbackCallback(t *testing.T) {
for _, finishError := range []error{nil, errors.New("host recording unavailable")} {
calls := 0
controller := &recordingFederationAccessController{decide: func(kata.FederationAccessRequest) (kata.FederationAccessDecision, error) {
return kata.FederationAccessDecision{TransactionFence: func(context.Context, kata.Transaction) error {
return kata.AfterTransactionRollback(kata.ErrAccessDenied, func(context.Context) error {
calls++
return finishError
})
}}, nil
}}
service, project, enrollment := newFederationAccessService(t, controller)
issue := createFederationAccessIssue(t, service, project.ID)
response := acquireFederationAccessClaim(t, service, project.ID, issue, enrollment.Token)
wantStatus := http.StatusForbidden
if finishError != nil {
wantStatus = http.StatusServiceUnavailable
}
assert.Equal(t, wantStatus, response.Code)
assert.Equal(t, 1, calls)
}
}