Skip to content

Commit b08314a

Browse files
committed
feat(azure): add WithClient variants and azfake unit tests for all modules
Add exported WithClient variants for Azure module functions that accept a pre-built SDK client, enabling unit testing with fake servers. Each existing ContextE function now delegates its SDK call to the corresponding WithClient variant. Pattern established across all modules: - GetDiskContextE(ctx, name, rg, sub) -- convenience, creates its own client - GetDiskWithClient(ctx, client, rg, name) -- testable, accepts injected client Functions that extract fields from response objects (not clients) drop the WithClient suffix since they don't take a client: - GetAvailabilitySetFaultDomainCount(avs) - GetLoadBalancerFrontendIPConfigNames(lb) - GetNetworkInterfacePrivateIPs(nic) - GetIPOfPublicIPAddressByName(pip) - GetVirtualNetworkDNSServerIPs(vnet) Nil-deref guards added on all newly-exported helpers that dereference response properties (avs.Properties, lb.Properties, nic.Properties, feConfig.Properties, vnet.Properties, pip.Name). Modules updated: disk, availabilityset, loadbalancer, networkinterface, publicaddress, virtualnetwork, cosmosdb, servicebus All WithClient functions are tested using Azure SDK's azfake framework (httptest for servicebus beta SDK), running in CI without credentials.
1 parent 57fac81 commit b08314a

16 files changed

Lines changed: 1578 additions & 96 deletions

modules/azure/availabilityset.go

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package azure
22

