1 package testutil_test
2
3 import (
4 "context"
5 "net"
6 "os"
7 "testing"
8
9 "github.com/cybertec-postgresql/pgwatch/v6/internal/testutil"
10 "github.com/stretchr/testify/assert"
11 "github.com/stretchr/testify/require"
12 "google.golang.org/grpc/codes"
13 "google.golang.org/grpc/metadata"
14 "google.golang.org/grpc/status"
15 )
16
17 func TestLoadServerTLSCredentials(t *testing.T) {
18 t.Run("valid credentials", func(t *testing.T) {
19 creds, err := testutil.LoadServerTLSCredentials()
20 assert.NoError(t, err)
21 assert.NotNil(t, creds)
22 assert.Equal(t, "tls", creds.Info().SecurityProtocol)
23 })
24
25 t.Run("invalid certificate", func(t *testing.T) {
26
27 origCert := testutil.Cert
28 origKey := testutil.PrivateKey
29 defer func() {
30 testutil.Cert = origCert
31 testutil.PrivateKey = origKey
32 }()
33
34
35 testutil.Cert = []byte("invalid cert")
36 testutil.PrivateKey = []byte("invalid key")
37
38 creds, err := testutil.LoadServerTLSCredentials()
39 assert.Error(t, err)
40 assert.Nil(t, creds)
41 })
42 }
43
44 func TestAuthInterceptor(t *testing.T) {
45 handler := func(context.Context, any) (any, error) {
46 return "success", nil
47 }
48
49 t.Run("valid credentials", func(t *testing.T) {
50 md := metadata.Pairs("username", "pgwatch", "password", "pgwatch")
51 ctx := metadata.NewIncomingContext(context.Background(), md)
52
53 result, err := testutil.AuthInterceptor(ctx, nil, nil, handler)
54 assert.NoError(t, err)
55 assert.Equal(t, "success", result)
56 })
57
58 t.Run("empty credentials", func(t *testing.T) {
59 md := metadata.Pairs("username", "", "password", "")
60 ctx := metadata.NewIncomingContext(context.Background(), md)
61
62 result, err := testutil.AuthInterceptor(ctx, nil, nil, handler)
63 assert.NoError(t, err)
64 assert.Equal(t, "success", result)
65 })
66
67 t.Run("invalid credentials", func(t *testing.T) {
68 md := metadata.Pairs("username", "wrong", "password", "wrong")
69 ctx := metadata.NewIncomingContext(context.Background(), md)
70
71 result, err := testutil.AuthInterceptor(ctx, nil, nil, handler)
72 assert.Error(t, err)
73 assert.Nil(t, result)
74
75 st, ok := status.FromError(err)
76 require.True(t, ok)
77 assert.Equal(t, codes.Unauthenticated, st.Code())
78 })
79 }
80
81 func TestSetupPostgresContainer(t *testing.T) {
82 if testing.Short() {
83 t.Skip("Skipping container test in short mode")
84 }
85
86 container, teardown, err := testutil.SetupPostgresContainer()
87 for i := range 2 {
88 if i == 1 {
89 container, teardown, err = testutil.SetupPostgresContainerWithInitScripts("../../docker/bootstrap/create_role_db.sql")
90 }
91
92 if err != nil {
93 t.Skipf("Skipping postgres container test: %v", err)
94 return
95 }
96 defer teardown()
97
98 assert.NotNil(t, container)
99
100
101 state, err := container.State(context.Background())
102 require.NoError(t, err)
103 assert.True(t, state.Running)
104
105
106 connStr, err := container.ConnectionString(context.Background())
107 require.NoError(t, err)
108 assert.NotEmpty(t, connStr)
109 }
110 }
111
112 func TestSetupPostgresContainerWithConfig(t *testing.T) {
113 if testing.Short() {
114 t.Skip("Skipping container test in short mode")
115 }
116
117
118 tempDir := t.TempDir()
119 configPath := tempDir + "/postgresql.conf"
120 configContent := `
121 listen_addresses = '*'
122 log_destination = 'csvlog'
123 logging_collector = on
124 log_directory = 'pg_log'
125 log_filename = 'pgwatch.csv'
126 `
127 err := os.WriteFile(configPath, []byte(configContent), 0644)
128 require.NoError(t, err)
129
130 container, teardown, err := testutil.SetupPostgresContainerWithConfig(configPath)
131 require.NoError(t, err)
132 defer teardown()
133
134 require.NotNil(t, container)
135
136
137 state, err := container.State(context.Background())
138 require.NoError(t, err)
139 assert.True(t, state.Running)
140
141
142 connStr, err := container.ConnectionString(context.Background())
143 assert.NoError(t, err)
144 assert.NotEmpty(t, connStr)
145 }
146
147 func TestSetupEtcdContainer(t *testing.T) {
148 if testing.Short() {
149 t.Skip("Skipping etcd container test in short mode")
150 }
151
152 etcdContainer, etcdTeardown, err := testutil.SetupEtcdContainer()
153 require.NoError(t, err)
154 defer etcdTeardown()
155
156
157 state, err := etcdContainer.State(context.Background())
158 require.NoError(t, err)
159 assert.True(t, state.Running)
160 }
161
162 func TestSetupRPCServers(t *testing.T) {
163 teardown, err := testutil.SetupRPCServers()
164 require.NoError(t, err)
165 require.NotNil(t, teardown)
166 defer teardown()
167
168
169 _, statErr := os.Stat(testutil.CAFile)
170 assert.NoError(t, statErr, "CA file should exist after SetupRPCServers")
171
172
173 conn, dialErr := net.Dial("tcp", testutil.PlainServerAddress)
174 require.NoError(t, dialErr, "plain gRPC server should be listening on %s", testutil.PlainServerAddress)
175 conn.Close()
176
177
178 conn, dialErr = net.Dial("tcp", testutil.TLSServerAddress)
179 require.NoError(t, dialErr, "TLS gRPC server should be listening on %s", testutil.TLSServerAddress)
180 conn.Close()
181 }
182