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("")) })) defer ok.Close() body, status, err := fetchSamlMetadata(context.Background(), Credentials{}, ok.URL) if err != nil || status != http.StatusOK || string(body) != "" { t.Fatalf("fetch = %q, %d, %v; want 正常返回", body, status, err) } }