Files
easyai-ai-gateway/apps/api/internal/httpapi/oidc_user_middleware_test.go
T
chengcheng 5c679ff13f feat(identity): 接入认证中心多租户登录
支持 Manifest V2 动态 tid 验证、Tenant Context 同步和租户内 JIT 投影,并保留 Manifest V1 与旧 Session 兼容。\n\n增加 tenantHint、租户切换、普通注册关闭及 application/principal/tenant 两级 SSF 撤销;迁移、定向安全测试和本地双租户跨仓 E2E 已通过。\n\nrelease_required=true;未执行 Release、Staging 或真实链路。
2026-07-28 17:28:35 +08:00

378 lines
16 KiB
Go

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)
}
})
}
}