@@ -0,0 +1,196 @@
|
||||
package oci
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/oracle/oci-go-sdk/v65/common"
|
||||
"github.com/oracle/oci-go-sdk/v65/identitydomains"
|
||||
)
|
||||
|
||||
type federationCreateStub struct {
|
||||
created identitydomains.IdentityProvider
|
||||
createErr error
|
||||
patchErr error
|
||||
deleteErr error
|
||||
patchCalls int
|
||||
deletedID string
|
||||
deleteCtxErr error
|
||||
}
|
||||
|
||||
func (s *federationCreateStub) ListGroups(context.Context, identitydomains.ListGroupsRequest) (identitydomains.ListGroupsResponse, error) {
|
||||
return identitydomains.ListGroupsResponse{}, nil
|
||||
}
|
||||
|
||||
func (s *federationCreateStub) CreateIdentityProvider(context.Context, identitydomains.CreateIdentityProviderRequest) (identitydomains.CreateIdentityProviderResponse, error) {
|
||||
return identitydomains.CreateIdentityProviderResponse{IdentityProvider: s.created}, s.createErr
|
||||
}
|
||||
|
||||
func (s *federationCreateStub) GetIdentityProvider(context.Context, identitydomains.GetIdentityProviderRequest) (identitydomains.GetIdentityProviderResponse, error) {
|
||||
return identitydomains.GetIdentityProviderResponse{IdentityProvider: s.created}, nil
|
||||
}
|
||||
|
||||
func (s *federationCreateStub) PatchMappedAttribute(context.Context, identitydomains.PatchMappedAttributeRequest) (identitydomains.PatchMappedAttributeResponse, error) {
|
||||
s.patchCalls++
|
||||
return identitydomains.PatchMappedAttributeResponse{}, s.patchErr
|
||||
}
|
||||
|
||||
func (s *federationCreateStub) DeleteIdentityProvider(ctx context.Context, request identitydomains.DeleteIdentityProviderRequest) (identitydomains.DeleteIdentityProviderResponse, error) {
|
||||
s.deletedID = deref(request.IdentityProviderId)
|
||||
s.deleteCtxErr = ctx.Err()
|
||||
return identitydomains.DeleteIdentityProviderResponse{}, s.deleteErr
|
||||
}
|
||||
|
||||
func createdJitIdp(id string) identitydomains.IdentityProvider {
|
||||
return identitydomains.IdentityProvider{
|
||||
Id: idPtr(id), PartnerName: common.String("test-idp"), Type: identitydomains.IdentityProviderTypeSaml,
|
||||
JitUserProvEnabled: common.Bool(true), JitUserProvAttributes: &identitydomains.IdentityProviderJitUserProvAttributes{Value: common.String("mapping-1")},
|
||||
}
|
||||
}
|
||||
|
||||
func idPtr(value string) *string {
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
return common.String(value)
|
||||
}
|
||||
|
||||
func testJitInput() CreateIdpInput {
|
||||
return CreateIdpInput{Name: "test-idp", JitEnabled: true, IconURL: "https://img.example/icon.png"}
|
||||
}
|
||||
|
||||
func TestCreateSamlIdentityProviderCreateFailure(t *testing.T) {
|
||||
createErr := errors.New("create rejected")
|
||||
stub := &federationCreateStub{createErr: createErr}
|
||||
got, err := createSamlIdentityProvider(context.Background(), stub, testJitInput())
|
||||
if !errors.Is(err, createErr) {
|
||||
t.Fatalf("err = %v, want create error", err)
|
||||
}
|
||||
if got.ID != "" || stub.patchCalls != 0 || stub.deletedID != "" {
|
||||
t.Errorf("got = %+v, patchCalls = %d, deletedID = %q", got, stub.patchCalls, stub.deletedID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateSamlIdentityProviderJitFailureRollbackSuccess(t *testing.T) {
|
||||
patchErr := errors.New("patch rejected")
|
||||
stub := &federationCreateStub{created: createdJitIdp("idp-1"), patchErr: patchErr}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
got, err := createSamlIdentityProvider(ctx, stub, testJitInput())
|
||||
var partial *PartialIdentityProviderCreateError
|
||||
if !errors.Is(err, patchErr) || errors.As(err, &partial) {
|
||||
t.Fatalf("err = %v, want ordinary wrapped patch error", err)
|
||||
}
|
||||
if got.ID != "" || stub.deletedID != "idp-1" || stub.deleteCtxErr != nil {
|
||||
t.Errorf("got.ID = %q, deletedID = %q, deleteCtxErr = %v", got.ID, stub.deletedID, stub.deleteCtxErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateSamlIdentityProviderPartialCreateContract(t *testing.T) {
|
||||
cases := []struct {
|
||||
name, id string
|
||||
deleteErr error
|
||||
wantDelete bool
|
||||
}{
|
||||
{"rollback fails", "idp-1", errors.New("rollback secret detail"), true},
|
||||
{"created id missing", "", nil, false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) { assertPartialCreate(t, tc.id, tc.deleteErr, tc.wantDelete) })
|
||||
}
|
||||
}
|
||||
|
||||
func assertPartialCreate(t *testing.T, id string, deleteErr error, wantDelete bool) {
|
||||
t.Helper()
|
||||
stub := &federationCreateStub{created: createdJitIdp(id), patchErr: errors.New("patch secret detail"), deleteErr: deleteErr}
|
||||
got, err := createSamlIdentityProvider(context.Background(), stub, testJitInput())
|
||||
var partial *PartialIdentityProviderCreateError
|
||||
if !errors.As(err, &partial) || partial.IdentityProvider.ID != id {
|
||||
t.Fatalf("got = %+v, err = %v, partial = %+v", got, err, partial)
|
||||
}
|
||||
if strings.Contains(err.Error(), "secret") || (stub.deletedID != "") != wantDelete {
|
||||
t.Errorf("unsafe err = %q or deletedID = %q, wantDelete = %v", err, stub.deletedID, wantDelete)
|
||||
}
|
||||
}
|
||||
|
||||
func ruleReturn(name, value string) identitydomains.RuleReturn {
|
||||
return identitydomains.RuleReturn{Name: common.String(name), Value: common.String(value)}
|
||||
}
|
||||
|
||||
type rebuildIdpReturnCase struct {
|
||||
name string
|
||||
items []identitydomains.RuleReturn
|
||||
show bool
|
||||
changed bool
|
||||
want string
|
||||
}
|
||||
|
||||
func TestRebuildSamlIdpsReturn(t *testing.T) {
|
||||
local := ruleReturn("LocalIDPs", `["UserNamePassword"]`)
|
||||
cases := []rebuildIdpReturnCase{
|
||||
{"无SamlIDPs项时添加须补建", []identitydomains.RuleReturn{local}, true, true, `["idp-1"]`},
|
||||
{"无SamlIDPs项时移除无变化", []identitydomains.RuleReturn{local}, false, false, ""},
|
||||
{"已有其他IdP时追加", []identitydomains.RuleReturn{local, ruleReturn("SamlIDPs", `["other"]`)}, true, true, `["other","idp-1"]`},
|
||||
{"已在列表中再添加无变化", []identitydomains.RuleReturn{ruleReturn("SamlIDPs", `["idp-1"]`)}, true, false, `["idp-1"]`},
|
||||
{"移除目标IdP", []identitydomains.RuleReturn{ruleReturn("SamlIDPs", `["idp-1","other"]`)}, false, true, `["other"]`},
|
||||
{"空值项添加", []identitydomains.RuleReturn{ruleReturn("SamlIDPs", "")}, true, true, `["idp-1"]`},
|
||||
}
|
||||
assertRebuildIdpReturns(t, cases)
|
||||
}
|
||||
|
||||
func assertRebuildIdpReturns(t *testing.T, cases []rebuildIdpReturnCase) {
|
||||
t.Helper()
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
returns, changed, err := rebuildSamlIdpsReturn(tc.items, "idp-1", tc.show)
|
||||
if err != nil {
|
||||
t.Fatalf("rebuildSamlIdpsReturn: %v", err)
|
||||
}
|
||||
if changed != tc.changed {
|
||||
t.Errorf("changed = %v, want %v", changed, tc.changed)
|
||||
}
|
||||
got := ""
|
||||
for _, r := range returns {
|
||||
m := r.(map[string]string)
|
||||
if m["name"] == "SamlIDPs" {
|
||||
got = m["value"]
|
||||
}
|
||||
}
|
||||
if got != tc.want {
|
||||
t.Errorf("SamlIDPs = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRebuildSamlIdpsReturnBadJSON(t *testing.T) {
|
||||
items := []identitydomains.RuleReturn{ruleReturn("SamlIDPs", "not-json")}
|
||||
if _, _, err := rebuildSamlIdpsReturn(items, "idp-1", true); err == nil {
|
||||
t.Fatal("坏 JSON 应报错而非静默覆盖")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFetchSamlMetadataLimitsBody 锁定匿名元数据请求的响应体上限:超限报错而非吞下。
|
||||
func TestFetchSamlMetadataLimitsBody(t *testing.T) {
|
||||
huge := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = w.Write(bytes.Repeat([]byte("x"), samlMetadataMaxBytes+1))
|
||||
}))
|
||||
defer huge.Close()
|
||||
if _, _, err := fetchSamlMetadata(context.Background(), Credentials{}, huge.URL); err == nil ||
|
||||
!strings.Contains(err.Error(), "exceeds") {
|
||||
t.Fatalf("err = %v, want 超限错误", err)
|
||||
}
|
||||
ok := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = w.Write([]byte("<EntityDescriptor/>"))
|
||||
}))
|
||||
defer ok.Close()
|
||||
body, status, err := fetchSamlMetadata(context.Background(), Credentials{}, ok.URL)
|
||||
if err != nil || status != http.StatusOK || string(body) != "<EntityDescriptor/>" {
|
||||
t.Fatalf("fetch = %q, %d, %v; want 正常返回", body, status, err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user