33
import (
44
"context"
5+
"errors"
56
"strings"
67

78
"github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/compute/armcompute/v6"
@@ -94,12 +95,21 @@ func CheckAvailabilitySetContainsVMContextE(t testing.TestingT, ctx context.Cont
9495
return false, err
9596
}
9697

98+
return CheckAvailabilitySetContainsVMWithClient(ctx, client, resGroupName, avsName, vmName)
99+
}
100+
101+
// CheckAvailabilitySetContainsVMWithClient checks if the Virtual Machine is contained in the Availability Set VMs
102+
// using the provided AvailabilitySetsClient.
103+
func CheckAvailabilitySetContainsVMWithClient(ctx context.Context, client *armcompute.AvailabilitySetsClient, resGroupName string, avsName string, vmName string) (bool, error) {
97104
resp, err := client.Get(ctx, resGroupName, avsName, nil)
98105
if err != nil {
99106
return false, err
100107
}
101108

102109
for _, vm := range resp.Properties.VirtualMachines {
110+
if vm.ID == nil {
111+
continue
112+
}
103113
// VM IDs are always ALL CAPS in this property so ignoring case
104114
if strings.EqualFold(vmName, GetNameFromResourceID(*vm.ID)) {
105115
return true, nil
@@ -148,6 +158,12 @@ func GetAvailabilitySetVMNamesInCapsContextE(t testing.TestingT, ctx context.Con
148158
return nil, err
149159
}
150160

161+
return GetAvailabilitySetVMNamesInCapsWithClient(ctx, client, resGroupName, avsName)
162+
}
163+
164+
// GetAvailabilitySetVMNamesInCapsWithClient gets a list of VM names in the specified Azure Availability Set
165+
// using the provided AvailabilitySetsClient.
166+
func GetAvailabilitySetVMNamesInCapsWithClient(ctx context.Context, client *armcompute.AvailabilitySetsClient, resGroupName string, avsName string) ([]string, error) {
151167
resp, err := client.Get(ctx, resGroupName, avsName, nil)
152168
if err != nil {
153169
return nil, err
@@ -156,6 +172,9 @@ func GetAvailabilitySetVMNamesInCapsContextE(t testing.TestingT, ctx context.Con
156172
vms := []string{}
157173

158174
for _, vm := range resp.Properties.VirtualMachines {
175+
if vm.ID == nil {
176+
continue
177+
}
159178
// IDs are returned in ALL CAPS for this property
160179
if vmName := GetNameFromResourceID(*vm.ID); len(vmName) > 0 {
161180
vms = append(vms, vmName)
@@ -204,6 +223,15 @@ func GetAvailabilitySetFaultDomainCountContextE(t testing.TestingT, ctx context.
204223
return -1, err
205224
}
206225

226+
return ExtractAvailabilitySetFaultDomainCount(avs)
227+
}
228+
229+
// ExtractAvailabilitySetFaultDomainCount gets the Fault Domain Count from the provided AvailabilitySet.
230+
func ExtractAvailabilitySetFaultDomainCount(avs *armcompute.AvailabilitySet) (int32, error) {
231+
if avs.Properties == nil || avs.Properties.PlatformFaultDomainCount == nil {
232+
return -1, errors.New("availability set has no fault domain count")
233+
}
234+
207235
return *avs.Properties.PlatformFaultDomainCount, nil
208236
}
209237

@@ -229,6 +257,11 @@ func GetAvailabilitySetContextE(t testing.TestingT, ctx context.Context, avsName
229257
return nil, err
230258
}
231259

260+
return GetAvailabilitySetWithClient(ctx, client, resGroupName, avsName)
261+
}
262+
263+
// GetAvailabilitySetWithClient gets an Availability Set using the provided AvailabilitySetsClient.
264+
func GetAvailabilitySetWithClient(ctx context.Context, client *armcompute.AvailabilitySetsClient, resGroupName string, avsName string) (*armcompute.AvailabilitySet, error) {
232265
resp, err := client.Get(ctx, resGroupName, avsName, nil)
233266
if err != nil {
234267
return nil, err
Lines changed: 149 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,149 @@
1+
package azure_test
2+
3+
import (
4+
"context"
5+
"net/http"
6+
"testing"
7+
8+
"github.com/Azure/azure-sdk-for-go/sdk/azcore/arm"
9+
azfake "github.com/Azure/azure-sdk-for-go/sdk/azcore/fake"
10+
"github.com/Azure/azure-sdk-for-go/sdk/azcore/policy"
11+
"github.com/Azure/azure-sdk-for-go/sdk/azcore/to"
12+
"github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/compute/armcompute/v6"
13+
computefake "github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/compute/armcompute/v6/fake"
14+
"github.com/gruntwork-io/terratest/modules/azure"
15+
"github.com/stretchr/testify/assert"
16+
"github.com/stretchr/testify/require"
17+
)
18+
19+
func newFakeAvailabilitySetsClient(t *testing.T, srv *computefake.AvailabilitySetsServer) *armcompute.AvailabilitySetsClient {
20+
t.Helper()
21+
22+
transport := computefake.NewAvailabilitySetsServerTransport(srv)
23+
client, err := armcompute.NewAvailabilitySetsClient("fake-sub", &azfake.TokenCredential{}, &arm.ClientOptions{
24+
ClientOptions: policy.ClientOptions{Transport: transport},
25+
})
26+
require.NoError(t, err)
27+
28+
return client
29+
}
30+
31+
func fakeAvsGetHandler(avsName string, vmIDs []string, faultDomainCount int32) func(context.Context, string, string, *armcompute.AvailabilitySetsClientGetOptions) (azfake.Responder[armcompute.AvailabilitySetsClientGetResponse], azfake.ErrorResponder) {
32+
return func(_ context.Context, _ string, _ string, _ *armcompute.AvailabilitySetsClientGetOptions) (resp azfake.Responder[armcompute.AvailabilitySetsClientGetResponse], errResp azfake.ErrorResponder) {
33+
vms := make([]*armcompute.SubResource, len(vmIDs))
34+
for i, id := range vmIDs {
35+
vms[i] = &armcompute.SubResource{ID: to.Ptr(id)}
36+
}
37+
38+
resp.SetResponse(http.StatusOK, armcompute.AvailabilitySetsClientGetResponse{
39+
AvailabilitySet: armcompute.AvailabilitySet{
40+
Name: to.Ptr(avsName),
41+
Properties: &armcompute.AvailabilitySetProperties{
42+
VirtualMachines: vms,
43+
PlatformFaultDomainCount: to.Ptr(faultDomainCount),
44+
},
45+
},
46+
}, nil)
47+
48+
return
49+
}
50+
}
51+
52+
func TestGetAvailabilitySetWithClient(t *testing.T) {
53+
t.Parallel()
54+
55+
srv := &computefake.AvailabilitySetsServer{
56+
Get: fakeAvsGetHandler("my-avs", nil, 2),
57+
}
58+
client := newFakeAvailabilitySetsClient(t, srv)
59+
60+
avs, err := azure.GetAvailabilitySetWithClient(t.Context(), client, "rg", "my-avs")
61+
require.NoError(t, err)
62+
assert.Equal(t, "my-avs", *avs.Name)
63+
}
64+
65+
func TestCheckAvailabilitySetContainsVMWithClient(t *testing.T) {
66+
t.Parallel()
67+
68+
vmIDs := []string{
69+
"/subscriptions/sub/resourceGroups/RG/providers/Microsoft.Compute/virtualMachines/VM-ONE",
70+
"/subscriptions/sub/resourceGroups/RG/providers/Microsoft.Compute/virtualMachines/VM-TWO",
71+
}
72+
73+
tests := []struct {
74+
name string
75+
vmName string
76+
found bool
77+
wantErr bool
78+
}{
79+
{name: "exact case match", vmName: "VM-ONE", found: true},
80+
{name: "case insensitive match", vmName: "vm-one", found: true},
81+
{name: "not found", vmName: "vm-three", found: false, wantErr: true},
82+
}
83+
84+
for _, tc := range tests {
85+
t.Run(tc.name, func(t *testing.T) {
86+
t.Parallel()
87+
88+
srv := &computefake.AvailabilitySetsServer{
89+
Get: fakeAvsGetHandler("avs", vmIDs, 2),
90+
}
91+
client := newFakeAvailabilitySetsClient(t, srv)
92+
93+
found, err := azure.CheckAvailabilitySetContainsVMWithClient(t.Context(), client, "rg", "avs", tc.vmName)
94+
95+
if tc.wantErr {
96+
require.Error(t, err)
97+
} else {
98+
require.NoError(t, err)
99+
}
100+
101+
assert.Equal(t, tc.found, found)
102+
})
103+
}
104+
}
105+
106+
func TestGetAvailabilitySetVMNamesInCapsWithClient(t *testing.T) {
107+
t.Parallel()
108+
109+
vmIDs := []string{
110+
"/subscriptions/sub/resourceGroups/RG/providers/Microsoft.Compute/virtualMachines/VM-ALPHA",
111+
"/subscriptions/sub/resourceGroups/RG/providers/Microsoft.Compute/virtualMachines/VM-BETA",
112+
}
113+
114+
srv := &computefake.AvailabilitySetsServer{
115+
Get: fakeAvsGetHandler("avs", vmIDs, 3),
116+
}
117+
client := newFakeAvailabilitySetsClient(t, srv)
118+
119+
names, err := azure.GetAvailabilitySetVMNamesInCapsWithClient(t.Context(), client, "rg", "avs")
120+
require.NoError(t, err)
121+
assert.Equal(t, []string{"VM-ALPHA", "VM-BETA"}, names)
122+
}
123+
124+
func TestExtractAvailabilitySetFaultDomainCount(t *testing.T) {
125+
t.Parallel()
126+
127+
t.Run("valid", func(t *testing.T) {
128+
t.Parallel()
129+
130+
avs := &armcompute.AvailabilitySet{
131+
Properties: &armcompute.AvailabilitySetProperties{
132+
PlatformFaultDomainCount: to.Ptr[int32](3),
133+
},
134+
}
135+
136+
count, err := azure.ExtractAvailabilitySetFaultDomainCount(avs)
137+
require.NoError(t, err)
138+
assert.Equal(t, int32(3), count)
139+
})
140+
141+
t.Run("nil properties", func(t *testing.T) {
142+
t.Parallel()
143+
144+
avs := &armcompute.AvailabilitySet{}
145+
146+
_, err := azure.ExtractAvailabilitySetFaultDomainCount(avs)
147+
require.Error(t, err)
148+
})
149+
}

modules/azure/cosmosdb.go

Lines changed: 32 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -73,13 +73,16 @@ func GetCosmosDBAccountContextE(ctx context.Context, subscriptionID string, reso
7373
return nil, err
7474
}
7575

76-
// Get the corresponding database account
77-
resp, err := cosmosClient.Get(ctx, resourceGroupName, accountName, nil)
76+
return GetCosmosDBAccountWithClient(ctx, cosmosClient, resourceGroupName, accountName)
77+
}
78+
79+
// GetCosmosDBAccountWithClient gets a database account using the provided DatabaseAccountsClient.
80+
func GetCosmosDBAccountWithClient(ctx context.Context, client *armcosmos.DatabaseAccountsClient, resourceGroupName string, accountName string) (*armcosmos.DatabaseAccountGetResults, error) {
81+
resp, err := client.Get(ctx, resourceGroupName, accountName, nil)
7882
if err != nil {
7983
return nil, err
8084
}
8185

82-
// Return DB
8386
return &resp.DatabaseAccountGetResults, nil
8487
}
8588

@@ -148,13 +151,16 @@ func GetCosmosDBSQLDatabaseContextE(ctx context.Context, subscriptionID string,
148151
return nil, err
149152
}
150153

151-
// Get the corresponding database
152-
resp, err := cosmosClient.GetSQLDatabase(ctx, resourceGroupName, accountName, databaseName, nil)
154+
return GetCosmosDBSQLDatabaseWithClient(ctx, cosmosClient, resourceGroupName, accountName, databaseName)
155+
}
156+
157+
// GetCosmosDBSQLDatabaseWithClient gets a SQL database using the provided SQLResourcesClient.
158+
func GetCosmosDBSQLDatabaseWithClient(ctx context.Context, client *armcosmos.SQLResourcesClient, resourceGroupName string, accountName string, databaseName string) (*armcosmos.SQLDatabaseGetResults, error) {
159+
resp, err := client.GetSQLDatabase(ctx, resourceGroupName, accountName, databaseName, nil)
153160
if err != nil {
154161
return nil, err
155162
}
156163

157-
// Return DB
158164
return &resp.SQLDatabaseGetResults, nil
159165
}
160166

@@ -192,13 +198,16 @@ func GetCosmosDBSQLContainerContextE(ctx context.Context, subscriptionID string,
192198
return nil, err
193199
}
194200

195-
// Get the corresponding SQL container
196-
resp, err := cosmosClient.GetSQLContainer(ctx, resourceGroupName, accountName, databaseName, containerName, nil)
201+
return GetCosmosDBSQLContainerWithClient(ctx, cosmosClient, resourceGroupName, accountName, databaseName, containerName)
202+
}
203+
204+
// GetCosmosDBSQLContainerWithClient gets a SQL container using the provided SQLResourcesClient.
205+
func GetCosmosDBSQLContainerWithClient(ctx context.Context, client *armcosmos.SQLResourcesClient, resourceGroupName string, accountName string, databaseName string, containerName string) (*armcosmos.SQLContainerGetResults, error) {
206+
resp, err := client.GetSQLContainer(ctx, resourceGroupName, accountName, databaseName, containerName, nil)
197207
if err != nil {
198208
return nil, err
199209
}
200210

201-
// Return container
202211
return &resp.SQLContainerGetResults, nil
203212
}
204213

@@ -236,13 +245,17 @@ func GetCosmosDBSQLDatabaseThroughputContextE(ctx context.Context, subscriptionI
236245
return nil, err
237246
}
238247

239-
// Get the corresponding database throughput config
240-
resp, err := cosmosClient.GetSQLDatabaseThroughput(ctx, resourceGroupName, accountName, databaseName, nil)
248+
return GetCosmosDBSQLDatabaseThroughputWithClient(ctx, cosmosClient, resourceGroupName, accountName, databaseName)
249+
}
250+
251+
// GetCosmosDBSQLDatabaseThroughputWithClient gets a SQL database throughput configuration
252+
// using the provided SQLResourcesClient.
253+
func GetCosmosDBSQLDatabaseThroughputWithClient(ctx context.Context, client *armcosmos.SQLResourcesClient, resourceGroupName string, accountName string, databaseName string) (*armcosmos.ThroughputSettingsGetResults, error) {
254+
resp, err := client.GetSQLDatabaseThroughput(ctx, resourceGroupName, accountName, databaseName, nil)
241255
if err != nil {
242256
return nil, err
243257
}
244258

245-
// Return throughput config
246259
return &resp.ThroughputSettingsGetResults, nil
247260
}
248261

@@ -280,12 +293,16 @@ func GetCosmosDBSQLContainerThroughputContextE(ctx context.Context, subscription
280293
return nil, err
281294
}
282295

283-
// Get the corresponding container throughput config
284-
resp, err := cosmosClient.GetSQLContainerThroughput(ctx, resourceGroupName, accountName, databaseName, containerName, nil)
296+
return GetCosmosDBSQLContainerThroughputWithClient(ctx, cosmosClient, resourceGroupName, accountName, databaseName, containerName)
297+
}
298+
299+
// GetCosmosDBSQLContainerThroughputWithClient gets a SQL container throughput configuration
300+
// using the provided SQLResourcesClient.
301+
func GetCosmosDBSQLContainerThroughputWithClient(ctx context.Context, client *armcosmos.SQLResourcesClient, resourceGroupName string, accountName string, databaseName string, containerName string) (*armcosmos.ThroughputSettingsGetResults, error) {
302+
resp, err := client.GetSQLContainerThroughput(ctx, resourceGroupName, accountName, databaseName, containerName, nil)
285303
if err != nil {
286304
return nil, err
287305
}
288306

289-
// Return throughput config
290307
return &resp.ThroughputSettingsGetResults, nil
291308
}

0 commit comments

Comments
 (0)