package httpapi import ( "context" "encoding/json" "errors" "io" "log/slog" "net/http" "net/http/httptest" "testing" "time" "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" "github.com/easyai/easyai-ai-gateway/apps/api/internal/identity" "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" ) type fakeOIDCUserResolver struct { result store.ResolveOrProvisionOIDCUserResult err error calls int input store.ResolveOrProvisionOIDCUserInput binding store.OIDCTenantBindingContext bindingErr error bindingCalls int } type fakeTenantContextReader struct { tenant identity.TenantContext err error calls *int etag *string } func (reader fakeTenantContextReader) Get(_ context.Context, _, etag string) (identity.TenantContext, bool, error) { if reader.calls != nil { (*reader.calls)++ } if reader.etag != nil { *reader.etag = etag } return reader.tenant, false, reader.err } func (f *fakeOIDCUserResolver) ResolveOrProvisionOIDCUser(_ context.Context, input store.ResolveOrProvisionOIDCUserInput) (store.ResolveOrProvisionOIDCUserResult, error) { f.calls++ f.input = input return f.result, f.err } func (f *fakeOIDCUserResolver) OIDCTenantBindingContext(context.Context, string, string, string) (store.OIDCTenantBindingContext, error) { f.bindingCalls++ if f.binding.ID == "" && f.bindingErr == nil { return store.OIDCTenantBindingContext{}, store.ErrOIDCTenantUnavailable } return f.binding, f.bindingErr } func TestResolveGatewayUserAddsLocalOIDCContext(t *testing.T) { resolver := &fakeOIDCUserResolver{result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{ ID: "platform-user", Username: "alice", Roles: []string{"basic"}, TenantID: "external-tenant", Source: "oidc", GatewayUserID: "21dd9ccb-3793-4023-ab31-4d04982ca4d3", GatewayTenantID: "8f17f3ac-136e-4d0f-b097-655e2a6240a3", TenantKey: "default", UserGroupID: "6dcf86f2-8eaf-4b43-8e69-181315db24f0", }}} server := &Server{ identityTestRevision: identity.Revision{ Issuer: "https://auth.test.example/issuer", LocalTenantKey: "default", JITEnabled: true, }, oidcUserResolver: resolver, logger: slog.New(slog.NewTextHandler(io.Discard, nil)), } next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { user, ok := auth.UserFromContext(r.Context()) if !ok || user.GatewayUserID == "" || user.GatewayTenantID == "" || user.UserGroupID == "" { t.Fatalf("resolved Gateway context missing: %+v", user) } writeJSON(w, http.StatusOK, user) }) request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) request = request.WithContext(auth.WithUser(request.Context(), &auth.User{ ID: "platform-user", Username: "alice", Roles: []string{"basic"}, TenantID: "external-tenant", Source: "oidc", })) recorder := httptest.NewRecorder() server.resolveGatewayUser(next).ServeHTTP(recorder, request) if recorder.Code != http.StatusOK { t.Fatalf("status = %d, want 200", recorder.Code) } if resolver.calls != 1 || resolver.input.Subject != "platform-user" || resolver.input.GatewayTenantKey != "default" || !resolver.input.ProvisioningEnabled { t.Fatalf("unexpected resolver call: calls=%d input=%+v", resolver.calls, resolver.input) } } func TestResolveGatewayUserLeavesNonOIDCIdentityChainsUnchanged(t *testing.T) { for _, source := range []string{"gateway", "api_key", "server-main"} { t.Run(source, func(t *testing.T) { resolver := &fakeOIDCUserResolver{} server := &Server{oidcUserResolver: resolver, logger: slog.New(slog.NewTextHandler(io.Discard, nil))} request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) original := &auth.User{ID: "local-user", Source: source, GatewayUserID: "local-user"} request = request.WithContext(auth.WithUser(request.Context(), original)) recorder := httptest.NewRecorder() server.resolveGatewayUser(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { resolved, _ := auth.UserFromContext(r.Context()) if resolved != original { t.Fatalf("non-OIDC identity context was replaced: %+v", resolved) } w.WriteHeader(http.StatusNoContent) })).ServeHTTP(recorder, request) if recorder.Code != http.StatusNoContent || resolver.calls != 0 { t.Fatalf("status=%d resolver calls=%d", recorder.Code, resolver.calls) } }) } } func TestResolveOIDCMultiTenantProjectionUsesRuntimeTenantContext(t *testing.T) { tenantID := "d9dcb4e7-6938-4547-af68-10ea404aa4b0" applicationID := "0a3565fa-c0b0-4862-be20-d3a7b37bb7c8" resolver := &fakeOIDCUserResolver{result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{GatewayUserID: "local-user"}}} server := &Server{oidcUserResolver: resolver} runtime := &identityRequestRuntime{ Revision: identity.Revision{ Issuer: "https://auth.test.example", ApplicationID: applicationID, TenantMode: "multi_tenant", JITEnabled: true, }, TenantContext: fakeTenantContextReader{tenant: identity.TenantContext{ ApplicationID: applicationID, TenantID: tenantID, DisplayName: "租户 A", Slug: "tenant-a", TenantStatus: "active", TenantApplicationStatus: "active", Version: "v7", UpdatedAt: time.Date(2026, 7, 28, 9, 0, 0, 0, time.UTC), ETag: `"tenant-v7"`, }}, } request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) _, err := server.resolveOIDCUserProjectionForRuntime(request.Context(), request, &auth.User{ ID: "shared-subject", TenantID: tenantID, Username: "alice", Roles: []string{"basic"}, }, runtime) if err != nil { t.Fatalf("resolve multi-tenant projection: %v", err) } if resolver.input.TenantMode != "multi_tenant" || resolver.input.ApplicationID != applicationID || resolver.input.TenantName != "租户 A" || resolver.input.TenantSlug != "tenant-a" || resolver.input.TenantMetadataStatus != "synced" || resolver.input.TenantMetadataVersion != "v7" || resolver.input.TenantMetadataETag != `"tenant-v7"` { t.Fatalf("unexpected tenant context projection input: %+v", resolver.input) } } func TestResolveOIDCMultiTenantProjectionTreatsTemporaryContextFailureAsPending(t *testing.T) { tenantID := "d9dcb4e7-6938-4547-af68-10ea404aa4b0" resolver := &fakeOIDCUserResolver{result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{GatewayUserID: "local-user"}}} server := &Server{oidcUserResolver: resolver} runtime := &identityRequestRuntime{ Revision: identity.Revision{ Issuer: "https://auth.test.example", ApplicationID: "0a3565fa-c0b0-4862-be20-d3a7b37bb7c8", TenantMode: "multi_tenant", JITEnabled: true, }, TenantContext: fakeTenantContextReader{err: identity.ErrTenantContextUnavailable}, } request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) if _, err := server.resolveOIDCUserProjectionForRuntime(request.Context(), request, &auth.User{ ID: "subject", TenantID: tenantID, }, runtime); err != nil { t.Fatalf("temporary tenant context failure should reach fail-closed store projection: %v", err) } if resolver.input.TenantMetadataStatus != "metadata_pending" || resolver.input.TenantName != "认证中心租户 d9dcb4e7" { t.Fatalf("unexpected pending projection input: %+v", resolver.input) } } func TestResolveOIDCMultiTenantProjectionUsesFreshLocalTenantCache(t *testing.T) { tenantID := "d9dcb4e7-6938-4547-af68-10ea404aa4b0" remoteCalls := 0 resolver := &fakeOIDCUserResolver{ result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{GatewayUserID: "local-user"}}, binding: store.OIDCTenantBindingContext{ ID: "binding-a", AccessStatus: "active", MetadataStatus: "synced", DisplayName: "缓存租户 A", Slug: "tenant-a", Version: "v8", ETag: `"tenant-v8"`, MetadataUpdatedAt: time.Date(2026, 7, 28, 10, 0, 0, 0, time.UTC), NextSyncAt: time.Now().Add(time.Minute), }, } server := &Server{oidcUserResolver: resolver} runtime := &identityRequestRuntime{ Revision: identity.Revision{ Issuer: "https://auth.test.example", ApplicationID: "0a3565fa-c0b0-4862-be20-d3a7b37bb7c8", TenantMode: "multi_tenant", JITEnabled: true, }, TenantContext: fakeTenantContextReader{ err: errors.New("fresh cache must avoid a remote request"), calls: &remoteCalls, }, } request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) if _, err := server.resolveOIDCUserProjectionForRuntime(request.Context(), request, &auth.User{ ID: "subject", TenantID: tenantID, }, runtime); err != nil { t.Fatal(err) } if remoteCalls != 0 || resolver.input.TenantName != "缓存租户 A" || resolver.input.TenantMetadataVersion != "v8" { t.Fatalf("remote calls=%d input=%+v", remoteCalls, resolver.input) } } func TestResolveOIDCMultiTenantProjectionFallsBackToSyncedCacheOnTemporaryFailure(t *testing.T) { tenantID := "d9dcb4e7-6938-4547-af68-10ea404aa4b0" resolver := &fakeOIDCUserResolver{ result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{GatewayUserID: "local-user"}}, binding: store.OIDCTenantBindingContext{ ID: "binding-a", AccessStatus: "active", MetadataStatus: "synced", DisplayName: "缓存租户 A", Slug: "tenant-a", Version: "v7", ETag: `"tenant-v7"`, NextSyncAt: time.Now().Add(-time.Minute), }, } server := &Server{oidcUserResolver: resolver} request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) if _, err := server.resolveOIDCUserProjectionForRuntime(request.Context(), request, &auth.User{ ID: "subject", TenantID: tenantID, }, &identityRequestRuntime{ Revision: identity.Revision{ Issuer: "https://auth.test.example", ApplicationID: "0a3565fa-c0b0-4862-be20-d3a7b37bb7c8", TenantMode: "multi_tenant", JITEnabled: true, }, TenantContext: fakeTenantContextReader{err: identity.ErrTenantContextUnavailable}, }); err != nil { t.Fatal(err) } if resolver.input.TenantName != "缓存租户 A" || resolver.input.TenantMetadataStatus != "synced" { t.Fatalf("cached fallback input=%+v", resolver.input) } } func TestResolveOIDCMultiTenantProjectionRevalidatesDisabledBindingBeforeReassignment(t *testing.T) { tenantID := "d9dcb4e7-6938-4547-af68-10ea404aa4b0" applicationID := "0a3565fa-c0b0-4862-be20-d3a7b37bb7c8" remoteETag := "not-called" resolver := &fakeOIDCUserResolver{ result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{GatewayUserID: "local-user"}}, binding: store.OIDCTenantBindingContext{ ID: "binding-a", AccessStatus: "disabled", MetadataStatus: "rejected", ETag: `"revoked-v7"`, }, } server := &Server{oidcUserResolver: resolver} request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) _, err := server.resolveOIDCUserProjectionForRuntime(request.Context(), request, &auth.User{ ID: "subject", TenantID: tenantID, }, &identityRequestRuntime{ Revision: identity.Revision{ Issuer: "https://auth.test.example", ApplicationID: applicationID, TenantMode: "multi_tenant", JITEnabled: true, }, TenantContext: fakeTenantContextReader{tenant: identity.TenantContext{ ApplicationID: applicationID, TenantID: tenantID, DisplayName: "恢复后的租户 A", Slug: "tenant-a", TenantStatus: "active", TenantApplicationStatus: "active", Version: "v8", UpdatedAt: time.Now(), }, etag: &remoteETag}, }) if err != nil { t.Fatalf("revalidate reassigned tenant: %v", err) } if remoteETag != "" || resolver.calls != 1 || resolver.input.TenantMetadataStatus != "synced" { t.Fatalf("etag=%q calls=%d input=%+v", remoteETag, resolver.calls, resolver.input) } } func TestResolveOIDCMultiTenantProjectionKeepsDisabledBindingClosedDuringContextFailure(t *testing.T) { resolver := &fakeOIDCUserResolver{binding: store.OIDCTenantBindingContext{ ID: "binding-a", AccessStatus: "disabled", MetadataStatus: "rejected", }} server := &Server{oidcUserResolver: resolver} request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) _, err := server.resolveOIDCUserProjectionForRuntime(request.Context(), request, &auth.User{ ID: "subject", TenantID: "d9dcb4e7-6938-4547-af68-10ea404aa4b0", }, &identityRequestRuntime{ Revision: identity.Revision{ Issuer: "https://auth.test.example", ApplicationID: "0a3565fa-c0b0-4862-be20-d3a7b37bb7c8", TenantMode: "multi_tenant", JITEnabled: true, }, TenantContext: fakeTenantContextReader{err: identity.ErrTenantContextUnavailable}, }) if !errors.Is(err, store.ErrOIDCTenantUnavailable) || resolver.calls != 0 { t.Fatalf("err=%v resolver calls=%d", err, resolver.calls) } } func TestResolveOIDCMultiTenantProjectionRejectsMissingOrInactiveTenant(t *testing.T) { tenantID := "d9dcb4e7-6938-4547-af68-10ea404aa4b0" tests := []struct { name string reader fakeTenantContextReader }{ {name: "not found", reader: fakeTenantContextReader{err: identity.ErrTenantContextNotFound}}, {name: "inactive", reader: fakeTenantContextReader{tenant: identity.TenantContext{ TenantStatus: "suspended", TenantApplicationStatus: "active", }}}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { resolver := &fakeOIDCUserResolver{} server := &Server{oidcUserResolver: resolver} request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) _, err := server.resolveOIDCUserProjectionForRuntime(request.Context(), request, &auth.User{ ID: "subject", TenantID: tenantID, }, &identityRequestRuntime{ Revision: identity.Revision{TenantMode: "multi_tenant"}, TenantContext: test.reader, }) if !errors.Is(err, store.ErrOIDCTenantUnavailable) || resolver.calls != 0 { t.Fatalf("err=%v resolver calls=%d", err, resolver.calls) } }) } } func TestResolveGatewayUserReturnsStructuredErrors(t *testing.T) { tests := []struct { name string err error status int code string }{ {name: "not provisioned", err: store.ErrOIDCUserNotProvisioned, status: http.StatusForbidden, code: "GATEWAY_USER_NOT_PROVISIONED"}, {name: "disabled", err: store.ErrOIDCUserDisabled, status: http.StatusForbidden, code: "GATEWAY_USER_DISABLED"}, {name: "tenant unavailable", err: store.ErrOIDCTenantUnavailable, status: http.StatusServiceUnavailable, code: "GATEWAY_TENANT_UNAVAILABLE"}, {name: "storage failure", err: errors.New("database unavailable"), status: http.StatusServiceUnavailable, code: "GATEWAY_USER_PROVISIONING_FAILED"}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { server := &Server{ identityTestRevision: identity.Revision{Issuer: "https://auth.test.example", LocalTenantKey: "default"}, oidcUserResolver: &fakeOIDCUserResolver{err: test.err}, logger: slog.New(slog.NewTextHandler(io.Discard, nil)), } request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil) request = request.WithContext(auth.WithUser(request.Context(), &auth.User{ ID: "platform-user", Source: "oidc", TenantID: "external-tenant", })) recorder := httptest.NewRecorder() server.resolveGatewayUser(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { t.Fatal("next handler must not run") })).ServeHTTP(recorder, request) if recorder.Code != test.status { t.Fatalf("status = %d, want %d", recorder.Code, test.status) } var envelope struct { Error struct { Code string `json:"code"` Message string `json:"message"` Status int `json:"status"` } `json:"error"` } if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil { t.Fatalf("decode error envelope: %v", err) } if envelope.Error.Code != test.code || envelope.Error.Status != test.status || envelope.Error.Message == "" { t.Fatalf("unexpected error envelope: %+v", envelope) } if envelope.Error.Message == test.err.Error() { t.Fatalf("internal error leaked to response: %q", envelope.Error.Message) } }) } }