197 lines
7.5 KiB
Go
197 lines
7.5 KiB
Go
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)
|
|
}
|
|
}
|