diff --git a/cns/NetworkContainerContract.go b/cns/NetworkContainerContract.go index 647e4ba97dd..6f2f11fdfb0 100644 --- a/cns/NetworkContainerContract.go +++ b/cns/NetworkContainerContract.go @@ -675,6 +675,7 @@ type NetworkInterface struct { // PublishNetworkContainerRequest specifies request to publish network container via NMAgent. type PublishNetworkContainerRequest struct { NetworkID string + SubnetName string NetworkContainerID string JoinNetworkURL string CreateNetworkContainerURL string @@ -683,8 +684,8 @@ type PublishNetworkContainerRequest struct { func (p PublishNetworkContainerRequest) String() string { // %q as a verb on a byte slice prints safely escaped text instead of individual bytes - return fmt.Sprintf("{NetworkID:%s NetworkContainerID:%s JoinNetworkURL:%s CreateNetworkContainerURL:%s CreateNetworkContainerRequestBody:%q}", - p.NetworkID, p.NetworkContainerID, p.JoinNetworkURL, p.CreateNetworkContainerURL, p.CreateNetworkContainerRequestBody) + return fmt.Sprintf("{NetworkID:%s SubnetName:%s NetworkContainerID:%s JoinNetworkURL:%s CreateNetworkContainerURL:%s CreateNetworkContainerRequestBody:%q}", + p.NetworkID, p.SubnetName, p.NetworkContainerID, p.JoinNetworkURL, p.CreateNetworkContainerURL, p.CreateNetworkContainerRequestBody) } // NetworkContainerParameters parameters available in network container operations @@ -711,6 +712,7 @@ func (p PublishNetworkContainerResponse) String() string { // UnpublishNetworkContainerRequest specifies request to unpublish network container via NMAgent. type UnpublishNetworkContainerRequest struct { NetworkID string + SubnetName string NetworkContainerID string JoinNetworkURL string DeleteNetworkContainerURL string @@ -718,8 +720,8 @@ type UnpublishNetworkContainerRequest struct { } func (u UnpublishNetworkContainerRequest) String() string { - return fmt.Sprintf("{NetworkID:%s NetworkContainerID:%s JoinNetworkURL:%s DeleteNetworkContainerURL:%s DeleteNetworkContainerRequestBody:%q}", - u.NetworkID, u.NetworkContainerID, u.JoinNetworkURL, u.DeleteNetworkContainerURL, u.DeleteNetworkContainerRequestBody) + return fmt.Sprintf("{NetworkID:%s SubnetName:%s NetworkContainerID:%s JoinNetworkURL:%s DeleteNetworkContainerURL:%s DeleteNetworkContainerRequestBody:%q}", + u.NetworkID, u.SubnetName, u.NetworkContainerID, u.JoinNetworkURL, u.DeleteNetworkContainerURL, u.DeleteNetworkContainerRequestBody) } // UnpublishNetworkContainerResponse specifies the response to unpublish network container request. diff --git a/cns/fakes/wireserverproxyfake.go b/cns/fakes/wireserverproxyfake.go index e6e52e7c54d..01b20f33e45 100644 --- a/cns/fakes/wireserverproxyfake.go +++ b/cns/fakes/wireserverproxyfake.go @@ -10,9 +10,10 @@ import ( ) type WireserverProxyFake struct { - JoinNetworkFunc func(context.Context, string) (*http.Response, error) - PublishNCFunc func(context.Context, cns.NetworkContainerParameters, []byte) (*http.Response, error) - UnpublishNCFunc func(context.Context, cns.NetworkContainerParameters, []byte) (*http.Response, error) + JoinNetworkFunc func(context.Context, string, bool) (*http.Response, error) + JoinSubnetFunc func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) + PublishNCFunc func(context.Context, cns.NetworkContainerParameters, []byte, bool) (*http.Response, error) + UnpublishNCFunc func(context.Context, cns.NetworkContainerParameters, []byte, bool) (*http.Response, error) } const defaultResponseBody = `{"httpStatusCode":"200"}` @@ -25,25 +26,33 @@ func defaultResponse() *http.Response { } } -func (w *WireserverProxyFake) JoinNetwork(ctx context.Context, vnetID string) (*http.Response, error) { +func (w *WireserverProxyFake) JoinNetwork(ctx context.Context, vnetID string, useRNCPublisher bool) (*http.Response, error) { if w.JoinNetworkFunc != nil { - return w.JoinNetworkFunc(ctx, vnetID) + return w.JoinNetworkFunc(ctx, vnetID, useRNCPublisher) } return defaultResponse(), nil } -func (w *WireserverProxyFake) PublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte) (*http.Response, error) { +func (w *WireserverProxyFake) JoinSubnet(ctx context.Context, vnetID, subnetName string, ncParams cns.NetworkContainerParameters) (*http.Response, error) { + if w.JoinSubnetFunc != nil { + return w.JoinSubnetFunc(ctx, vnetID, subnetName, ncParams) + } + + return defaultResponse(), nil +} + +func (w *WireserverProxyFake) PublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte, useRNCPublisher bool) (*http.Response, error) { if w.PublishNCFunc != nil { - return w.PublishNCFunc(ctx, ncParams, payload) + return w.PublishNCFunc(ctx, ncParams, payload, useRNCPublisher) } return defaultResponse(), nil } -func (w *WireserverProxyFake) UnpublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte) (*http.Response, error) { +func (w *WireserverProxyFake) UnpublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte, useRNCPublisher bool) (*http.Response, error) { if w.UnpublishNCFunc != nil { - return w.UnpublishNCFunc(ctx, ncParams, payload) + return w.UnpublishNCFunc(ctx, ncParams, payload, useRNCPublisher) } return defaultResponse(), nil diff --git a/cns/restserver/api.go b/cns/restserver/api.go index ac5ba5994ce..f3b0124d9f6 100644 --- a/cns/restserver/api.go +++ b/cns/restserver/api.go @@ -21,7 +21,6 @@ import ( "github.com/Azure/azure-container-networking/cns/types" "github.com/Azure/azure-container-networking/cns/wireserver" "github.com/Azure/azure-container-networking/common" - "github.com/Azure/azure-container-networking/nmagent" "github.com/pkg/errors" ) @@ -40,6 +39,16 @@ const ( ncURLExpectedMatches = 5 ) +type ncPublishBody struct { + UseRNCPublisher bool `json:"useRNCPublisher"` +} + +type ncUnpublishBody struct { + UseRNCPublisher bool `json:"useRNCPublisher"` + AzID uint `json:"azID"` + AZREnabled bool `json:"azrEnabled"` +} + // This file contains implementation of all HTTP APIs which are exposed to external clients. // TODO: break it even further per module (network, nc, etc) like it is done for ipam @@ -948,7 +957,20 @@ func (service *HTTPRestService) publishNetworkContainer(w http.ResponseWriter, r ctx := r.Context() - joinResp, err := service.wsproxy.JoinNetwork(ctx, req.NetworkID) //nolint:govet // ok to shadow + var publishBody ncPublishBody + var useRNCPublisher bool + + err = json.Unmarshal(req.CreateNetworkContainerRequestBody, &publishBody) + if err != nil { + http.Error(w, fmt.Sprintf("could not unmarshal create network container body: %v", err), http.StatusBadRequest) + return + } + + if publishBody.UseRNCPublisher { + useRNCPublisher = true + } + + joinResp, err := service.wsproxy.JoinNetwork(ctx, req.NetworkID, useRNCPublisher) //nolint:govet // ok to shadow if err != nil { resp := cns.PublishNetworkContainerResponse{ Response: cns.Response{ @@ -982,7 +1004,48 @@ func (service *HTTPRestService) publishNetworkContainer(w http.ResponseWriter, r service.setNetworkStateJoined(req.NetworkID) logger.Printf("[Azure-CNS] joined vnet %s during nc %s publish. wireserver response: %v", req.NetworkID, req.NetworkContainerID, string(joinBytes)) - publishResp, err := service.wsproxy.PublishNC(ctx, ncParams, req.CreateNetworkContainerRequestBody) + if useRNCPublisher { + joinSubnetResp, errSubnetJoin := service.wsproxy.JoinSubnet(ctx, req.NetworkID, req.SubnetName, ncParams) //nolint:govet // ok to shadow + if errSubnetJoin != nil { + resp := cns.PublishNetworkContainerResponse{ + Response: cns.Response{ + ReturnCode: types.SubnetJoinFailed, + Message: fmt.Sprintf("failed to join subnet %s in network %s: %v", req.SubnetName, req.NetworkID, errSubnetJoin), + }, + PublishErrorStr: errSubnetJoin.Error(), + } + respondJSON(w, http.StatusOK, resp) // legacy behavior + logger.Response(service.Name, resp, resp.Response.ReturnCode, errSubnetJoin) //nolint:staticcheck // match existing logger usage in this handler + return + } + + subnetJoinBytes, _ := io.ReadAll(joinSubnetResp.Body) + _ = joinSubnetResp.Body.Close() + + if joinSubnetResp.StatusCode != http.StatusOK { + resp := cns.PublishNetworkContainerResponse{ + Response: cns.Response{ + ReturnCode: types.SubnetJoinFailed, + Message: fmt.Sprintf("failed to join subnet %s in network %s. did not get 200 from wireserver", req.SubnetName, req.NetworkID), + }, + PublishStatusCode: joinSubnetResp.StatusCode, + PublishResponseBody: subnetJoinBytes, + } + respondJSON(w, http.StatusOK, resp) // legacy behavior + logger.Response(service.Name, resp, resp.Response.ReturnCode, nil) //nolint:staticcheck // match existing logger usage in this handler + return + } + + logger.Printf( //nolint:staticcheck // match existing logger usage in this handler + "[Azure-CNS] joined subnet %s in vnet %s during nc %s publish. wireserver response: %v", + req.SubnetName, + req.NetworkID, + req.NetworkContainerID, + string(subnetJoinBytes), + ) + } + + publishResp, err := service.wsproxy.PublishNC(ctx, ncParams, req.CreateNetworkContainerRequestBody, useRNCPublisher) if err != nil { resp := cns.PublishNetworkContainerResponse{ Response: cns.Response{ @@ -1044,8 +1107,9 @@ func (service *HTTPRestService) unpublishNetworkContainer(w http.ResponseWriter, ctx := r.Context() - var unpublishBody nmagent.DeleteContainerRequest + var unpublishBody ncUnpublishBody var azrNC bool + var useRNCPublisher bool err = json.Unmarshal(req.DeleteNetworkContainerRequestBody, &unpublishBody) if err != nil { // If the body contains only `""\n`, it is non-AZR NC @@ -1059,6 +1123,9 @@ func (service *HTTPRestService) unpublishNetworkContainer(w http.ResponseWriter, } else { // If unmarshalling was successful, it is an AZR NC azrNC = true + if unpublishBody.UseRNCPublisher { + useRNCPublisher = true + } } /* For AZR scenarios, if NMAgent is restarted, it loses state and does not know what VNETs to subscribe to. @@ -1066,7 +1133,7 @@ func (service *HTTPRestService) unpublishNetworkContainer(w http.ResponseWriter, nc unpublish calls just like publish nc calls. */ if azrNC || !service.isNetworkJoined(req.NetworkID) { - joinResp, err := service.wsproxy.JoinNetwork(ctx, req.NetworkID) //nolint:govet // ok to shadow + joinResp, err := service.wsproxy.JoinNetwork(ctx, req.NetworkID, useRNCPublisher) //nolint:govet // ok to shadow if err != nil { resp := cns.UnpublishNetworkContainerResponse{ Response: cns.Response{ @@ -1101,7 +1168,49 @@ func (service *HTTPRestService) unpublishNetworkContainer(w http.ResponseWriter, logger.Printf("[Azure-CNS] joined vnet %s during nc %s unpublish. AZREnabled: %t, wireserver response: %v", req.NetworkID, req.NetworkContainerID, unpublishBody.AZREnabled, string(joinBytes)) } - publishResp, err := service.wsproxy.UnpublishNC(ctx, ncParams, req.DeleteNetworkContainerRequestBody) + if useRNCPublisher { + joinSubnetResp, err := service.wsproxy.JoinSubnet(ctx, req.NetworkID, req.SubnetName, ncParams) //nolint:govet // ok to shadow + if err != nil { + resp := cns.UnpublishNetworkContainerResponse{ + Response: cns.Response{ + ReturnCode: types.SubnetJoinFailed, + Message: fmt.Sprintf("failed to join subnet %s in network %s: %v", req.SubnetName, req.NetworkID, err), + }, + UnpublishErrorStr: err.Error(), + } + respondJSON(w, http.StatusOK, resp) // legacy behavior + logger.Response(service.Name, resp, resp.Response.ReturnCode, err) //nolint:staticcheck // match existing logger usage in this handler + return + } + + subnetJoinBytes, _ := io.ReadAll(joinSubnetResp.Body) + _ = joinSubnetResp.Body.Close() + + if joinSubnetResp.StatusCode != http.StatusOK { + resp := cns.UnpublishNetworkContainerResponse{ + Response: cns.Response{ + ReturnCode: types.SubnetJoinFailed, + Message: fmt.Sprintf("failed to join subnet %s in network %s. did not get 200 from wireserver", req.SubnetName, req.NetworkID), + }, + UnpublishStatusCode: joinSubnetResp.StatusCode, + UnpublishResponseBody: subnetJoinBytes, + } + respondJSON(w, http.StatusOK, resp) // legacy behavior + logger.Response(service.Name, resp, resp.Response.ReturnCode, nil) //nolint:staticcheck // match existing logger usage in this handler + return + } + + logger.Printf( //nolint:staticcheck // match existing logger usage in this handler + "[Azure-CNS] joined subnet %s in vnet %s during nc %s unpublish. AZREnabled: %t, wireserver response: %v", + req.SubnetName, + req.NetworkID, + req.NetworkContainerID, + unpublishBody.AZREnabled, + string(subnetJoinBytes), + ) + } + + unpublishResp, err := service.wsproxy.UnpublishNC(ctx, ncParams, req.DeleteNetworkContainerRequestBody, useRNCPublisher) if err != nil { resp := cns.UnpublishNetworkContainerResponse{ Response: cns.Response{ @@ -1115,15 +1224,15 @@ func (service *HTTPRestService) unpublishNetworkContainer(w http.ResponseWriter, return } - publishBytes, _ := io.ReadAll(publishResp.Body) - _ = publishResp.Body.Close() + unpublishBytes, _ := io.ReadAll(unpublishResp.Body) + _ = unpublishResp.Body.Close() resp := cns.UnpublishNetworkContainerResponse{ - UnpublishStatusCode: publishResp.StatusCode, - UnpublishResponseBody: publishBytes, + UnpublishStatusCode: unpublishResp.StatusCode, + UnpublishResponseBody: unpublishBytes, } - if publishResp.StatusCode != http.StatusOK { + if unpublishResp.StatusCode != http.StatusOK { resp.Response = cns.Response{ ReturnCode: types.NetworkContainerUnpublishFailed, Message: fmt.Sprintf("failed to unpublish nc %s. did not get 200 from wireserver", req.NetworkContainerID), diff --git a/cns/restserver/api_test.go b/cns/restserver/api_test.go index f68fb7facb3..8c217d74532 100644 --- a/cns/restserver/api_test.go +++ b/cns/restserver/api_test.go @@ -315,7 +315,6 @@ func TestSetOrchestratorType_NCsPresent(t *testing.T) { }, } for _, tt := range tests { - tt := tt t.Run(tt.name, func(t *testing.T) { var resp cns.Response // Since this is global, we have to replace the state @@ -645,11 +644,11 @@ func TestGetNetworkContainerVersionStatus(t *testing.T) { return rsp, errors.New("boom") //nolint:goerr113 // it's just a test } - wsproxy.JoinNetworkFunc = func(ctx context.Context, s string) (*http.Response, error) { + wsproxy.JoinNetworkFunc = func(_ context.Context, _ string, _ bool) (*http.Response, error) { return nil, errors.New("boom") //nolint:goerr113 // it's just a test } - wsproxy.PublishNCFunc = func(ctx context.Context, parameters cns.NetworkContainerParameters, i []byte) (*http.Response, error) { + wsproxy.PublishNCFunc = func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, _ bool) (*http.Response, error) { return nil, errors.New("boom") //nolint:goerr113 // it's just a test } @@ -764,7 +763,7 @@ func TestPublishNC_NMAgentApplicationErrors(t *testing.T) { wireserverBody := `{"httpStatusCode":"401"}` wsproxy := fakes.WireserverProxyFake{ - PublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte) (*http.Response, error) { + PublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, _ bool) (*http.Response, error) { return &http.Response{ StatusCode: http.StatusOK, Body: io.NopCloser(bytes.NewBufferString(wireserverBody)), @@ -839,6 +838,429 @@ func TestPublishNC_NMAgentApplicationErrors(t *testing.T) { } } +func TestPublishNCAllowsEmptyRequestBody(t *testing.T) { + var ( + joinUsedRNCPublisher bool + publishUsedRNCPublisher bool + ) + + wsproxy := fakes.WireserverProxyFake{ + JoinNetworkFunc: func(_ context.Context, _ string, useRNCPublisher bool) (*http.Response, error) { + joinUsedRNCPublisher = useRNCPublisher + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + PublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, useRNCPublisher bool) (*http.Response, error) { + publishUsedRNCPublisher = useRNCPublisher + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + defer cleanup() + + joinNetworkURL := "http://" + nmagentEndpoint + "/dummyVnetURL" + createNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1" + publishNCRequest := &cns.PublishNetworkContainerRequest{ + NetworkID: "foo", + NetworkContainerID: "bar", + JoinNetworkURL: joinNetworkURL, + CreateNetworkContainerURL: createNetworkContainerURL, + CreateNetworkContainerRequestBody: []byte("{}"), + } + + body := encodeRequestBody(t, publishNCRequest) + + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.PublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code) + + var resp cns.PublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.Success, resp.Response.ReturnCode) + require.False(t, joinUsedRNCPublisher) + require.False(t, publishUsedRNCPublisher) +} + +func TestPublishNCRequestBodyParsingMatrix(t *testing.T) { + const ( + networkID = "vnet-publish-body-matrix" + subnetName = "subnet-publish-body-matrix" + networkContainerID = "nc-publish-body-matrix" + ) + + tests := []struct { + name string + body []byte + wantHTTPStatus int + wantReturnCode types.ResponseCode + wantUseRNCPublisher bool + wantJoinSubnetCalls int + wantPublishCalls int + }{ + { + name: "empty object succeeds without rnc", + body: []byte(`{}`), + wantHTTPStatus: http.StatusOK, + wantReturnCode: types.Success, + wantUseRNCPublisher: false, + wantJoinSubnetCalls: 0, + wantPublishCalls: 1, + }, + { + name: "unknown fields succeed without rnc", + body: []byte(`{"someField":"someValue"}`), + wantHTTPStatus: http.StatusOK, + wantReturnCode: types.Success, + wantUseRNCPublisher: false, + wantJoinSubnetCalls: 0, + wantPublishCalls: 1, + }, + { + name: "invalid version type no longer blocks request", + body: []byte(`{"version":"bad"}`), + wantHTTPStatus: http.StatusOK, + wantReturnCode: types.Success, + wantUseRNCPublisher: false, + wantJoinSubnetCalls: 0, + wantPublishCalls: 1, + }, + { + name: "rnc body triggers subnet join", + body: []byte(`{"useRNCPublisher":true}`), + wantHTTPStatus: http.StatusOK, + wantReturnCode: types.Success, + wantUseRNCPublisher: true, + wantJoinSubnetCalls: 1, + wantPublishCalls: 1, + }, + { + name: "invalid json returns bad request", + body: []byte("invalid\n"), + wantHTTPStatus: http.StatusBadRequest, + wantJoinSubnetCalls: 0, + wantPublishCalls: 0, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var ( + joinSubnetCalls int + publishCalls int + capturedPublishRNC bool + ) + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(_ context.Context, vnetID, gotSubnetName string, _ cns.NetworkContainerParameters) (*http.Response, error) { + joinSubnetCalls++ + require.Equal(t, networkID, vnetID) + require.Equal(t, subnetName, gotSubnetName) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + PublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, useRNCPublisher bool) (*http.Response, error) { + publishCalls++ + capturedPublishRNC = useRNCPublisher + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + createNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1" + publishNCRequest := &cns.PublishNetworkContainerRequest{ + NetworkID: networkID, + SubnetName: subnetName, + NetworkContainerID: networkContainerID, + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + CreateNetworkContainerURL: createNetworkContainerURL, + CreateNetworkContainerRequestBody: tt.body, + } + + body := encodeRequestBody(t, publishNCRequest) + + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.PublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + require.Equal(t, tt.wantHTTPStatus, w.Code) + + if tt.wantHTTPStatus == http.StatusOK { + var resp cns.PublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, tt.wantReturnCode, resp.Response.ReturnCode) + require.Equal(t, tt.wantUseRNCPublisher, capturedPublishRNC) + } + + require.Equal(t, tt.wantJoinSubnetCalls, joinSubnetCalls) + require.Equal(t, tt.wantPublishCalls, publishCalls) + }) + } +} + +func TestPublishNCWithRNCPublisherJoinsSubnetEveryTime(t *testing.T) { + const ( + networkID = "vnet-rnc-publish" + subnetName = "subnet-rnc-publish" + networkContainerID = "nc-rnc-publish" + ) + + var ( + joinSubnetCalls int + publishCalls int + ) + + wsproxy := fakes.WireserverProxyFake{ + JoinNetworkFunc: func(_ context.Context, vnetID string, useRNCPublisher bool) (*http.Response, error) { + require.Equal(t, networkID, vnetID) + require.True(t, useRNCPublisher) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + JoinSubnetFunc: func(_ context.Context, vnetID, gotSubnetName string, _ cns.NetworkContainerParameters) (*http.Response, error) { + joinSubnetCalls++ + require.Equal(t, networkID, vnetID) + require.Equal(t, subnetName, gotSubnetName) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + PublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, useRNCPublisher bool) (*http.Response, error) { + publishCalls++ + require.True(t, useRNCPublisher) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + createNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1" + publishNCRequest := &cns.PublishNetworkContainerRequest{ + NetworkID: networkID, + SubnetName: subnetName, + NetworkContainerID: networkContainerID, + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + CreateNetworkContainerURL: createNetworkContainerURL, + CreateNetworkContainerRequestBody: []byte(`{"useRNCPublisher":true}`), + } + + for i := 0; i < 2; i++ { + body := encodeRequestBody(t, publishNCRequest) + + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.PublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.PublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.Success, resp.Response.ReturnCode) + } + + require.Equal(t, 2, joinSubnetCalls) + require.Equal(t, 2, publishCalls) +} + +func TestPublishNCWithRNCPublisherSubnetJoinFailure(t *testing.T) { + const ( + networkID = "vnet-rnc-subnet-failure" + subnetName = "subnet-rnc-subnet-failure" + networkContainerID = "nc-rnc-subnet-failure" + ) + + var publishCalls int + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) { + return nil, errors.New("subnet join failed") + }, + PublishNCFunc: func(context.Context, cns.NetworkContainerParameters, []byte, bool) (*http.Response, error) { + publishCalls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + createNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1" + publishNCRequest := &cns.PublishNetworkContainerRequest{ + NetworkID: networkID, + SubnetName: subnetName, + NetworkContainerID: networkContainerID, + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + CreateNetworkContainerURL: createNetworkContainerURL, + CreateNetworkContainerRequestBody: []byte(`{"useRNCPublisher":true}`), + } + + body := encodeRequestBody(t, publishNCRequest) + + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.PublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.PublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.SubnetJoinFailed, resp.Response.ReturnCode) + require.Contains(t, resp.PublishErrorStr, "subnet join failed") + require.Zero(t, publishCalls) +} + +func TestPublishNCWithRNCPublisherSubnetJoinNon200(t *testing.T) { + const ( + networkID = "vnet-rnc-subnet-status-failure" + subnetName = "subnet-rnc-subnet-status-failure" + networkContainerID = "nc-rnc-subnet-status-failure" + ) + + var publishCalls int + const subnetJoinStatusCode = http.StatusInternalServerError + subnetJoinBody := []byte(`{"httpStatusCode":"500"}`) + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) { + return &http.Response{ + StatusCode: subnetJoinStatusCode, + Body: io.NopCloser(bytes.NewBuffer(subnetJoinBody)), + }, nil + }, + PublishNCFunc: func(context.Context, cns.NetworkContainerParameters, []byte, bool) (*http.Response, error) { + publishCalls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + createNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1" + publishNCRequest := &cns.PublishNetworkContainerRequest{ + NetworkID: networkID, + SubnetName: subnetName, + NetworkContainerID: networkContainerID, + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + CreateNetworkContainerURL: createNetworkContainerURL, + CreateNetworkContainerRequestBody: []byte(`{"useRNCPublisher":true}`), + } + + body := encodeRequestBody(t, publishNCRequest) + + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.PublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.PublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.SubnetJoinFailed, resp.Response.ReturnCode) + require.Equal(t, subnetJoinStatusCode, resp.PublishStatusCode) + require.Equal(t, subnetJoinBody, resp.PublishResponseBody) + require.Zero(t, publishCalls) +} + +func TestPublishNCWithRNCPublisherDisabledSkipsSubnetJoin(t *testing.T) { + var ( + joinSubnetCalls int + publishCalls int + ) + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) { + joinSubnetCalls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + PublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, useRNCPublisher bool) (*http.Response, error) { + publishCalls++ + require.False(t, useRNCPublisher) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + createNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1" + publishNCRequest := &cns.PublishNetworkContainerRequest{ + NetworkID: "vnet-rnc-disabled-publish", + SubnetName: "subnet-rnc-disabled-publish", + NetworkContainerID: "nc-rnc-disabled-publish", + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + CreateNetworkContainerURL: createNetworkContainerURL, + CreateNetworkContainerRequestBody: []byte(`{"useRNCPublisher":false}`), + } + + body := encodeRequestBody(t, publishNCRequest) + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.PublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.PublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.Success, resp.Response.ReturnCode) + require.Zero(t, joinSubnetCalls) + require.Equal(t, 1, publishCalls) +} + func publishNCViaCNS( networkID, networkContainerID, @@ -981,10 +1403,21 @@ func TestUnpublishViaCNSRequestBody(t *testing.T) { body: []byte(`{"azID":1,"azrEnabled":true}`), requireError: false, }, + { + name: "Delete NC with invalid AZR azID type", + ncID: "ncID4", + body: []byte(`{"azID":"bad","azrEnabled":true}`), + requireError: true, + }, + { + name: "Delete NC with empty object body", + ncID: "ncID5", + body: []byte(`{}`), + requireError: false, + }, } for _, tt := range tests { - tt := tt t.Run(tt.name, func(t *testing.T) { errPublish := publishNCViaCNS(vnet, tt.ncID, createNetworkContainerURL) require.NoError(t, errPublish) @@ -1001,7 +1434,7 @@ func TestUnpublishViaCNSRequestBody(t *testing.T) { func TestUnpublishNCViaCNS401(t *testing.T) { wsproxy := fakes.WireserverProxyFake{ - UnpublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, i []byte) (*http.Response, error) { + UnpublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, _ bool) (*http.Response, error) { return &http.Response{ StatusCode: http.StatusOK, Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"401"}`)), @@ -1078,6 +1511,222 @@ func TestUnpublishNCViaCNS401(t *testing.T) { } } +func TestUnpublishNCWithRNCPublisherJoinsSubnet(t *testing.T) { + const ( + networkID = "vnet-rnc-unpublish" + subnetName = "subnet-rnc-unpublish" + networkContainerID = "nc-rnc-unpublish" + ) + + var ( + joinSubnetCalls int + unpublishCalls int + ) + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) { + joinSubnetCalls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + UnpublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, useRNCPublisher bool) (*http.Response, error) { + unpublishCalls++ + require.True(t, useRNCPublisher) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + deleteNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1/method/DELETE" + unpublishNCRequest := &cns.UnpublishNetworkContainerRequest{ + NetworkID: networkID, + SubnetName: subnetName, + NetworkContainerID: networkContainerID, + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + DeleteNetworkContainerURL: deleteNetworkContainerURL, + DeleteNetworkContainerRequestBody: []byte(`{"azrEnabled":true,"useRNCPublisher":true}`), + } + + for i := 0; i < 2; i++ { + body := encodeRequestBody(t, unpublishNCRequest) + + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.UnpublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.UnpublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.Success, resp.Response.ReturnCode) + } + + require.Equal(t, 2, joinSubnetCalls) + require.Equal(t, 2, unpublishCalls) +} + +func TestUnpublishNCWithRNCPublisherDisabledSkipsSubnetJoin(t *testing.T) { + var ( + joinSubnetCalls int + unpublishCalls int + ) + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) { + joinSubnetCalls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + UnpublishNCFunc: func(_ context.Context, _ cns.NetworkContainerParameters, _ []byte, useRNCPublisher bool) (*http.Response, error) { + unpublishCalls++ + require.False(t, useRNCPublisher) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + deleteNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1/method/DELETE" + unpublishNCRequest := &cns.UnpublishNetworkContainerRequest{ + NetworkID: "vnet-rnc-disabled-unpublish", + SubnetName: "subnet-rnc-disabled-unpublish", + NetworkContainerID: "nc-rnc-disabled-unpublish", + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + DeleteNetworkContainerURL: deleteNetworkContainerURL, + DeleteNetworkContainerRequestBody: []byte(`{"azID":1,"azrEnabled":true,"useRNCPublisher":false}`), + } + + body := encodeRequestBody(t, unpublishNCRequest) + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.UnpublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.UnpublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.Success, resp.Response.ReturnCode) + require.Zero(t, joinSubnetCalls) + require.Equal(t, 1, unpublishCalls) +} + +func TestUnpublishNCWithRNCPublisherSubnetJoinFailure(t *testing.T) { + var unpublishCalls int + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) { + return nil, errors.New("subnet join failed") + }, + UnpublishNCFunc: func(context.Context, cns.NetworkContainerParameters, []byte, bool) (*http.Response, error) { + unpublishCalls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + deleteNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1/method/DELETE" + unpublishNCRequest := &cns.UnpublishNetworkContainerRequest{ + NetworkID: "vnet-rnc-unpublish-subnet-failure", + SubnetName: "subnet-rnc-unpublish-subnet-failure", + NetworkContainerID: "nc-rnc-unpublish-subnet-failure", + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + DeleteNetworkContainerURL: deleteNetworkContainerURL, + DeleteNetworkContainerRequestBody: []byte(`{"azID":1,"azrEnabled":true,"useRNCPublisher":true}`), + } + + body := encodeRequestBody(t, unpublishNCRequest) + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.UnpublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.UnpublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.SubnetJoinFailed, resp.Response.ReturnCode) + require.Contains(t, resp.UnpublishErrorStr, "subnet join failed") + require.Zero(t, unpublishCalls) +} + +func TestUnpublishNCWithRNCPublisherSubnetJoinNon200(t *testing.T) { + var unpublishCalls int + const subnetJoinStatusCode = http.StatusInternalServerError + subnetJoinBody := []byte(`{"httpStatusCode":"500"}`) + + wsproxy := fakes.WireserverProxyFake{ + JoinSubnetFunc: func(context.Context, string, string, cns.NetworkContainerParameters) (*http.Response, error) { + return &http.Response{ + StatusCode: subnetJoinStatusCode, + Body: io.NopCloser(bytes.NewBuffer(subnetJoinBody)), + }, nil + }, + UnpublishNCFunc: func(context.Context, cns.NetworkContainerParameters, []byte, bool) (*http.Response, error) { + unpublishCalls++ + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{"httpStatusCode":"200"}`)), + }, nil + }, + } + + cleanup := setWireserverProxy(svc, &wsproxy) + t.Cleanup(cleanup) + + deleteNetworkContainerURL := "http://" + nmagentEndpoint + + "/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/dummyIntf/networkContainers/dummyNCURL/authenticationToken/dummyT/api-version/1/method/DELETE" + unpublishNCRequest := &cns.UnpublishNetworkContainerRequest{ + NetworkID: "vnet-rnc-unpublish-subnet-status-failure", + SubnetName: "subnet-rnc-unpublish-subnet-status-failure", + NetworkContainerID: "nc-rnc-unpublish-subnet-status-failure", + JoinNetworkURL: "http://" + nmagentEndpoint + "/dummyVnetURL", + DeleteNetworkContainerURL: deleteNetworkContainerURL, + DeleteNetworkContainerRequestBody: []byte(`{"azID":1,"azrEnabled":true,"useRNCPublisher":true}`), + } + + body := encodeRequestBody(t, unpublishNCRequest) + //nolint:noctx // not needed in test + req, err := http.NewRequest(http.MethodPost, cns.UnpublishNetworkContainer, &body) + require.NoError(t, err) + + w := httptest.NewRecorder() + mux.ServeHTTP(w, req) + + var resp cns.UnpublishNetworkContainerResponse + err = decodeResponse(w, &resp) + require.NoError(t, err) + require.Equal(t, types.SubnetJoinFailed, resp.Response.ReturnCode) + require.Equal(t, subnetJoinStatusCode, resp.UnpublishStatusCode) + require.Equal(t, subnetJoinBody, resp.UnpublishResponseBody) + require.Zero(t, unpublishCalls) +} + func unpublishNCViaCNS(networkID, networkContainerID, deleteNetworkContainerURL string, bodyBytes []byte) error { joinNetworkURL := "http://" + nmagentEndpoint + "/dummyVnetURL" @@ -1638,6 +2287,16 @@ func decodeResponse(w *httptest.ResponseRecorder, response interface{}) error { return json.NewDecoder(w.Body).Decode(&response) } +func encodeRequestBody(t *testing.T, request any) bytes.Buffer { + t.Helper() + + var body bytes.Buffer + err := json.NewEncoder(&body).Encode(request) + require.NoError(t, err) + + return body +} + func setEnv(t *testing.T) *httptest.ResponseRecorder { envRequest := cns.SetEnvironmentRequest{Location: "Azure", NetworkType: "Underlay"} envRequestJSON := new(bytes.Buffer) diff --git a/cns/restserver/restserver.go b/cns/restserver/restserver.go index ec0f3827482..93b15582db4 100644 --- a/cns/restserver/restserver.go +++ b/cns/restserver/restserver.go @@ -47,9 +47,10 @@ type nmagentClient interface { } type wireserverProxy interface { - JoinNetwork(ctx context.Context, vnetID string) (*http.Response, error) - PublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte) (*http.Response, error) - UnpublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte) (*http.Response, error) + JoinNetwork(ctx context.Context, vnetID string, useRNCPublisher bool) (*http.Response, error) + JoinSubnet(ctx context.Context, vnetID, subnetName string, ncParams cns.NetworkContainerParameters) (*http.Response, error) + PublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte, useRNCPublisher bool) (*http.Response, error) + UnpublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte, useRNCPublisher bool) (*http.Response, error) } type imdsClient interface { diff --git a/cns/restserver/util_test.go b/cns/restserver/util_test.go index c3c7ad76eb9..5faaff3a76e 100644 --- a/cns/restserver/util_test.go +++ b/cns/restserver/util_test.go @@ -185,6 +185,37 @@ func TestRestoreState(t *testing.T) { } } +func TestRestoreStateIgnoresLegacyJoinedSubnetsField(t *testing.T) { + const ( + underlayNetworkType = "Underlay" + vnetID = "vnet1" + ) + + mainStore := store.NewMockStore("") + require.NoError(t, mainStore.Write(storeKey, map[string]any{ + "NetworkType": underlayNetworkType, + "joinedNetworks": map[string]struct{}{ + vnetID: {}, + }, + "joinedSubnets": map[string]struct{}{ + "vnet1_subnet1": {}, + }, + })) + + svc := HTTPRestService{ + Service: &cns.Service{ + Service: &common.Service{Options: map[string]any{}}, + }, + store: mainStore, + state: &httpRestServiceState{}, + } + + svc.restoreState() + + require.Equal(t, underlayNetworkType, svc.state.NetworkType) + require.Nil(t, svc.state.joinedNetworks) +} + // test to check if nc can be deleted from ncList for Delete() method func TestDeleteNCs(t *testing.T) { var ncs ncList diff --git a/cns/types/codes.go b/cns/types/codes.go index 9492e92a596..98389b70a7e 100644 --- a/cns/types/codes.go +++ b/cns/types/codes.go @@ -47,6 +47,7 @@ const ( ConnectionError ResponseCode = 45 UnexpectedError ResponseCode = 99 NmAgentNCVersionListError ResponseCode = 100 + SubnetJoinFailed ResponseCode = 101 ) // nolint:gocyclo @@ -128,6 +129,8 @@ func (c ResponseCode) String() string { return "StatusUnauthorized" case FailedToAllocateBackendConfig: return "FailedToAllocateBackendConfig" + case SubnetJoinFailed: + return "SubnetJoinFailed" default: return "UnknownError" } diff --git a/cns/wireserver/proxy.go b/cns/wireserver/proxy.go index d180ff9cccc..60fccf37945 100644 --- a/cns/wireserver/proxy.go +++ b/cns/wireserver/proxy.go @@ -12,6 +12,7 @@ import ( const ( joinNetworkURLFmt = `http://%s/machine/plugins/?comp=nmagent&type=NetworkManagement/joinedVirtualNetworks/%s/api-version/1` + joinSubnetURLFmt = `http://%s/machine/plugins/?comp=nmagent&type=NetworkManagement/joinedVirtualNetworks/%s/joinedSubnets/%s/authenticationToken/%s/api-version/1&useLegacyChannel=false` publishNCURLFmt = `http://%s/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/%s/networkContainers/%s/authenticationToken/%s/api-version/1` unpublishNCURLFmt = `http://%s/machine/plugins/?comp=nmagent&type=NetworkManagement/interfaces/%s/networkContainers/%s/authenticationToken/%s/api-version/1/method/DELETE` ) @@ -21,8 +22,14 @@ type Proxy struct { HTTPClient do } -func (p *Proxy) JoinNetwork(ctx context.Context, vnetID string) (*http.Response, error) { - reqURL := fmt.Sprintf(joinNetworkURLFmt, p.Host, vnetID) +func (p *Proxy) JoinNetwork(ctx context.Context, vnetID string, useRNCPublisher bool) (*http.Response, error) { + var joinNetworkURLFormat string + if useRNCPublisher { + joinNetworkURLFormat = joinNetworkURLFmt + "&useLegacyChannel=false" + } else { + joinNetworkURLFormat = joinNetworkURLFmt + } + reqURL := fmt.Sprintf(joinNetworkURLFormat, p.Host, vnetID) req, err := http.NewRequestWithContext(ctx, http.MethodPost, reqURL, bytes.NewBufferString(`""`)) if err != nil { @@ -39,8 +46,32 @@ func (p *Proxy) JoinNetwork(ctx context.Context, vnetID string) (*http.Response, return resp, nil } -func (p *Proxy) PublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte) (*http.Response, error) { - reqURL := fmt.Sprintf(publishNCURLFmt, p.Host, ncParams.AssociatedInterfaceID, ncParams.NCID, ncParams.AuthToken) +func (p *Proxy) JoinSubnet(ctx context.Context, vnetID, subnetName string, ncParams cns.NetworkContainerParameters) (*http.Response, error) { + reqURL := fmt.Sprintf(joinSubnetURLFmt, p.Host, vnetID, subnetName, ncParams.AuthToken) + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, reqURL, bytes.NewBufferString(`""`)) + if err != nil { + return nil, errors.Wrap(err, "wireserver proxy: join subnet: could not build http request") + } + + req.Header.Set("Content-Type", "application/json") + + resp, err := p.HTTPClient.Do(req) + if err != nil { + return nil, errors.Wrap(err, "wireserver proxy: join subnet: could not perform http request") + } + + return resp, nil +} + +func (p *Proxy) PublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte, useRNCPublisher bool) (*http.Response, error) { + var publishNCURLFormat string + if useRNCPublisher { + publishNCURLFormat = publishNCURLFmt + "&useLegacyChannel=false" + } else { + publishNCURLFormat = publishNCURLFmt + } + reqURL := fmt.Sprintf(publishNCURLFormat, p.Host, ncParams.AssociatedInterfaceID, ncParams.NCID, ncParams.AuthToken) req, err := http.NewRequestWithContext(ctx, http.MethodPost, reqURL, bytes.NewBuffer(payload)) if err != nil { @@ -57,8 +88,14 @@ func (p *Proxy) PublishNC(ctx context.Context, ncParams cns.NetworkContainerPara return resp, nil } -func (p *Proxy) UnpublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte) (*http.Response, error) { - reqURL := fmt.Sprintf(unpublishNCURLFmt, p.Host, ncParams.AssociatedInterfaceID, ncParams.NCID, ncParams.AuthToken) +func (p *Proxy) UnpublishNC(ctx context.Context, ncParams cns.NetworkContainerParameters, payload []byte, useRNCPublisher bool) (*http.Response, error) { + var unpublishNCURLFormat string + if useRNCPublisher { + unpublishNCURLFormat = unpublishNCURLFmt + "&useLegacyChannel=false" + } else { + unpublishNCURLFormat = unpublishNCURLFmt + } + reqURL := fmt.Sprintf(unpublishNCURLFormat, p.Host, ncParams.AssociatedInterfaceID, ncParams.NCID, ncParams.AuthToken) // a POST to wireserver must contain a body. For legacy purposes, // an empty json string (two quote characters) should be sent by default. diff --git a/cns/wireserver/proxy_test.go b/cns/wireserver/proxy_test.go new file mode 100644 index 00000000000..b94dac051d5 --- /dev/null +++ b/cns/wireserver/proxy_test.go @@ -0,0 +1,157 @@ +package wireserver + +import ( + "bytes" + "context" + "io" + "net/http" + "net/url" + "testing" + + "github.com/Azure/azure-container-networking/cns" + "github.com/stretchr/testify/require" +) + +type testDo struct { + do func(*http.Request) (*http.Response, error) +} + +func (t *testDo) Do(req *http.Request) (*http.Response, error) { + return t.do(req) +} + +const ( + useLegacyChannelFalse = "false" + interfaceID = "iface-1" + networkContainerID = "nc-1" + authToken = "token-1" +) + +func TestProxyRNCPublisherQueryParam(t *testing.T) { + tests := []struct { + name string + call func(*Proxy) (*http.Response, error) + expectedFlag string + expectLegacySwitch bool + expectedTypePath string + }{ + { + name: "JoinNetwork adds useLegacyChannel=false for RNC", + call: func(p *Proxy) (*http.Response, error) { + return p.JoinNetwork(context.Background(), "vnet-1", true) + }, + expectedFlag: useLegacyChannelFalse, + expectLegacySwitch: true, + expectedTypePath: "NetworkManagement/joinedVirtualNetworks/vnet-1/api-version/1", + }, + { + name: "PublishNC adds useLegacyChannel=false for RNC", + call: func(p *Proxy) (*http.Response, error) { + return p.PublishNC(context.Background(), cns.NetworkContainerParameters{ + AssociatedInterfaceID: interfaceID, + NCID: networkContainerID, + AuthToken: authToken, + }, []byte(`{}`), true) + }, + expectedFlag: useLegacyChannelFalse, + expectLegacySwitch: true, + expectedTypePath: "NetworkManagement/interfaces/iface-1/networkContainers/nc-1/authenticationToken/token-1/api-version/1", + }, + { + name: "UnpublishNC adds useLegacyChannel=false for RNC", + call: func(p *Proxy) (*http.Response, error) { + return p.UnpublishNC(context.Background(), cns.NetworkContainerParameters{ + AssociatedInterfaceID: interfaceID, + NCID: networkContainerID, + AuthToken: authToken, + }, []byte(`{}`), true) + }, + expectedFlag: useLegacyChannelFalse, + expectLegacySwitch: true, + expectedTypePath: "NetworkManagement/interfaces/iface-1/networkContainers/nc-1/authenticationToken/token-1/api-version/1/method/DELETE", + }, + { + name: "JoinSubnet includes useLegacyChannel=false", + call: func(p *Proxy) (*http.Response, error) { + return p.JoinSubnet(context.Background(), "vnet-1", "subnet-1", cns.NetworkContainerParameters{ + AuthToken: authToken, + }) + }, + expectedFlag: useLegacyChannelFalse, + expectLegacySwitch: true, + expectedTypePath: "NetworkManagement/joinedVirtualNetworks/vnet-1/joinedSubnets/subnet-1/authenticationToken/token-1/api-version/1", + }, + { + name: "JoinNetwork does not include useLegacyChannel when RNC disabled", + call: func(p *Proxy) (*http.Response, error) { + return p.JoinNetwork(context.Background(), "vnet-1", false) + }, + expectLegacySwitch: false, + expectedTypePath: "NetworkManagement/joinedVirtualNetworks/vnet-1/api-version/1", + }, + { + name: "PublishNC does not include useLegacyChannel when RNC disabled", + call: func(p *Proxy) (*http.Response, error) { + return p.PublishNC(context.Background(), cns.NetworkContainerParameters{ + AssociatedInterfaceID: interfaceID, + NCID: networkContainerID, + AuthToken: authToken, + }, []byte(`{}`), false) + }, + expectLegacySwitch: false, + expectedTypePath: "NetworkManagement/interfaces/iface-1/networkContainers/nc-1/authenticationToken/token-1/api-version/1", + }, + { + name: "UnpublishNC does not include useLegacyChannel when RNC disabled", + call: func(p *Proxy) (*http.Response, error) { + return p.UnpublishNC(context.Background(), cns.NetworkContainerParameters{ + AssociatedInterfaceID: interfaceID, + NCID: networkContainerID, + AuthToken: authToken, + }, []byte(`{}`), false) + }, + expectLegacySwitch: false, + expectedTypePath: "NetworkManagement/interfaces/iface-1/networkContainers/nc-1/authenticationToken/token-1/api-version/1/method/DELETE", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var reqURL *url.URL + + p := &Proxy{ + Host: "127.0.0.1:9001", + HTTPClient: &testDo{ + do: func(req *http.Request) (*http.Response, error) { + reqURL = req.URL + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewBufferString(`{}`)), + }, nil + }, + }, + } + + resp, err := tt.call(p) + require.NoError(t, err) + t.Cleanup(func() { + if resp != nil && resp.Body != nil { + _ = resp.Body.Close() + } + }) + require.NotNil(t, reqURL) + + q := reqURL.Query() + typeVal := q.Get("type") + require.NotContains(t, typeVal, "?useLegacyChannel=false") + require.Equal(t, tt.expectedTypePath, typeVal) + + if tt.expectLegacySwitch { + require.Equal(t, tt.expectedFlag, q.Get("useLegacyChannel")) + } else { + _, exists := q["useLegacyChannel"] + require.False(t, exists) + } + }) + } +} diff --git a/nmagent/requests.go b/nmagent/requests.go index 6a173080fa7..7482cf065fd 100644 --- a/nmagent/requests.go +++ b/nmagent/requests.go @@ -74,6 +74,9 @@ type PutNetworkContainerRequest struct { // AZREnabled denotes whether AZR is enabled for network container or not AZREnabled bool + + // UseRNCPublisher denotes whether NC should be published via RNC (used for auth) + UseRNCPublisher bool } type internalNC struct { @@ -83,27 +86,29 @@ type internalNC struct { Version string `json:"version"` // The rest of these are copied verbatim from the above struct and should be kept in sync. - VNetID string `json:"virtualNetworkId"` - SubnetName string `json:"subnetName"` - IPv4Addrs []string `json:"ipV4Addresses"` - Policies []Policy `json:"policies"` - VlanID int `json:"vlanId"` - GREKey uint16 `json:"greKey"` - AzID uint `json:"azID"` - AZREnabled bool `json:"azrEnabled"` + VNetID string `json:"virtualNetworkId"` + SubnetName string `json:"subnetName"` + IPv4Addrs []string `json:"ipV4Addresses"` + Policies []Policy `json:"policies"` + VlanID int `json:"vlanId"` + GREKey uint16 `json:"greKey"` + AzID uint `json:"azID"` + AZREnabled bool `json:"azrEnabled"` + UseRNCPublisher bool `json:"useRNCPublisher"` } func (p *PutNetworkContainerRequest) MarshalJSON() ([]byte, error) { pBody := internalNC{ - Version: strconv.Itoa(int(p.Version)), - VNetID: p.VNetID, - SubnetName: p.SubnetName, - IPv4Addrs: p.IPv4Addrs, - Policies: p.Policies, - VlanID: p.VlanID, - GREKey: p.GREKey, - AzID: p.AzID, - AZREnabled: p.AZREnabled, + Version: strconv.FormatUint(p.Version, 10), + VNetID: p.VNetID, + SubnetName: p.SubnetName, + IPv4Addrs: p.IPv4Addrs, + Policies: p.Policies, + VlanID: p.VlanID, + GREKey: p.GREKey, + AzID: p.AzID, + AZREnabled: p.AZREnabled, + UseRNCPublisher: p.UseRNCPublisher, } body, err := json.Marshal(pBody) @@ -135,6 +140,7 @@ func (p *PutNetworkContainerRequest) UnmarshalJSON(in []byte) error { p.GREKey = req.GREKey p.AzID = req.AzID p.AZREnabled = req.AZREnabled + p.UseRNCPublisher = req.UseRNCPublisher return nil } @@ -317,9 +323,10 @@ var _ Request = DeleteContainerRequest{} // DeleteContainerRequest represents all information necessary to request that // NMAgent delete a particular network container type DeleteContainerRequest struct { - NCID string `json:"-"` // the Network Container ID - AzID uint `json:"azID"` // home AZ of the Network Container - AZREnabled bool `json:"azrEnabled"` // whether AZR is enabled or not + NCID string `json:"-"` // the Network Container ID + AzID uint `json:"azID"` // home AZ of the Network Container + AZREnabled bool `json:"azrEnabled"` // whether AZR is enabled or not + UseRNCPublisher bool `json:"useRNCPublisher"` // PrimaryAddress is the primary customer address of the interface in the // management VNET