Files
easyai-ai-gateway/apps/api/internal/store/oidc_users_integration_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

481 lines
18 KiB
Go

package store
import (
"context"
"os"
"path/filepath"
"runtime"
"sort"
"strings"
"sync"
"testing"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identity"
"github.com/google/uuid"
"github.com/jackc/pgx/v5/pgxpool"
)
func TestResolveOrProvisionOIDCMultiTenantUserIsIdempotentIsolatedAndReusable(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run OIDC multi-tenant PostgreSQL integration tests")
}
ctx := context.Background()
applyOIDCJITTestMigrations(t, ctx, databaseURL)
db, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatal(err)
}
defer db.Close()
issuer := "https://auth.test.example/issuer/shared"
applicationID, tenantA, tenantB, subject := uuid.NewString(), uuid.NewString(), uuid.NewString(), uuid.NewString()
input := func(tenantID, name, slug string) ResolveOrProvisionOIDCUserInput {
return ResolveOrProvisionOIDCUserInput{
Issuer: issuer, ApplicationID: applicationID, Subject: subject, Username: "shared-subject",
Roles: []string{"basic"}, TenantID: tenantID, TenantMode: "multi_tenant",
TenantName: name, TenantSlug: slug, TenantMetadataStatus: "synced",
TenantMetadataVersion: "v1", TenantMetadataETag: `"v1"`,
TenantMetadataUpdatedAt: time.Unix(1_780_000_000, 0).UTC(),
OIDCClientID: "gateway-browser", ProvisioningEnabled: true,
}
}
const concurrentLogins = 8
results := make([]ResolveOrProvisionOIDCUserResult, concurrentLogins)
errs := make([]error, concurrentLogins)
var wait sync.WaitGroup
for index := range concurrentLogins {
wait.Add(1)
go func(index int) {
defer wait.Done()
results[index], errs[index] = db.ResolveOrProvisionOIDCUser(ctx, input(tenantA, "Tenant A", "tenant-a"))
}(index)
}
wait.Wait()
var userA *auth.User
created := 0
for index, result := range results {
if errs[index] != nil || result.User == nil {
t.Fatalf("tenant A login %d result=%#v error=%v", index, result, errs[index])
}
if userA == nil {
userA = result.User
}
if result.User.GatewayUserID != userA.GatewayUserID || result.User.GatewayTenantID != userA.GatewayTenantID {
t.Fatalf("concurrent tenant A projection diverged: first=%#v current=%#v", userA, result.User)
}
if result.Created {
created++
}
}
if created != 1 {
t.Fatalf("tenant A created projections=%d", created)
}
resultB, err := db.ResolveOrProvisionOIDCUser(ctx, input(tenantB, "Tenant B", "tenant-b"))
if err != nil || resultB.User == nil {
t.Fatalf("tenant B projection=%#v error=%v", resultB, err)
}
userB := resultB.User
if userA.GatewayUserID == userB.GatewayUserID || userA.GatewayTenantID == userB.GatewayTenantID ||
userA.ID != userB.ID {
t.Fatalf("same subject was not isolated by tenant: A=%#v B=%#v", userA, userB)
}
var tenants, users, wallets, audits int
if err := db.pool.QueryRow(ctx, `SELECT
(SELECT count(*) FROM gateway_oidc_tenant_bindings WHERE application_id=$1 AND external_tenant_id IN ($2,$3)),
(SELECT count(*) FROM gateway_oidc_user_bindings binding
JOIN gateway_oidc_tenant_bindings tenant ON tenant.id=binding.tenant_binding_id
WHERE tenant.application_id=$1 AND binding.subject=$4),
(SELECT count(*) FROM gateway_wallet_accounts WHERE gateway_user_id IN ($5::uuid,$6::uuid) AND currency='resource'),
(SELECT count(*) FROM gateway_audit_logs WHERE action='identity.oidc_user.provisioned'
AND target_gateway_user_id IN ($5::uuid,$6::uuid))`,
applicationID, tenantA, tenantB, subject, userA.GatewayUserID, userB.GatewayUserID,
).Scan(&tenants, &users, &wallets, &audits); err != nil {
t.Fatal(err)
}
if tenants != 2 || users != 2 || wallets != 2 || audits != 2 {
t.Fatalf("tenants=%d users=%d wallets=%d audits=%d", tenants, users, wallets, audits)
}
if _, err := db.CreateAPIKey(ctx, CreateAPIKeyInput{Name: "Tenant A only"}, userA); err != nil {
t.Fatal(err)
}
keysB, err := db.ListAPIKeys(ctx, userB)
if err != nil || len(keysB) != 0 {
t.Fatalf("tenant B observed tenant A API keys: keys=%#v error=%v", keysB, err)
}
if _, err := db.CreateTask(ctx, CreateTaskInput{
Kind: "multi-tenant-isolation", Model: "local-fixture", RunMode: "async", Async: true,
Request: map[string]any{"prompt": "tenant-a"},
}, userA); err != nil {
t.Fatal(err)
}
tasksB, err := db.ListTasks(ctx, userB, TaskListFilter{})
if err != nil || len(tasksB.Items) != 0 {
t.Fatalf("tenant B observed tenant A tasks: tasks=%#v error=%v", tasksB.Items, err)
}
bindingA, err := db.OIDCTenantBindingContext(ctx, issuer, applicationID, tenantA)
if err != nil {
t.Fatal(err)
}
if err := db.RejectOIDCTenantBinding(ctx, bindingA.ID, "tenant_application_revoked", time.Now().UTC()); err != nil {
t.Fatal(err)
}
pending := input(tenantA, "", "")
pending.TenantMetadataStatus = "metadata_pending"
if _, err := db.ResolveOrProvisionOIDCUser(ctx, pending); !errorsIs(err, ErrOIDCTenantUnavailable) {
t.Fatalf("disabled binding accepted pending context: %v", err)
}
reactivated, err := db.ResolveOrProvisionOIDCUser(ctx, input(tenantA, "Tenant A restored", "tenant-a"))
if err != nil {
t.Fatal(err)
}
if reactivated.User.GatewayTenantID != userA.GatewayTenantID || reactivated.User.GatewayUserID != userA.GatewayUserID {
t.Fatalf("reassignment did not reuse local projection: before=%#v after=%#v", userA, reactivated.User)
}
if err := db.ApplyOIDCTenantBindingSync(ctx, bindingA.ID, identity.TenantContext{
ApplicationID: applicationID, TenantID: tenantA, DisplayName: "Tenant A renamed", Slug: "tenant-a",
TenantStatus: "active", TenantApplicationStatus: "active", Version: "v2", ETag: `"v2"`,
UpdatedAt: time.Unix(1_780_000_100, 0).UTC(),
}, false, time.Now().UTC()); err != nil {
t.Fatal(err)
}
synced, err := db.OIDCTenantBindingContext(ctx, issuer, applicationID, tenantA)
if err != nil || synced.DisplayName != "Tenant A renamed" || synced.Version != "v2" {
t.Fatalf("synced tenant context=%#v error=%v", synced, err)
}
}
func TestResolveOrProvisionOIDCUserLifecycleAndConcurrency(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run OIDC JIT PostgreSQL integration tests")
}
ctx := context.Background()
applyOIDCJITTestMigrations(t, ctx, databaseURL)
db, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect store: %v", err)
}
defer db.Close()
suffix := time.Now().UTC().Format("20060102150405.000000000")
subject := "platform-jit-" + suffix
input := ResolveOrProvisionOIDCUserInput{
Issuer: "https://auth.test.example/issuer/shared",
Subject: subject,
Username: "jit-user-" + suffix,
Roles: []string{"basic"},
TenantID: "auth-center-test-tenant",
GatewayTenantKey: "default",
ProvisioningEnabled: true,
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `
DELETE FROM gateway_audit_logs
WHERE target_gateway_user_id IN (
SELECT id FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1
)`, subject)
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1`, subject)
})
const callers = 12
results := make([]ResolveOrProvisionOIDCUserResult, callers)
errs := make([]error, callers)
var wg sync.WaitGroup
for index := 0; index < callers; index++ {
wg.Add(1)
go func(index int) {
defer wg.Done()
results[index], errs[index] = db.ResolveOrProvisionOIDCUser(ctx, input)
}(index)
}
wg.Wait()
firstID := ""
createdCount := 0
auditID := ""
for index, err := range errs {
if err != nil {
t.Fatalf("concurrent resolve %d: %v", index, err)
}
result := results[index]
if result.User == nil || result.User.GatewayUserID == "" {
t.Fatalf("concurrent resolve %d returned no local user: %+v", index, result)
}
if firstID == "" {
firstID = result.User.GatewayUserID
}
if result.User.GatewayUserID != firstID {
t.Fatalf("concurrent resolve returned different users: %q and %q", firstID, result.User.GatewayUserID)
}
if result.Created {
createdCount++
auditID = result.AuditID
}
}
if createdCount != 1 {
t.Fatalf("created count = %d, want 1", createdCount)
}
if auditID == "" {
t.Fatal("first provision must return an audit ID")
}
var users, wallets, audits int
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1`, subject).Scan(&users); err != nil {
t.Fatalf("count OIDC users: %v", err)
}
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_wallet_accounts WHERE gateway_user_id = $1::uuid AND currency = 'resource'`, firstID).Scan(&wallets); err != nil {
t.Fatalf("count wallets: %v", err)
}
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_audit_logs WHERE action = 'identity.oidc_user.provisioned' AND target_gateway_user_id = $1::uuid`, firstID).Scan(&audits); err != nil {
t.Fatalf("count audits: %v", err)
}
if users != 1 || wallets != 1 || audits != 1 {
t.Fatalf("users=%d wallets=%d audits=%d, want one of each", users, wallets, audits)
}
var auditProjection string
if err := db.pool.QueryRow(ctx, `
SELECT COALESCE(actor_user_id, '') || metadata::text || after_state::text
FROM gateway_audit_logs
WHERE id = $1::uuid`, auditID).Scan(&auditProjection); err != nil {
t.Fatalf("read OIDC provisioning audit: %v", err)
}
if strings.Contains(auditProjection, subject) || strings.Contains(auditProjection, input.Issuer) {
t.Fatal("OIDC provisioning audit exposed raw external identity claims")
}
if _, err := db.pool.Exec(ctx, `
UPDATE gateway_users
SET display_name = 'Manual Display Name',
email = 'manual-profile@example.test',
metadata = metadata || '{"manualProfile":true}'::jsonb
WHERE id = $1::uuid`, firstID); err != nil {
t.Fatalf("seed manually managed profile fields: %v", err)
}
input.Username = "jit-user-renamed-" + suffix
input.Roles = []string{"basic", "admin"}
input.ProvisioningEnabled = false
repeated, err := db.ResolveOrProvisionOIDCUser(ctx, input)
if err != nil {
t.Fatalf("repeat resolve: %v", err)
}
if repeated.Created || repeated.AuditID != "" || repeated.User.GatewayUserID != firstID {
t.Fatalf("unexpected repeat result: %+v", repeated)
}
if repeated.User.Username != input.Username || !containsOIDCTestRole(repeated.User.Roles, "admin") {
t.Fatalf("repeat resolve did not sync token projection: %+v", repeated.User)
}
var displayName, email string
var manualProfile bool
if err := db.pool.QueryRow(ctx, `
SELECT COALESCE(display_name, ''), COALESCE(email, ''), COALESCE((metadata->>'manualProfile')::boolean, false)
FROM gateway_users WHERE id = $1::uuid`, firstID).Scan(&displayName, &email, &manualProfile); err != nil {
t.Fatalf("read manually managed profile fields: %v", err)
}
if displayName != "Manual Display Name" || email != "manual-profile@example.test" || !manualProfile {
t.Fatalf("repeat resolve overwrote manually managed profile fields")
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_users SET status = 'disabled' WHERE id = $1::uuid`, firstID); err != nil {
t.Fatalf("disable OIDC user: %v", err)
}
if _, err := db.ResolveOrProvisionOIDCUser(ctx, input); !errorsIs(err, ErrOIDCUserDisabled) {
t.Fatalf("disabled resolve error = %v, want ErrOIDCUserDisabled", err)
}
var status string
if err := db.pool.QueryRow(ctx, `SELECT status FROM gateway_users WHERE id = $1::uuid`, firstID).Scan(&status); err != nil {
t.Fatalf("read disabled status: %v", err)
}
if status != "disabled" {
t.Fatalf("disabled OIDC user was reactivated: %q", status)
}
}
func TestResolveOrProvisionOIDCUserRejectsMissingMappingWithoutWrites(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run OIDC JIT PostgreSQL integration tests")
}
ctx := context.Background()
applyOIDCJITTestMigrations(t, ctx, databaseURL)
db, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect store: %v", err)
}
defer db.Close()
subject := "platform-jit-missing-" + time.Now().UTC().Format("20060102150405.000000000")
input := ResolveOrProvisionOIDCUserInput{
Issuer: "https://auth.test.example/issuer/shared",
Subject: subject,
Username: "missing-user",
Roles: []string{"basic"},
TenantID: "auth-center-test-tenant",
GatewayTenantKey: "missing-tenant-key",
}
if _, err := db.ResolveOrProvisionOIDCUser(ctx, input); !errorsIs(err, ErrOIDCUserNotProvisioned) {
t.Fatalf("disabled JIT error = %v, want ErrOIDCUserNotProvisioned", err)
}
input.ProvisioningEnabled = true
if _, err := db.ResolveOrProvisionOIDCUser(ctx, input); !errorsIs(err, ErrOIDCTenantUnavailable) {
t.Fatalf("missing tenant error = %v, want ErrOIDCTenantUnavailable", err)
}
var users int
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1`, subject).Scan(&users); err != nil {
t.Fatalf("count rejected users: %v", err)
}
if users != 0 {
t.Fatalf("rejected OIDC request created %d users", users)
}
suffix := time.Now().UTC().Format("20060102150405.000000000")
groupKey := "jit-disabled-group-" + suffix
tenantKey := "jit-disabled-tenant-" + suffix
var groupID string
if err := db.pool.QueryRow(ctx, `
INSERT INTO gateway_user_groups (group_key, name, status)
VALUES ($1, 'OIDC JIT disabled group test', 'active')
RETURNING id::text`, groupKey).Scan(&groupID); err != nil {
t.Fatalf("create disabled-mapping test group: %v", err)
}
var tenantID string
if err := db.pool.QueryRow(ctx, `
INSERT INTO gateway_tenants (tenant_key, name, default_user_group_id, status)
VALUES ($1, 'OIDC JIT disabled tenant test', $2::uuid, 'disabled')
RETURNING id::text`, tenantKey, groupID).Scan(&tenantID); err != nil {
t.Fatalf("create disabled-mapping test tenant: %v", err)
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_tenants WHERE id = $1::uuid`, tenantID)
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_user_groups WHERE id = $1::uuid`, groupID)
})
disabledMappingInput := input
disabledMappingInput.Subject += "-disabled-mapping"
disabledMappingInput.GatewayTenantKey = tenantKey
if _, err := db.ResolveOrProvisionOIDCUser(ctx, disabledMappingInput); !errorsIs(err, ErrOIDCTenantUnavailable) {
t.Fatalf("disabled tenant error = %v, want ErrOIDCTenantUnavailable", err)
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_tenants SET status = 'active' WHERE id = $1::uuid`, tenantID); err != nil {
t.Fatalf("enable test tenant: %v", err)
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_user_groups SET status = 'disabled' WHERE id = $1::uuid`, groupID); err != nil {
t.Fatalf("disable test group: %v", err)
}
if _, err := db.ResolveOrProvisionOIDCUser(ctx, disabledMappingInput); !errorsIs(err, ErrOIDCTenantUnavailable) {
t.Fatalf("disabled user group error = %v, want ErrOIDCTenantUnavailable", err)
}
}
func TestOIDCUserKeyIsStableAndDoesNotExposeClaims(t *testing.T) {
issuer := "https://auth.test.example/issuer/shared"
subject := "platform-sensitive-subject"
first := deriveOIDCUserKey(issuer, subject)
second := deriveOIDCUserKey(issuer+"/", subject)
if first == "" || first != second {
t.Fatalf("OIDC user key is not stable: %q != %q", first, second)
}
if strings.Contains(first, subject) || strings.Contains(first, issuer) {
t.Fatalf("OIDC user key exposes raw claims: %q", first)
}
if first == deriveOIDCUserKey(issuer, subject+"-other") {
t.Fatal("different subjects produced the same OIDC user key")
}
}
func errorsIs(err error, target error) bool {
for err != nil {
if err == target {
return true
}
type unwrapper interface{ Unwrap() error }
wrapped, ok := err.(unwrapper)
if !ok {
return false
}
err = wrapped.Unwrap()
}
return false
}
func containsOIDCTestRole(roles []string, expected string) bool {
for _, role := range roles {
if role == expected {
return true
}
}
return false
}
func applyOIDCJITTestMigrations(t *testing.T, ctx context.Context, databaseURL string) {
t.Helper()
_, filename, _, _ := runtime.Caller(0)
migrationFiles, err := filepath.Glob(filepath.Join(filepath.Dir(filename), "..", "..", "migrations", "*.sql"))
if err != nil {
t.Fatalf("read migration files: %v", err)
}
sort.Strings(migrationFiles)
pool, err := pgxpool.New(ctx, databaseURL)
if err != nil {
t.Fatalf("connect migration db: %v", err)
}
defer pool.Close()
if _, err := pool.Exec(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (version text PRIMARY KEY, applied_at timestamptz NOT NULL DEFAULT now())`); err != nil {
t.Fatalf("ensure schema migrations: %v", err)
}
for _, migrationPath := range migrationFiles {
version := strings.TrimSuffix(filepath.Base(migrationPath), filepath.Ext(migrationPath))
var exists bool
if err := pool.QueryRow(ctx, `SELECT EXISTS (SELECT 1 FROM schema_migrations WHERE version = $1)`, version).Scan(&exists); err != nil {
t.Fatalf("check migration %s: %v", version, err)
}
if exists {
continue
}
migration, err := os.ReadFile(migrationPath)
if err != nil {
t.Fatalf("read migration %s: %v", version, err)
}
migrationSQL := string(migration)
const noTransactionMarker = "-- easyai:migration:no-transaction"
if strings.HasPrefix(strings.TrimSpace(migrationSQL), noTransactionMarker) {
migrationSQL = strings.TrimSpace(strings.TrimPrefix(strings.TrimSpace(migrationSQL), noTransactionMarker))
for _, statement := range strings.Split(migrationSQL, "-- easyai:migration:statement") {
if statement = strings.TrimSpace(statement); statement != "" {
if _, err := pool.Exec(ctx, statement); err != nil {
t.Fatalf("apply non-transaction migration %s: %v", version, err)
}
}
}
if _, err := pool.Exec(ctx, `INSERT INTO schema_migrations(version) VALUES($1)`, version); err != nil {
t.Fatalf("record non-transaction migration %s: %v", version, err)
}
continue
}
tx, err := pool.Begin(ctx)
if err != nil {
t.Fatalf("begin migration %s: %v", version, err)
}
if _, err := tx.Exec(ctx, string(migration)); err != nil {
_ = tx.Rollback(ctx)
t.Fatalf("apply migration %s: %v", version, err)
}
if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version) VALUES($1)`, version); err != nil {
_ = tx.Rollback(ctx)
t.Fatalf("record migration %s: %v", version, err)
}
if err := tx.Commit(ctx); err != nil {
t.Fatalf("commit migration %s: %v", version, err)
}
}
}