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 或真实链路。
This commit is contained in:
@@ -34,27 +34,31 @@ const (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type User struct {
|
type User struct {
|
||||||
ID string `json:"sub"`
|
ID string `json:"sub"`
|
||||||
Username string `json:"username"`
|
Username string `json:"username"`
|
||||||
Roles []string `json:"role,omitempty"`
|
Roles []string `json:"role,omitempty"`
|
||||||
TenantID string `json:"tenantId,omitempty"`
|
TenantID string `json:"tenantId,omitempty"`
|
||||||
GatewayTenantID string `json:"gatewayTenantId,omitempty"`
|
TenantName string `json:"tenantName,omitempty"`
|
||||||
TenantKey string `json:"tenantKey,omitempty"`
|
GatewayTenantID string `json:"gatewayTenantId,omitempty"`
|
||||||
SSOID string `json:"sso_id,omitempty"`
|
TenantKey string `json:"tenantKey,omitempty"`
|
||||||
Source string `json:"source,omitempty"`
|
SSOID string `json:"sso_id,omitempty"`
|
||||||
GatewayUserID string `json:"gatewayUserId,omitempty"`
|
Source string `json:"source,omitempty"`
|
||||||
UserGroupID string `json:"userGroupId,omitempty"`
|
GatewayUserID string `json:"gatewayUserId,omitempty"`
|
||||||
UserGroupKey string `json:"userGroupKey,omitempty"`
|
UserGroupID string `json:"userGroupId,omitempty"`
|
||||||
UserGroupKeys []string `json:"userGroupKeys,omitempty"`
|
UserGroupKey string `json:"userGroupKey,omitempty"`
|
||||||
APIKeyID string `json:"apiKeyId,omitempty"`
|
UserGroupKeys []string `json:"userGroupKeys,omitempty"`
|
||||||
APIKeySecret string `json:"apiKeySecret,omitempty"`
|
APIKeyID string `json:"apiKeyId,omitempty"`
|
||||||
APIKeyName string `json:"apiKeyName,omitempty"`
|
APIKeySecret string `json:"apiKeySecret,omitempty"`
|
||||||
APIKeyPrefix string `json:"apiKeyPrefix,omitempty"`
|
APIKeyName string `json:"apiKeyName,omitempty"`
|
||||||
APIKeyScopes []string `json:"apiKeyScopes,omitempty"`
|
APIKeyPrefix string `json:"apiKeyPrefix,omitempty"`
|
||||||
TokenExpiresAt time.Time `json:"-"`
|
APIKeyScopes []string `json:"apiKeyScopes,omitempty"`
|
||||||
TokenIssuedAt time.Time `json:"-"`
|
TokenExpiresAt time.Time `json:"-"`
|
||||||
Issuer string `json:"-"`
|
TokenIssuedAt time.Time `json:"-"`
|
||||||
TokenPurpose string `json:"-"`
|
Issuer string `json:"-"`
|
||||||
|
ApplicationID string `json:"-"`
|
||||||
|
OIDCClientID string `json:"-"`
|
||||||
|
OIDCUserBindingID string `json:"-"`
|
||||||
|
TokenPurpose string `json:"-"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type contextKey string
|
type contextKey string
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/golang-jwt/jwt/v5"
|
"github.com/golang-jwt/jwt/v5"
|
||||||
|
"github.com/google/uuid"
|
||||||
)
|
)
|
||||||
|
|
||||||
const maxOIDCResponseBytes = 1 << 20
|
const maxOIDCResponseBytes = 1 << 20
|
||||||
@@ -30,6 +31,9 @@ type OIDCConfig struct {
|
|||||||
Issuer string
|
Issuer string
|
||||||
Audience string
|
Audience string
|
||||||
TenantID string
|
TenantID string
|
||||||
|
TenantMode string
|
||||||
|
ApplicationID string
|
||||||
|
ClientID string
|
||||||
RolePrefix string
|
RolePrefix string
|
||||||
RequiredScopes []string
|
RequiredScopes []string
|
||||||
JWKSCacheTTL time.Duration
|
JWKSCacheTTL time.Duration
|
||||||
@@ -44,10 +48,11 @@ type OIDCConfig struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type OIDCSecurityEventIdentity struct {
|
type OIDCSecurityEventIdentity struct {
|
||||||
Issuer string
|
Issuer string
|
||||||
TenantID string
|
ApplicationID string
|
||||||
Subject string
|
TenantID string
|
||||||
IssuedAt time.Time
|
Subject string
|
||||||
|
IssuedAt time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
type OIDCSecurityEventEvaluation struct {
|
type OIDCSecurityEventEvaluation struct {
|
||||||
@@ -91,8 +96,16 @@ func NewOIDCVerifier(config OIDCConfig) (*OIDCVerifier, error) {
|
|||||||
config.Issuer = strings.TrimRight(strings.TrimSpace(config.Issuer), "/")
|
config.Issuer = strings.TrimRight(strings.TrimSpace(config.Issuer), "/")
|
||||||
config.Audience = strings.TrimSpace(config.Audience)
|
config.Audience = strings.TrimSpace(config.Audience)
|
||||||
config.TenantID = strings.TrimSpace(config.TenantID)
|
config.TenantID = strings.TrimSpace(config.TenantID)
|
||||||
|
config.TenantMode = strings.TrimSpace(config.TenantMode)
|
||||||
|
config.ApplicationID = strings.TrimSpace(config.ApplicationID)
|
||||||
|
config.ClientID = strings.TrimSpace(config.ClientID)
|
||||||
config.RolePrefix = strings.TrimSpace(config.RolePrefix)
|
config.RolePrefix = strings.TrimSpace(config.RolePrefix)
|
||||||
if err := validatePublicURL(config.Issuer, config.AppEnv); err != nil || config.Audience == "" || config.TenantID == "" || config.RolePrefix == "" {
|
if config.TenantMode == "" {
|
||||||
|
config.TenantMode = "single_tenant"
|
||||||
|
}
|
||||||
|
validTenantMode := config.TenantMode == "single_tenant" && config.TenantID != "" ||
|
||||||
|
config.TenantMode == "multi_tenant" && config.TenantID == "" && config.ApplicationID != ""
|
||||||
|
if err := validatePublicURL(config.Issuer, config.AppEnv); err != nil || config.Audience == "" || !validTenantMode || config.RolePrefix == "" {
|
||||||
return nil, errors.New("issuer, audience, tenant and role prefix are required")
|
return nil, errors.New("issuer, audience, tenant and role prefix are required")
|
||||||
}
|
}
|
||||||
if config.IntrospectionEnabled && config.IntrospectionCredentialProvider == nil &&
|
if config.IntrospectionEnabled && config.IntrospectionCredentialProvider == nil &&
|
||||||
@@ -166,7 +179,14 @@ func (v *OIDCVerifier) Verify(ctx context.Context, raw string) (*User, error) {
|
|||||||
return nil, oidcUnauthorized(registeredClaimsValidationCategory(err), "signature or registered claims are invalid", err)
|
return nil, oidcUnauthorized(registeredClaimsValidationCategory(err), "signature or registered claims are invalid", err)
|
||||||
}
|
}
|
||||||
claims, ok := token.Claims.(jwt.MapClaims)
|
claims, ok := token.Claims.(jwt.MapClaims)
|
||||||
if !ok || stringClaim(claims, "sub") == "" || stringClaim(claims, "tid") != v.config.TenantID {
|
tenantID := stringClaim(claims, "tid")
|
||||||
|
validTenant := tenantID == v.config.TenantID
|
||||||
|
if v.config.TenantMode == "multi_tenant" {
|
||||||
|
validTenant = uuid.Validate(tenantID) == nil
|
||||||
|
}
|
||||||
|
clientID := stringClaim(claims, "client_id")
|
||||||
|
if !ok || stringClaim(claims, "sub") == "" || !validTenant ||
|
||||||
|
v.config.ClientID != "" && clientID != v.config.ClientID {
|
||||||
return nil, oidcUnauthorized("STABLE_IDENTITY_CLAIMS_INVALID", "stable identity claims are invalid", nil)
|
return nil, oidcUnauthorized("STABLE_IDENTITY_CLAIMS_INVALID", "stable identity claims are invalid", nil)
|
||||||
}
|
}
|
||||||
expiresAt, ok := numericDateClaim(claims["exp"])
|
expiresAt, ok := numericDateClaim(claims["exp"])
|
||||||
@@ -190,7 +210,8 @@ func (v *OIDCVerifier) Verify(ctx context.Context, raw string) (*User, error) {
|
|||||||
}
|
}
|
||||||
if v.config.SecurityEventEvaluator != nil {
|
if v.config.SecurityEventEvaluator != nil {
|
||||||
evaluation, evaluateErr := v.config.SecurityEventEvaluator(ctx, OIDCSecurityEventIdentity{
|
evaluation, evaluateErr := v.config.SecurityEventEvaluator(ctx, OIDCSecurityEventIdentity{
|
||||||
Issuer: v.config.Issuer, TenantID: v.config.TenantID, Subject: stringClaim(claims, "sub"), IssuedAt: issuedAt,
|
Issuer: v.config.Issuer, ApplicationID: v.config.ApplicationID,
|
||||||
|
TenantID: tenantID, Subject: stringClaim(claims, "sub"), IssuedAt: issuedAt,
|
||||||
})
|
})
|
||||||
if evaluateErr != nil {
|
if evaluateErr != nil {
|
||||||
return nil, NewRequestAuthError(http.StatusServiceUnavailable, "OIDC_SECURITY_EVENT_STATE_UNAVAILABLE", "认证撤销状态暂时不可用")
|
return nil, NewRequestAuthError(http.StatusServiceUnavailable, "OIDC_SECURITY_EVENT_STATE_UNAVAILABLE", "认证撤销状态暂时不可用")
|
||||||
@@ -227,7 +248,8 @@ func (v *OIDCVerifier) Verify(ctx context.Context, raw string) (*User, error) {
|
|||||||
}
|
}
|
||||||
return &User{
|
return &User{
|
||||||
ID: stringClaim(claims, "sub"), Username: username, Roles: roles,
|
ID: stringClaim(claims, "sub"), Username: username, Roles: roles,
|
||||||
TenantID: v.config.TenantID, Source: "oidc", TokenExpiresAt: expiresAt, TokenIssuedAt: issuedAt, Issuer: v.config.Issuer,
|
TenantID: tenantID, Source: "oidc", TokenExpiresAt: expiresAt, TokenIssuedAt: issuedAt, Issuer: v.config.Issuer,
|
||||||
|
ApplicationID: v.config.ApplicationID, OIDCClientID: clientID,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/coreos/go-oidc/v3/oidc"
|
"github.com/coreos/go-oidc/v3/oidc"
|
||||||
|
"github.com/google/uuid"
|
||||||
"golang.org/x/oauth2"
|
"golang.org/x/oauth2"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -77,15 +78,22 @@ func (c *OIDCPublicClient) ValidateConfiguration(ctx context.Context) error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OIDCPublicClient) AuthorizationURL(ctx context.Context, state, nonce, pkceVerifier string) (string, error) {
|
func (c *OIDCPublicClient) AuthorizationURL(ctx context.Context, state, nonce, pkceVerifier, tenantHint string) (string, error) {
|
||||||
if strings.TrimSpace(state) == "" || strings.TrimSpace(nonce) == "" || !validPKCEVerifier(pkceVerifier) {
|
if strings.TrimSpace(state) == "" || strings.TrimSpace(nonce) == "" || !validPKCEVerifier(pkceVerifier) {
|
||||||
return "", errors.New("state, nonce and PKCE verifier are required")
|
return "", errors.New("state, nonce and PKCE verifier are required")
|
||||||
}
|
}
|
||||||
|
if strings.TrimSpace(tenantHint) != "" && uuid.Validate(strings.TrimSpace(tenantHint)) != nil {
|
||||||
|
return "", errors.New("tenant hint must be a UUID")
|
||||||
|
}
|
||||||
config, _, err := c.configuration(ctx)
|
config, _, err := c.configuration(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
return config.AuthCodeURL(state, oidc.Nonce(nonce), oauth2.S256ChallengeOption(pkceVerifier)), nil
|
options := []oauth2.AuthCodeOption{oidc.Nonce(nonce), oauth2.S256ChallengeOption(pkceVerifier)}
|
||||||
|
if strings.TrimSpace(tenantHint) != "" {
|
||||||
|
options = append(options, oauth2.SetAuthURLParam("tenant_hint", strings.TrimSpace(tenantHint)))
|
||||||
|
}
|
||||||
|
return config.AuthCodeURL(state, options...), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OIDCPublicClient) ExchangeCode(ctx context.Context, code, verifier string) (OIDCTokenResponse, error) {
|
func (c *OIDCPublicClient) ExchangeCode(ctx context.Context, code, verifier string) (OIDCTokenResponse, error) {
|
||||||
|
|||||||
@@ -67,7 +67,8 @@ func TestOIDCPublicClientUsesAuthorizationCodePKCES256WithoutSecret(t *testing.T
|
|||||||
if err := client.ValidateConfiguration(context.Background()); err != nil {
|
if err := client.ValidateConfiguration(context.Background()); err != nil {
|
||||||
t.Fatalf("ValidateConfiguration() error = %v", err)
|
t.Fatalf("ValidateConfiguration() error = %v", err)
|
||||||
}
|
}
|
||||||
authorizationURL, err := client.AuthorizationURL(context.Background(), "state", "nonce", pkceVerifier)
|
tenantHint := "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa"
|
||||||
|
authorizationURL, err := client.AuthorizationURL(context.Background(), "state", "nonce", pkceVerifier, tenantHint)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
@@ -78,6 +79,9 @@ func TestOIDCPublicClientUsesAuthorizationCodePKCES256WithoutSecret(t *testing.T
|
|||||||
if query.Get("response_type") != "code" || query.Get("code_challenge_method") != "S256" || query.Get("code_challenge") != expectedChallenge {
|
if query.Get("response_type") != "code" || query.Get("code_challenge_method") != "S256" || query.Get("code_challenge") != expectedChallenge {
|
||||||
t.Fatalf("authorization request is not PKCE S256: %v", query)
|
t.Fatalf("authorization request is not PKCE S256: %v", query)
|
||||||
}
|
}
|
||||||
|
if query.Get("tenant_hint") != tenantHint {
|
||||||
|
t.Fatalf("tenant_hint = %q, want %q", query.Get("tenant_hint"), tenantHint)
|
||||||
|
}
|
||||||
if _, err := client.ExchangeCode(context.Background(), "authorization-code", pkceVerifier); err != nil {
|
if _, err := client.ExchangeCode(context.Background(), "authorization-code", pkceVerifier); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import (
|
|||||||
|
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identity"
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/netproxy"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/netproxy"
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/runner"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/runner"
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
@@ -90,8 +91,8 @@ func (s *Server) me(w http.ResponseWriter, r *http.Request) {
|
|||||||
// @Failure 500 {object} ErrorEnvelope
|
// @Failure 500 {object} ErrorEnvelope
|
||||||
// @Router /api/v1/auth/register [post]
|
// @Router /api/v1/auth/register [post]
|
||||||
func (s *Server) register(w http.ResponseWriter, r *http.Request) {
|
func (s *Server) register(w http.ResponseWriter, r *http.Request) {
|
||||||
if !s.localIdentityEnabled() || !s.ordinaryLocalJWTEnabled() {
|
if !s.localRegistrationEnabled() {
|
||||||
writeError(w, http.StatusForbidden, "local registration is disabled")
|
writeError(w, http.StatusForbidden, "统一认证启用后,本地注册已关闭", "LOCAL_REGISTRATION_DISABLED")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
var input store.LocalRegisterInput
|
var input store.LocalRegisterInput
|
||||||
@@ -185,6 +186,12 @@ func (s *Server) localIdentityEnabled() bool {
|
|||||||
return mode == "" || mode == "standalone" || mode == "hybrid"
|
return mode == "" || mode == "standalone" || mode == "hybrid"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Server) localRegistrationEnabled() bool {
|
||||||
|
runtime := s.currentIdentityRuntime()
|
||||||
|
return (runtime == nil || runtime.Revision.State != identity.RevisionActive) &&
|
||||||
|
s.localIdentityEnabled() && s.ordinaryLocalJWTEnabled()
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Server) localLoginAllowed(user store.GatewayUser) bool {
|
func (s *Server) localLoginAllowed(user store.GatewayUser) bool {
|
||||||
for _, role := range user.Roles {
|
for _, role := range user.Roles {
|
||||||
if role == "manager" || role == "admin" {
|
if role == "manager" || role == "admin" {
|
||||||
|
|||||||
@@ -0,0 +1,41 @@
|
|||||||
|
package httpapi
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identity"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRegisterIsStructurallyForbiddenWhenIdentityRevisionIsActive(t *testing.T) {
|
||||||
|
server := &Server{
|
||||||
|
cfg: config.Config{IdentityMode: "hybrid"},
|
||||||
|
identityTestRevision: identity.Revision{
|
||||||
|
State: identity.RevisionActive, Issuer: "https://auth.example.test",
|
||||||
|
LegacyJWTEnabled: true,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
request := httptest.NewRequest(http.MethodPost, "/api/v1/auth/register", strings.NewReader(`{}`))
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
|
||||||
|
server.register(recorder, request)
|
||||||
|
|
||||||
|
if recorder.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("status=%d, want 403", recorder.Code)
|
||||||
|
}
|
||||||
|
var envelope struct {
|
||||||
|
Error struct {
|
||||||
|
Code string `json:"code"`
|
||||||
|
} `json:"error"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if envelope.Error.Code != "LOCAL_REGISTRATION_DISABLED" {
|
||||||
|
t.Fatalf("code=%q", envelope.Error.Code)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -331,7 +331,7 @@ func (s *Server) updateIdentityDraftPolicy(w http.ResponseWriter, r *http.Reques
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
policy := identity.RevisionPolicy{
|
policy := identity.RevisionPolicy{
|
||||||
LocalTenantKey: revision.LocalTenantKey, RolePrefix: revision.RolePrefix, JITEnabled: revision.JITEnabled,
|
TenantMode: revision.TenantMode, LocalTenantKey: revision.LocalTenantKey, RolePrefix: revision.RolePrefix, JITEnabled: revision.JITEnabled,
|
||||||
LegacyJWTEnabled: revision.LegacyJWTEnabled, SessionIdleSeconds: revision.SessionIdleSeconds,
|
LegacyJWTEnabled: revision.LegacyJWTEnabled, SessionIdleSeconds: revision.SessionIdleSeconds,
|
||||||
SessionAbsoluteSeconds: revision.SessionAbsoluteSeconds, SessionRefreshSeconds: revision.SessionRefreshSeconds,
|
SessionAbsoluteSeconds: revision.SessionAbsoluteSeconds, SessionRefreshSeconds: revision.SessionRefreshSeconds,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -15,6 +15,10 @@ type oidcTokenVerifier interface {
|
|||||||
Verify(context.Context, string) (*auth.User, error)
|
Verify(context.Context, string) (*auth.User, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type tenantContextReader interface {
|
||||||
|
Get(context.Context, string, string) (identity.TenantContext, bool, error)
|
||||||
|
}
|
||||||
|
|
||||||
// identityRequestRuntime is an immutable request-level snapshot. A handler that
|
// identityRequestRuntime is an immutable request-level snapshot. A handler that
|
||||||
// starts with one runtime keeps using it even when an administrator activates a
|
// starts with one runtime keeps using it even when an administrator activates a
|
||||||
// new revision while that request is in flight.
|
// new revision while that request is in flight.
|
||||||
@@ -25,6 +29,7 @@ type identityRequestRuntime struct {
|
|||||||
Sessions oidcSessionManager
|
Sessions oidcSessionManager
|
||||||
SessionCipher *oidcsession.Cipher
|
SessionCipher *oidcsession.Cipher
|
||||||
SecurityEvents *ssfreceiver.ConnectionManager
|
SecurityEvents *ssfreceiver.ConnectionManager
|
||||||
|
TenantContext tenantContextReader
|
||||||
CookieSecure bool
|
CookieSecure bool
|
||||||
BrowserEnabled bool
|
BrowserEnabled bool
|
||||||
}
|
}
|
||||||
@@ -38,7 +43,8 @@ func (s *Server) currentIdentityRuntime() *identityRequestRuntime {
|
|||||||
return &identityRequestRuntime{
|
return &identityRequestRuntime{
|
||||||
Revision: runtime.Revision, Verifier: runtime.Verifier, PublicClient: runtime.PublicClient,
|
Revision: runtime.Revision, Verifier: runtime.Verifier, PublicClient: runtime.PublicClient,
|
||||||
Sessions: runtime.Sessions, SessionCipher: runtime.SessionCipher, SecurityEvents: runtime.SecurityEvents,
|
Sessions: runtime.Sessions, SessionCipher: runtime.SessionCipher, SecurityEvents: runtime.SecurityEvents,
|
||||||
CookieSecure: runtime.CookieSecure, BrowserEnabled: runtime.PublicClient != nil,
|
TenantContext: runtime.TenantContext,
|
||||||
|
CookieSecure: runtime.CookieSecure, BrowserEnabled: runtime.PublicClient != nil,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,548 @@
|
|||||||
|
package httpapi
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"crypto/ecdsa"
|
||||||
|
"crypto/elliptic"
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"log/slog"
|
||||||
|
"net/http"
|
||||||
|
"net/http/cookiejar"
|
||||||
|
"net/http/httptest"
|
||||||
|
"net/url"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identity"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identityruntime"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
|
"github.com/golang-jwt/jwt/v5"
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLocalCrossRepositoryMultiTenantGatewayFlow(t *testing.T) {
|
||||||
|
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
|
||||||
|
onboardingContract := strings.TrimSpace(os.Getenv("AUTH_CENTER_ONBOARDING_CONTRACT"))
|
||||||
|
tenantRuntimeContract := strings.TrimSpace(os.Getenv("AUTH_CENTER_TENANT_RUNTIME_CONTRACT"))
|
||||||
|
if databaseURL == "" || onboardingContract == "" || tenantRuntimeContract == "" {
|
||||||
|
t.Skip("set the local Gateway database and Auth Center contract paths to run the cross-repository E2E")
|
||||||
|
}
|
||||||
|
assertAuthCenterMultiTenantContracts(t, onboardingContract, tenantRuntimeContract)
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
applyMigration(t, ctx, databaseURL)
|
||||||
|
db, err := store.Connect(ctx, databaseURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("connect Gateway store: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(db.Close)
|
||||||
|
breakGlassCreated := false
|
||||||
|
var breakGlassUserID string
|
||||||
|
hasBreakGlass, err := db.HasBreakGlassManager(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if !hasBreakGlass {
|
||||||
|
manager, registerErr := db.RegisterLocalUser(ctx, store.LocalRegisterInput{
|
||||||
|
Username: "multi-tenant-e2e-manager-" + uuid.NewString(),
|
||||||
|
Password: uuid.NewString() + "-local-only",
|
||||||
|
})
|
||||||
|
if registerErr != nil {
|
||||||
|
t.Fatalf("create local Break-glass Manager fixture: %v", registerErr)
|
||||||
|
}
|
||||||
|
breakGlassCreated, breakGlassUserID = true, manager.ID
|
||||||
|
}
|
||||||
|
|
||||||
|
applicationID := uuid.NewString()
|
||||||
|
tenantA, tenantB := uuid.NewString(), uuid.NewString()
|
||||||
|
subject := "shared-cross-repository-subject-" + uuid.NewString()
|
||||||
|
browserClientID := "gateway-browser-" + uuid.NewString()
|
||||||
|
machineClientID := "gateway-machine-" + uuid.NewString()
|
||||||
|
machineSecret := "local-cross-repository-machine-secret"
|
||||||
|
audience := "urn:easyai:gateway:" + applicationID
|
||||||
|
|
||||||
|
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var issuer, ssfIssuer, gatewayBaseURL string
|
||||||
|
ssfAudience := "urn:easyai:ssf:receiver:" + applicationID
|
||||||
|
var revokedRefreshTokens atomic.Int64
|
||||||
|
type authorization struct {
|
||||||
|
TenantID string
|
||||||
|
Nonce string
|
||||||
|
}
|
||||||
|
var authorizationMu sync.Mutex
|
||||||
|
authorizations := map[string]authorization{}
|
||||||
|
var ssfStreamMu sync.Mutex
|
||||||
|
var ssfStream map[string]any
|
||||||
|
tenantNames := map[string]string{tenantA: "本地租户 A", tenantB: "本地租户 B"}
|
||||||
|
tenantSlugs := map[string]string{tenantA: "local-tenant-a", tenantB: "local-tenant-b"}
|
||||||
|
|
||||||
|
authCenter := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
switch {
|
||||||
|
case r.URL.Path == "/.well-known/openid-configuration":
|
||||||
|
writeTestJSON(w, map[string]any{
|
||||||
|
"issuer": issuer, "jwks_uri": issuer + "/jwks",
|
||||||
|
"authorization_endpoint": issuer + "/authorize", "token_endpoint": issuer + "/token",
|
||||||
|
"revocation_endpoint": issuer + "/revoke", "end_session_endpoint": issuer + "/logout",
|
||||||
|
"introspection_endpoint": issuer + "/introspect",
|
||||||
|
})
|
||||||
|
case r.URL.Path == "/.well-known/ssf-configuration/ssf":
|
||||||
|
writeTestJSON(w, map[string]any{
|
||||||
|
"spec_version": "1_0", "issuer": ssfIssuer, "jwks_uri": issuer + "/jwks",
|
||||||
|
"configuration_endpoint": issuer + "/ssf/streams",
|
||||||
|
"status_endpoint": issuer + "/ssf/status",
|
||||||
|
"verification_endpoint": issuer + "/ssf/v1/verify",
|
||||||
|
"delivery_methods_supported": []string{"urn:ietf:rfc:8935"},
|
||||||
|
})
|
||||||
|
case r.URL.Path == "/jwks":
|
||||||
|
writeTestJSON(w, map[string]any{"keys": []any{oidcJITECJWK("multi-tenant-key", &key.PublicKey)}})
|
||||||
|
case r.URL.Path == "/authorize":
|
||||||
|
tenantID := r.URL.Query().Get("tenant_hint")
|
||||||
|
if tenantID != tenantA && tenantID != tenantB ||
|
||||||
|
r.URL.Query().Get("client_id") != browserClientID ||
|
||||||
|
r.URL.Query().Get("code_challenge_method") != "S256" {
|
||||||
|
http.Error(w, "authorization request rejected", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
code := uuid.NewString()
|
||||||
|
authorizationMu.Lock()
|
||||||
|
authorizations[code] = authorization{TenantID: tenantID, Nonce: r.URL.Query().Get("nonce")}
|
||||||
|
authorizationMu.Unlock()
|
||||||
|
callback := gatewayBaseURL + "/api/v1/auth/oidc/callback?code=" + url.QueryEscape(code) +
|
||||||
|
"&state=" + url.QueryEscape(r.URL.Query().Get("state"))
|
||||||
|
http.Redirect(w, r, callback, http.StatusSeeOther)
|
||||||
|
case r.URL.Path == "/token":
|
||||||
|
if err := r.ParseForm(); err != nil {
|
||||||
|
http.Error(w, "invalid token request", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if r.Form.Get("grant_type") == "client_credentials" {
|
||||||
|
clientID, secret, ok := r.BasicAuth()
|
||||||
|
if !ok || clientID != machineClientID || secret != machineSecret || r.Form.Get("scope") == "" {
|
||||||
|
http.Error(w, "machine token rejected", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
writeTestJSON(w, map[string]any{
|
||||||
|
"access_token": "opaque-local-machine-token", "token_type": "Bearer", "expires_in": 300,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
code := r.Form.Get("code")
|
||||||
|
authorizationMu.Lock()
|
||||||
|
requested, ok := authorizations[code]
|
||||||
|
delete(authorizations, code)
|
||||||
|
authorizationMu.Unlock()
|
||||||
|
if !ok || r.Form.Get("grant_type") != "authorization_code" ||
|
||||||
|
r.Form.Get("client_id") != browserClientID || r.Form.Get("client_secret") != "" {
|
||||||
|
http.Error(w, "authorization code rejected", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
accessToken := signedMultiTenantE2EToken(
|
||||||
|
t, key, issuer, audience, browserClientID, applicationID, requested.TenantID, subject, "",
|
||||||
|
)
|
||||||
|
idToken := signedMultiTenantE2EToken(
|
||||||
|
t, key, issuer, browserClientID, browserClientID, applicationID, requested.TenantID, subject, requested.Nonce,
|
||||||
|
)
|
||||||
|
writeTestJSON(w, map[string]any{
|
||||||
|
"access_token": accessToken, "refresh_token": "opaque-refresh-" + requested.TenantID,
|
||||||
|
"id_token": idToken, "token_type": "Bearer", "expires_in": 300,
|
||||||
|
})
|
||||||
|
case r.URL.Path == "/introspect":
|
||||||
|
clientID, secret, ok := r.BasicAuth()
|
||||||
|
if !ok || clientID != machineClientID || secret != machineSecret {
|
||||||
|
http.Error(w, "introspection rejected", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
writeTestJSON(w, map[string]any{"active": true})
|
||||||
|
case r.URL.Path == "/revoke":
|
||||||
|
revokedRefreshTokens.Add(1)
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
case r.URL.Path == "/logout":
|
||||||
|
http.Redirect(w, r, r.URL.Query().Get("post_logout_redirect_uri"), http.StatusSeeOther)
|
||||||
|
case strings.HasPrefix(r.URL.Path, "/api/v1/runtime/tenants/"):
|
||||||
|
tenantID := strings.TrimPrefix(r.URL.Path, "/api/v1/runtime/tenants/")
|
||||||
|
if r.Header.Get("Authorization") != "Bearer opaque-local-machine-token" {
|
||||||
|
http.Error(w, "tenant context unauthorized", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
displayName, ok := tenantNames[tenantID]
|
||||||
|
if !ok {
|
||||||
|
http.NotFound(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
etag := `"` + tenantID + `-v1"`
|
||||||
|
if r.Header.Get("If-None-Match") == etag {
|
||||||
|
w.WriteHeader(http.StatusNotModified)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.Header().Set("ETag", etag)
|
||||||
|
writeTestJSON(w, identity.TenantContext{
|
||||||
|
ApplicationID: applicationID, TenantID: tenantID, DisplayName: displayName, Slug: tenantSlugs[tenantID],
|
||||||
|
TenantStatus: "active", TenantApplicationStatus: "active", Version: "1", UpdatedAt: time.Now().UTC(),
|
||||||
|
})
|
||||||
|
case r.URL.Path == "/ssf/streams":
|
||||||
|
if r.Header.Get("Authorization") != "Bearer opaque-local-machine-token" {
|
||||||
|
http.Error(w, "SSF management unauthorized", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ssfStreamMu.Lock()
|
||||||
|
defer ssfStreamMu.Unlock()
|
||||||
|
if r.Method == http.MethodGet {
|
||||||
|
if ssfStream == nil {
|
||||||
|
writeTestJSON(w, []any{})
|
||||||
|
} else {
|
||||||
|
writeTestJSON(w, []any{ssfStream})
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if r.Method != http.MethodPost {
|
||||||
|
http.Error(w, "SSF method rejected", http.StatusMethodNotAllowed)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var input struct {
|
||||||
|
EventsRequested []string `json:"events_requested"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
Delivery struct {
|
||||||
|
Method string `json:"method"`
|
||||||
|
EndpointURL string `json:"endpoint_url"`
|
||||||
|
} `json:"delivery"`
|
||||||
|
}
|
||||||
|
if json.NewDecoder(r.Body).Decode(&input) != nil {
|
||||||
|
http.Error(w, "SSF stream rejected", http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ssfStream = map[string]any{
|
||||||
|
"stream_id": uuid.NewString(), "iss": ssfIssuer, "aud": ssfAudience,
|
||||||
|
"events_requested": input.EventsRequested, "description": input.Description,
|
||||||
|
"delivery": map[string]any{"method": input.Delivery.Method, "endpoint_url": input.Delivery.EndpointURL},
|
||||||
|
}
|
||||||
|
writeTestJSON(w, ssfStream)
|
||||||
|
case r.URL.Path == "/ssf/status":
|
||||||
|
writeTestJSON(w, map[string]any{})
|
||||||
|
case r.URL.Path == "/ssf/v1/verify":
|
||||||
|
w.WriteHeader(http.StatusNoContent)
|
||||||
|
default:
|
||||||
|
http.NotFound(w, r)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer authCenter.Close()
|
||||||
|
issuer = authCenter.URL
|
||||||
|
ssfIssuer = issuer + "/ssf"
|
||||||
|
|
||||||
|
cfg := config.Config{
|
||||||
|
AppEnv: "test", HTTPAddr: ":0", DatabaseURL: databaseURL, IdentityMode: "hybrid",
|
||||||
|
JWTSecret: "local-cross-repository-jwt-secret", IdentitySecretStore: "file",
|
||||||
|
IdentitySecretDir: t.TempDir(), LocalGeneratedStorageDir: t.TempDir(), LocalUploadedStorageDir: t.TempDir(),
|
||||||
|
LocalTempAssetTTLHours: 1, CORSAllowedOrigin: "http://localhost:5178", TaskProgressCallbackEnabled: false,
|
||||||
|
}
|
||||||
|
previous, previousErr := db.ActiveIdentityConfigurationRevision(ctx)
|
||||||
|
if previousErr != nil && !errors.Is(previousErr, identity.ErrRevisionNotFound) {
|
||||||
|
t.Fatal(previousErr)
|
||||||
|
}
|
||||||
|
revision := prepareMultiTenantE2ERevision(
|
||||||
|
t, ctx, db, cfg, issuer, ssfIssuer, ssfAudience,
|
||||||
|
applicationID, audience, browserClientID, machineClientID, machineSecret,
|
||||||
|
)
|
||||||
|
revisionID := revision.ID
|
||||||
|
t.Cleanup(func() {
|
||||||
|
_, _ = db.Pool().Exec(context.Background(), `DELETE FROM gateway_security_event_connections WHERE transmitter_issuer=$1`, ssfIssuer)
|
||||||
|
_, _ = db.Pool().Exec(context.Background(), `DELETE FROM gateway_security_event_stream_state WHERE issuer=$1`, ssfIssuer)
|
||||||
|
restoreOIDCJITIdentityRevision(t, context.Background(), db, previous, previousErr == nil, []string{revisionID})
|
||||||
|
if breakGlassCreated {
|
||||||
|
if err := db.DeleteGatewayUser(context.Background(), breakGlassUserID); err != nil {
|
||||||
|
t.Errorf("delete local Break-glass Manager fixture: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
activeRevision, _, err := db.ActivateIdentityRevision(ctx, revision.ID, revision.Version, "multi-tenant-e2e", "multi-tenant-e2e")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("activate multi-tenant revision: %v", err)
|
||||||
|
}
|
||||||
|
revision = activeRevision
|
||||||
|
gateway := httptest.NewServer(NewServerWithContext(ctx, cfg, db, slog.New(slog.NewTextHandler(io.Discard, nil))))
|
||||||
|
defer gateway.Close()
|
||||||
|
gatewayBaseURL = gateway.URL
|
||||||
|
|
||||||
|
tokenA := signedMultiTenantE2EToken(t, key, issuer, audience, browserClientID, applicationID, tenantA, subject, "")
|
||||||
|
tokenB := signedMultiTenantE2EToken(t, key, issuer, audience, browserClientID, applicationID, tenantB, subject, "")
|
||||||
|
var meA, meB auth.User
|
||||||
|
doOIDCJITJSON(t, gateway.URL, http.MethodGet, "/api/v1/me", tokenA, nil, http.StatusOK, &meA)
|
||||||
|
doOIDCJITJSON(t, gateway.URL, http.MethodGet, "/api/v1/me", tokenB, nil, http.StatusOK, &meB)
|
||||||
|
if meA.ID != subject || meB.ID != subject || meA.GatewayTenantID == meB.GatewayTenantID ||
|
||||||
|
meA.GatewayUserID == meB.GatewayUserID || meA.TenantName != tenantNames[tenantA] || meB.TenantName != tenantNames[tenantB] {
|
||||||
|
t.Fatalf("cross-tenant projection mismatch: A=%+v B=%+v", meA, meB)
|
||||||
|
}
|
||||||
|
|
||||||
|
var createdKey struct {
|
||||||
|
APIKey struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
} `json:"apiKey"`
|
||||||
|
}
|
||||||
|
doOIDCJITJSON(t, gateway.URL, http.MethodPost, "/api/v1/api-keys", tokenA,
|
||||||
|
map[string]any{"name": "multi-tenant-local-e2e"}, http.StatusCreated, &createdKey)
|
||||||
|
var keysB struct {
|
||||||
|
Items []struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
} `json:"items"`
|
||||||
|
}
|
||||||
|
doOIDCJITJSON(t, gateway.URL, http.MethodGet, "/api/v1/api-keys", tokenB, nil, http.StatusOK, &keysB)
|
||||||
|
for _, item := range keysB.Items {
|
||||||
|
if item.ID == createdKey.APIKey.ID {
|
||||||
|
t.Fatal("Tenant A API Key crossed into Tenant B")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
switchJar := newMultiTenantE2ECookieJar(t)
|
||||||
|
loginMultiTenantE2E(t, switchJar, gateway.URL, tenantA)
|
||||||
|
loginMultiTenantE2E(t, switchJar, gateway.URL, tenantB)
|
||||||
|
var switched auth.User
|
||||||
|
doMultiTenantE2ECookieJSON(t, switchJar, gateway.URL+"/api/v1/me", http.StatusOK, &switched)
|
||||||
|
if switched.TenantID != tenantB || revokedRefreshTokens.Load() != 1 {
|
||||||
|
t.Fatalf("tenant switch did not replace A session: me=%+v revoked=%d", switched, revokedRefreshTokens.Load())
|
||||||
|
}
|
||||||
|
var switchedASessions int
|
||||||
|
if err := db.Pool().QueryRow(ctx, `SELECT count(*) FROM gateway_oidc_sessions WHERE gateway_tenant_id=$1::uuid`, meA.GatewayTenantID).Scan(&switchedASessions); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if switchedASessions != 0 {
|
||||||
|
t.Fatalf("tenant switch retained %d old Tenant A sessions", switchedASessions)
|
||||||
|
}
|
||||||
|
|
||||||
|
jarA, jarB := newMultiTenantE2ECookieJar(t), newMultiTenantE2ECookieJar(t)
|
||||||
|
loginMultiTenantE2E(t, jarA, gateway.URL, tenantA)
|
||||||
|
loginMultiTenantE2E(t, jarB, gateway.URL, tenantB)
|
||||||
|
principalResult, err := db.ApplySessionRevoked(ctx, store.ApplySessionRevokedInput{
|
||||||
|
Issuer: issuer + "/ssf", Audience: "urn:easyai:ssf:" + applicationID, JTI: uuid.NewString(),
|
||||||
|
SubjectIssuer: issuer, ApplicationID: applicationID, SubjectType: "principal",
|
||||||
|
TenantID: tenantA, Subject: subject, EventTimestamp: time.Now().UTC().Add(-2 * time.Second), InitiatingEntity: "local-e2e",
|
||||||
|
})
|
||||||
|
if err != nil || principalResult.SessionsDeleted != 1 {
|
||||||
|
t.Fatalf("principal revocation result=%+v err=%v", principalResult, err)
|
||||||
|
}
|
||||||
|
doMultiTenantE2ECookieJSON(t, jarA, gateway.URL+"/api/v1/me", http.StatusUnauthorized, nil)
|
||||||
|
doMultiTenantE2ECookieJSON(t, jarB, gateway.URL+"/api/v1/me", http.StatusOK, nil)
|
||||||
|
|
||||||
|
loginMultiTenantE2E(t, jarA, gateway.URL, tenantA)
|
||||||
|
tenantResult, err := db.ApplySessionRevoked(ctx, store.ApplySessionRevokedInput{
|
||||||
|
Issuer: issuer + "/ssf", Audience: "urn:easyai:ssf:" + applicationID, JTI: uuid.NewString(),
|
||||||
|
SubjectIssuer: issuer, ApplicationID: applicationID, SubjectType: "tenant",
|
||||||
|
TenantID: tenantA, Subject: tenantA, EventTimestamp: time.Now().UTC().Add(-2 * time.Second), InitiatingEntity: "local-e2e",
|
||||||
|
})
|
||||||
|
if err != nil || tenantResult.SessionsDeleted != 1 {
|
||||||
|
t.Fatalf("tenant revocation result=%+v err=%v", tenantResult, err)
|
||||||
|
}
|
||||||
|
doMultiTenantE2ECookieJSON(t, jarA, gateway.URL+"/api/v1/me", http.StatusUnauthorized, nil)
|
||||||
|
doMultiTenantE2ECookieJSON(t, jarB, gateway.URL+"/api/v1/me", http.StatusOK, nil)
|
||||||
|
|
||||||
|
var restored auth.User
|
||||||
|
doOIDCJITJSON(t, gateway.URL, http.MethodGet, "/api/v1/me", tokenA, nil, http.StatusOK, &restored)
|
||||||
|
if restored.GatewayTenantID != meA.GatewayTenantID || restored.GatewayUserID != meA.GatewayUserID {
|
||||||
|
t.Fatalf("reassignment did not reuse projection: before=%+v after=%+v", meA, restored)
|
||||||
|
}
|
||||||
|
var survivingAPIKey int
|
||||||
|
if err := db.Pool().QueryRow(ctx, `SELECT count(*) FROM gateway_api_keys WHERE id=$1::uuid`, createdKey.APIKey.ID).Scan(&survivingAPIKey); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if survivingAPIKey != 1 {
|
||||||
|
t.Fatal("OIDC tenant revocation unexpectedly removed a Gateway API Key")
|
||||||
|
}
|
||||||
|
|
||||||
|
doOIDCJITJSON(t, gateway.URL, http.MethodPost, "/api/v1/auth/register", "",
|
||||||
|
map[string]any{"username": "must-not-register", "password": "not-a-real-password"}, http.StatusForbidden, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertAuthCenterMultiTenantContracts(t *testing.T, onboardingPath, runtimePath string) {
|
||||||
|
t.Helper()
|
||||||
|
for path, required := range map[string][]string{
|
||||||
|
onboardingPath: {"schema_version", "multi_tenant", "tenant_context", "machine_to_machine"},
|
||||||
|
runtimePath: {"/api/v1/runtime/tenants/{tenantId}", "tenant.context.read", "If-None-Match", "ETag"},
|
||||||
|
} {
|
||||||
|
raw, err := os.ReadFile(filepath.Clean(path))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("read Auth Center contract %s: %v", filepath.Base(path), err)
|
||||||
|
}
|
||||||
|
for _, marker := range required {
|
||||||
|
if !bytes.Contains(raw, []byte(marker)) {
|
||||||
|
t.Fatalf("Auth Center contract %s is missing %q", filepath.Base(path), marker)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func prepareMultiTenantE2ERevision(
|
||||||
|
t *testing.T,
|
||||||
|
ctx context.Context,
|
||||||
|
db *store.Store,
|
||||||
|
cfg config.Config,
|
||||||
|
issuer, ssfIssuer, ssfAudience, applicationID, audience, browserClientID, machineClientID, machineSecret string,
|
||||||
|
) identity.Revision {
|
||||||
|
t.Helper()
|
||||||
|
draft, err := identity.NewDraft(identity.PairingInput{
|
||||||
|
AuthCenterURL: issuer, PublicBaseURL: "http://localhost", WebBaseURL: "http://localhost:5178",
|
||||||
|
LegacyJWTEnabled: true,
|
||||||
|
}, "test")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
draft, err = db.CreateIdentityConfigurationRevision(ctx, draft)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
secrets, err := identitySecretStore(cfg)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
machineRef, sessionRef := "multi-tenant-machine-"+draft.ID, "multi-tenant-session-"+draft.ID
|
||||||
|
if err := secrets.Put(ctx, machineRef, []byte(machineSecret)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := secrets.Put(ctx, sessionRef, bytes.Repeat([]byte{9}, 32)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for _, ref := range []string{machineRef, sessionRef} {
|
||||||
|
if err := db.QueueIdentitySecretCleanup(ctx, ref, time.Now().Add(10*time.Minute)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
draft, err = db.ApplyIdentityManifest(ctx, draft.ID, draft.Version, identity.ManifestApplication{
|
||||||
|
Manifest: identity.ManifestV2{
|
||||||
|
SchemaVersion: 2, TenantMode: "multi_tenant", Issuer: issuer, ApplicationID: applicationID,
|
||||||
|
Capabilities: []string{"oidc_login", "api_access", "machine_to_machine", "token_introspection", "session_revocation"},
|
||||||
|
Audience: audience, Scopes: []string{"gateway.access"},
|
||||||
|
Clients: identity.ManifestClients{
|
||||||
|
BrowserLogin: &identity.ManifestClient{ClientID: browserClientID},
|
||||||
|
MachineToMachine: &identity.ManifestClient{ClientID: machineClientID},
|
||||||
|
},
|
||||||
|
TenantContext: &identity.ManifestTenantContext{
|
||||||
|
Endpoint: issuer + "/api/v1/runtime/tenants/{tenantId}",
|
||||||
|
Audience: "urn:easyai:auth-center:tenant-context", Scope: "tenant.context.read",
|
||||||
|
},
|
||||||
|
SecurityEvents: &identity.ManifestSecurityEvents{
|
||||||
|
TransmitterIssuer: ssfIssuer, ConfigurationEndpoint: issuer + "/ssf/streams", Audience: ssfAudience,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
MachineCredentialRef: machineRef, SessionEncryptionKeyRef: sessionRef,
|
||||||
|
TraceID: "multi-tenant-local-e2e", AuditID: "multi-tenant-local-e2e", AppEnv: "test",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
draft, err = db.MarkIdentityRevisionValidated(ctx, draft.ID, draft.Version, "multi-tenant-local-e2e", "multi-tenant-local-e2e")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
builder := identityruntime.NewRuntimeBuilder(ctx, db, secrets, identityruntime.RuntimeBuilderConfig{
|
||||||
|
AppEnv: "test", HeartbeatInterval: time.Hour, StaleAfter: 3 * time.Hour, ClockSkew: time.Minute,
|
||||||
|
}, nil)
|
||||||
|
if err := builder.PrepareSecurityEvents(ctx, draft, []byte(machineSecret)); err != nil {
|
||||||
|
t.Fatalf("prepare application-scoped SSF receiver: %v", err)
|
||||||
|
}
|
||||||
|
return draft
|
||||||
|
}
|
||||||
|
|
||||||
|
func signedMultiTenantE2EToken(
|
||||||
|
t *testing.T,
|
||||||
|
key *ecdsa.PrivateKey,
|
||||||
|
issuer, audience, clientID, applicationID, tenantID, subject, nonce string,
|
||||||
|
) string {
|
||||||
|
t.Helper()
|
||||||
|
now := time.Now().UTC()
|
||||||
|
claims := jwt.MapClaims{
|
||||||
|
"iss": issuer, "aud": audience, "sub": subject, "tid": tenantID, "client_id": clientID,
|
||||||
|
"application_id": applicationID, "preferred_username": "shared-local-user",
|
||||||
|
"roles": []string{"gateway.user"}, "scope": "openid gateway.access",
|
||||||
|
"iat": now.Unix(), "nbf": now.Add(-time.Second).Unix(), "exp": now.Add(time.Hour).Unix(),
|
||||||
|
}
|
||||||
|
if nonce != "" {
|
||||||
|
claims["nonce"] = nonce
|
||||||
|
}
|
||||||
|
token := jwt.NewWithClaims(jwt.SigningMethodES256, claims)
|
||||||
|
token.Header["kid"] = "multi-tenant-key"
|
||||||
|
raw, err := token.SignedString(key)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return raw
|
||||||
|
}
|
||||||
|
|
||||||
|
func newMultiTenantE2ECookieJar(t *testing.T) *cookiejar.Jar {
|
||||||
|
t.Helper()
|
||||||
|
jar, err := cookiejar.New(nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return jar
|
||||||
|
}
|
||||||
|
|
||||||
|
func loginMultiTenantE2E(t *testing.T, jar *cookiejar.Jar, gatewayURL, tenantID string) {
|
||||||
|
t.Helper()
|
||||||
|
gateway, err := url.Parse(gatewayURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
client := &http.Client{Jar: jar, CheckRedirect: func(request *http.Request, via []*http.Request) error {
|
||||||
|
if len(via) > 10 {
|
||||||
|
return errors.New("too many redirects")
|
||||||
|
}
|
||||||
|
if request.URL.Host != gateway.Host && request.URL.Path != "/authorize" {
|
||||||
|
return http.ErrUseLastResponse
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}}
|
||||||
|
response, err := client.Get(gatewayURL + "/api/v1/auth/oidc/login?returnTo=%2F&tenantHint=" + url.QueryEscape(tenantID))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("complete tenant %s login: %v", tenantID[:8], err)
|
||||||
|
}
|
||||||
|
defer response.Body.Close()
|
||||||
|
if response.StatusCode != http.StatusSeeOther && response.StatusCode != http.StatusOK {
|
||||||
|
raw, _ := io.ReadAll(io.LimitReader(response.Body, 4096))
|
||||||
|
t.Fatalf("tenant %s login status=%d body=%s", tenantID[:8], response.StatusCode, raw)
|
||||||
|
}
|
||||||
|
for _, cookie := range jar.Cookies(gateway) {
|
||||||
|
if cookie.Name == auth.OIDCSessionCookieName && cookie.Value != "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
t.Fatalf("tenant %s login returned no opaque session", tenantID[:8])
|
||||||
|
}
|
||||||
|
|
||||||
|
func doMultiTenantE2ECookieJSON(t *testing.T, jar *cookiejar.Jar, requestURL string, expectedStatus int, output any) {
|
||||||
|
t.Helper()
|
||||||
|
client := &http.Client{Jar: jar}
|
||||||
|
response, err := client.Get(requestURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer response.Body.Close()
|
||||||
|
raw, err := io.ReadAll(io.LimitReader(response.Body, 1<<20))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if response.StatusCode != expectedStatus {
|
||||||
|
t.Fatalf("GET %s status=%d want=%d body=%s", requestURL, response.StatusCode, expectedStatus, raw)
|
||||||
|
}
|
||||||
|
if output != nil {
|
||||||
|
if err := json.Unmarshal(raw, output); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeTestJSON(w http.ResponseWriter, value any) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
_ = json.NewEncoder(w).Encode(value)
|
||||||
|
}
|
||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/oidcsession"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/oidcsession"
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
|
"github.com/google/uuid"
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -49,7 +50,12 @@ func (s *Server) startOIDCLogin(w http.ResponseWriter, r *http.Request) {
|
|||||||
if returnTo == "" {
|
if returnTo == "" {
|
||||||
returnTo = "/"
|
returnTo = "/"
|
||||||
}
|
}
|
||||||
transaction, err := oidcsession.NewLoginTransaction(returnTo, time.Now())
|
tenantHint := strings.TrimSpace(r.URL.Query().Get("tenantHint"))
|
||||||
|
if tenantHint != "" && (runtime.Revision.TenantMode != "multi_tenant" || uuid.Validate(tenantHint) != nil) {
|
||||||
|
writeError(w, http.StatusBadRequest, "租户提示无效", errorCodeOIDCLoginInvalid)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
transaction, err := oidcsession.NewLoginTransactionWithTenantHint(returnTo, tenantHint, time.Now())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
writeError(w, http.StatusBadRequest, "登录后返回地址无效", errorCodeOIDCLoginInvalid)
|
writeError(w, http.StatusBadRequest, "登录后返回地址无效", errorCodeOIDCLoginInvalid)
|
||||||
return
|
return
|
||||||
@@ -60,7 +66,9 @@ func (s *Server) startOIDCLogin(w http.ResponseWriter, r *http.Request) {
|
|||||||
writeError(w, http.StatusServiceUnavailable, "登录会话初始化失败,请稍后重试", errorCodeOIDCSessionStoreUnavailable)
|
writeError(w, http.StatusServiceUnavailable, "登录会话初始化失败,请稍后重试", errorCodeOIDCSessionStoreUnavailable)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
authorizationURL, err := runtime.PublicClient.AuthorizationURL(r.Context(), transaction.State, transaction.Nonce, transaction.PKCEVerifier)
|
authorizationURL, err := runtime.PublicClient.AuthorizationURL(
|
||||||
|
r.Context(), transaction.State, transaction.Nonce, transaction.PKCEVerifier, transaction.TenantHint,
|
||||||
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.logger.ErrorContext(r.Context(), "load OIDC authorization endpoint failed", "error", err)
|
s.logger.ErrorContext(r.Context(), "load OIDC authorization endpoint failed", "error", err)
|
||||||
writeError(w, http.StatusServiceUnavailable, "认证中心暂时不可用,请稍后重试", "OIDC_AUTHORIZATION_UNAVAILABLE")
|
writeError(w, http.StatusServiceUnavailable, "认证中心暂时不可用,请稍后重试", "OIDC_AUTHORIZATION_UNAVAILABLE")
|
||||||
@@ -124,7 +132,7 @@ func (s *Server) completeOIDCLogin(w http.ResponseWriter, r *http.Request) {
|
|||||||
s.writeOIDCTokenFailure(w, r, "ID_TOKEN_INVALID", auth.OIDCValidationCategory(err), "认证中心身份令牌校验失败")
|
s.writeOIDCTokenFailure(w, r, "ID_TOKEN_INVALID", auth.OIDCValidationCategory(err), "认证中心身份令牌校验失败")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
projection, err := s.resolveOIDCUserProjectionForRevision(r.Context(), r, identity, runtime.Revision)
|
projection, err := s.resolveOIDCUserProjectionForRuntime(r.Context(), r, identity, runtime)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.writeOIDCCallbackProjectionError(w, r, err)
|
s.writeOIDCCallbackProjectionError(w, r, err)
|
||||||
return
|
return
|
||||||
@@ -140,6 +148,13 @@ func (s *Server) completeOIDCLogin(w http.ResponseWriter, r *http.Request) {
|
|||||||
s.writeOIDCCallbackError(w, r, http.StatusServiceUnavailable, "登录会话保存失败,请稍后重试", errorCodeOIDCSessionStoreUnavailable)
|
s.writeOIDCCallbackError(w, r, http.StatusServiceUnavailable, "登录会话保存失败,请稍后重试", errorCodeOIDCSessionStoreUnavailable)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if err := s.replaceOIDCBrowserSession(r.Context(), r, runtime, rawSession); err != nil {
|
||||||
|
if _, cleanupErr := runtime.Sessions.Delete(r.Context(), rawSession); cleanupErr != nil {
|
||||||
|
s.logger.WarnContext(r.Context(), "cleanup replacement OIDC session failed", "error", cleanupErr)
|
||||||
|
}
|
||||||
|
s.writeOIDCCallbackError(w, r, http.StatusServiceUnavailable, "旧登录会话清理失败,请稍后重试", errorCodeOIDCSessionStoreUnavailable)
|
||||||
|
return
|
||||||
|
}
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
http.SetCookie(w, &http.Cookie{
|
http.SetCookie(w, &http.Cookie{
|
||||||
Name: auth.OIDCSessionCookieName, Value: rawSession, Path: "/",
|
Name: auth.OIDCSessionCookieName, Value: rawSession, Path: "/",
|
||||||
@@ -151,6 +166,34 @@ func (s *Server) completeOIDCLogin(w http.ResponseWriter, r *http.Request) {
|
|||||||
http.Redirect(w, r, oidcReturnLocation(runtime.Revision.WebBaseURL, transaction.ReturnTo), http.StatusSeeOther)
|
http.Redirect(w, r, oidcReturnLocation(runtime.Revision.WebBaseURL, transaction.ReturnTo), http.StatusSeeOther)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (s *Server) replaceOIDCBrowserSession(
|
||||||
|
ctx context.Context,
|
||||||
|
r *http.Request,
|
||||||
|
runtime *identityRequestRuntime,
|
||||||
|
newRawSession string,
|
||||||
|
) error {
|
||||||
|
previous, err := r.Cookie(auth.OIDCSessionCookieName)
|
||||||
|
if errors.Is(err, http.ErrNoCookie) {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if previous.Value == "" || previous.Value == newRawSession {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
bundle, err := runtime.Sessions.Delete(ctx, previous.Value)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if bundle.RefreshToken != "" {
|
||||||
|
if err := runtime.PublicClient.RevokeRefreshToken(ctx, bundle.RefreshToken); err != nil {
|
||||||
|
s.logger.WarnContext(ctx, "revoke replaced OIDC refresh token failed", "error", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
// logoutOIDCSession godoc
|
// logoutOIDCSession godoc
|
||||||
// @Summary 注销 OIDC 登录会话
|
// @Summary 注销 OIDC 登录会话
|
||||||
// @Description 删除 Gateway Session、撤销公共 Client Refresh Token,并跳转认证中心退出地址。
|
// @Description 删除 Gateway Session、撤销公共 Client Refresh Token,并跳转认证中心退出地址。
|
||||||
|
|||||||
@@ -60,6 +60,57 @@ func TestStartOIDCLoginRejectsOpenRedirect(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestStartOIDCLoginEncryptsAndForwardsMultiTenantHint(t *testing.T) {
|
||||||
|
cipher, _ := oidcsession.NewCipher(bytes.Repeat([]byte{3}, 32))
|
||||||
|
client := &fakeOIDCClient{authorizationURL: "https://auth.example.com/authorize"}
|
||||||
|
server := &Server{
|
||||||
|
auth: &auth.Authenticator{OIDCVerifier: &auth.OIDCVerifier{}}, oidcClient: client,
|
||||||
|
oidcSessions: &fakeOIDCSessions{}, oidcSessionCipher: cipher,
|
||||||
|
identityTestRevision: identity.Revision{TenantMode: "multi_tenant"},
|
||||||
|
logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
|
||||||
|
}
|
||||||
|
tenantHint := "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa"
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
server.startOIDCLogin(recorder, httptest.NewRequest(
|
||||||
|
http.MethodGet, "/api/v1/auth/oidc/login?tenantHint="+tenantHint, nil,
|
||||||
|
))
|
||||||
|
if recorder.Code != http.StatusSeeOther || client.tenantHint != tenantHint {
|
||||||
|
t.Fatalf("status=%d forwarded tenantHint=%q", recorder.Code, client.tenantHint)
|
||||||
|
}
|
||||||
|
cookies := recorder.Result().Cookies()
|
||||||
|
transaction, err := cipher.DecodeLoginTransaction(cookies[0].Value, time.Now())
|
||||||
|
if err != nil || transaction.TenantHint != tenantHint {
|
||||||
|
t.Fatalf("transaction=%+v err=%v", transaction, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStartOIDCLoginRejectsInvalidOrSingleTenantHint(t *testing.T) {
|
||||||
|
for _, test := range []struct {
|
||||||
|
name string
|
||||||
|
revision identity.Revision
|
||||||
|
hint string
|
||||||
|
}{
|
||||||
|
{name: "invalid UUID", revision: identity.Revision{TenantMode: "multi_tenant"}, hint: "not-a-uuid"},
|
||||||
|
{name: "single tenant", revision: identity.Revision{TenantMode: "single_tenant"}, hint: "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa"},
|
||||||
|
} {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
cipher, _ := oidcsession.NewCipher(bytes.Repeat([]byte{3}, 32))
|
||||||
|
server := &Server{
|
||||||
|
auth: &auth.Authenticator{OIDCVerifier: &auth.OIDCVerifier{}}, oidcClient: &fakeOIDCClient{},
|
||||||
|
oidcSessions: &fakeOIDCSessions{}, oidcSessionCipher: cipher,
|
||||||
|
identityTestRevision: test.revision, logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
|
||||||
|
}
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
server.startOIDCLogin(recorder, httptest.NewRequest(
|
||||||
|
http.MethodGet, "/api/v1/auth/oidc/login?tenantHint="+test.hint, nil,
|
||||||
|
))
|
||||||
|
if recorder.Code != http.StatusBadRequest || recorder.Header().Get("Set-Cookie") != "" {
|
||||||
|
t.Fatalf("status=%d cookie=%q", recorder.Code, recorder.Header().Get("Set-Cookie"))
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestCompleteOIDCLoginReportsSafeTransactionFailureReason(t *testing.T) {
|
func TestCompleteOIDCLoginReportsSafeTransactionFailureReason(t *testing.T) {
|
||||||
cipher, err := oidcsession.NewCipher(bytes.Repeat([]byte{3}, 32))
|
cipher, err := oidcsession.NewCipher(bytes.Repeat([]byte{3}, 32))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -191,6 +242,42 @@ func TestDeleteOIDCBrowserSessionIsIdempotentAndExpiresCookie(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestReplaceOIDCBrowserSessionDeletesPreviousSessionAndRevokesRefreshToken(t *testing.T) {
|
||||||
|
sessions := &fakeOIDCSessions{
|
||||||
|
deleteBundle: oidcsession.TokenBundle{RefreshToken: "previous-refresh-token"},
|
||||||
|
}
|
||||||
|
client := &fakeOIDCClient{}
|
||||||
|
server := &Server{}
|
||||||
|
runtime := &identityRequestRuntime{Sessions: sessions, PublicClient: client}
|
||||||
|
request := httptest.NewRequest(http.MethodGet, "/api/v1/auth/oidc/callback", nil)
|
||||||
|
request.AddCookie(&http.Cookie{Name: auth.OIDCSessionCookieName, Value: "previous-session"})
|
||||||
|
|
||||||
|
if err := server.replaceOIDCBrowserSession(request.Context(), request, runtime, "new-session"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if sessions.deleted != "previous-session" {
|
||||||
|
t.Fatalf("deleted session=%q", sessions.deleted)
|
||||||
|
}
|
||||||
|
if client.revokedRefreshToken != "previous-refresh-token" {
|
||||||
|
t.Fatalf("revoked refresh token=%q", client.revokedRefreshToken)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReplaceOIDCBrowserSessionDoesNotDeleteNewSession(t *testing.T) {
|
||||||
|
sessions := &fakeOIDCSessions{}
|
||||||
|
server := &Server{}
|
||||||
|
runtime := &identityRequestRuntime{Sessions: sessions, PublicClient: &fakeOIDCClient{}}
|
||||||
|
request := httptest.NewRequest(http.MethodGet, "/api/v1/auth/oidc/callback", nil)
|
||||||
|
request.AddCookie(&http.Cookie{Name: auth.OIDCSessionCookieName, Value: "same-session"})
|
||||||
|
|
||||||
|
if err := server.replaceOIDCBrowserSession(request.Context(), request, runtime, "same-session"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if sessions.deleted != "" {
|
||||||
|
t.Fatalf("unexpected deleted session=%q", sessions.deleted)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestOIDCSessionCSRFAcceptsOnlyAllowedOriginForCookieWrites(t *testing.T) {
|
func TestOIDCSessionCSRFAcceptsOnlyAllowedOriginForCookieWrites(t *testing.T) {
|
||||||
server := &Server{
|
server := &Server{
|
||||||
cfg: config.Config{CORSAllowedOrigin: "https://gateway.example.com"},
|
cfg: config.Config{CORSAllowedOrigin: "https://gateway.example.com"},
|
||||||
@@ -329,10 +416,13 @@ func TestOIDCSessionCSRFIsInactiveWhenOIDCIsDisabled(t *testing.T) {
|
|||||||
type fakeOIDCClient struct {
|
type fakeOIDCClient struct {
|
||||||
authorizationURL string
|
authorizationURL string
|
||||||
state, nonce, challenge string
|
state, nonce, challenge string
|
||||||
|
tenantHint string
|
||||||
|
revokedRefreshToken string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (f *fakeOIDCClient) AuthorizationURL(_ context.Context, state, nonce, challenge string) (string, error) {
|
func (f *fakeOIDCClient) AuthorizationURL(_ context.Context, state, nonce, challenge, tenantHint string) (string, error) {
|
||||||
f.state, f.nonce, f.challenge = state, nonce, challenge
|
f.state, f.nonce, f.challenge = state, nonce, challenge
|
||||||
|
f.tenantHint = tenantHint
|
||||||
return f.authorizationURL, nil
|
return f.authorizationURL, nil
|
||||||
}
|
}
|
||||||
func (f *fakeOIDCClient) ExchangeCode(context.Context, string, string) (auth.OIDCTokenResponse, error) {
|
func (f *fakeOIDCClient) ExchangeCode(context.Context, string, string) (auth.OIDCTokenResponse, error) {
|
||||||
@@ -344,12 +434,18 @@ func (f *fakeOIDCClient) VerifyIDToken(context.Context, string, string) (string,
|
|||||||
func (f *fakeOIDCClient) Refresh(context.Context, string) (auth.OIDCTokenResponse, error) {
|
func (f *fakeOIDCClient) Refresh(context.Context, string) (auth.OIDCTokenResponse, error) {
|
||||||
return auth.OIDCTokenResponse{}, nil
|
return auth.OIDCTokenResponse{}, nil
|
||||||
}
|
}
|
||||||
func (f *fakeOIDCClient) RevokeRefreshToken(context.Context, string) error { return nil }
|
func (f *fakeOIDCClient) RevokeRefreshToken(_ context.Context, refreshToken string) error {
|
||||||
|
f.revokedRefreshToken = refreshToken
|
||||||
|
return nil
|
||||||
|
}
|
||||||
func (f *fakeOIDCClient) EndSessionURL(context.Context, string) (string, error) {
|
func (f *fakeOIDCClient) EndSessionURL(context.Context, string) (string, error) {
|
||||||
return "https://gateway.example.com/", nil
|
return "https://gateway.example.com/", nil
|
||||||
}
|
}
|
||||||
|
|
||||||
type fakeOIDCSessions struct{ deleted string }
|
type fakeOIDCSessions struct {
|
||||||
|
deleted string
|
||||||
|
deleteBundle oidcsession.TokenBundle
|
||||||
|
}
|
||||||
|
|
||||||
func (f *fakeOIDCSessions) Create(context.Context, oidcsession.TokenBundle, *auth.User) (string, error) {
|
func (f *fakeOIDCSessions) Create(context.Context, oidcsession.TokenBundle, *auth.User) (string, error) {
|
||||||
return "opaque-session", nil
|
return "opaque-session", nil
|
||||||
@@ -357,6 +453,6 @@ func (f *fakeOIDCSessions) Create(context.Context, oidcsession.TokenBundle, *aut
|
|||||||
func (f *fakeOIDCSessions) Resolve(context.Context, string) (*auth.User, error) { return nil, nil }
|
func (f *fakeOIDCSessions) Resolve(context.Context, string) (*auth.User, error) { return nil, nil }
|
||||||
func (f *fakeOIDCSessions) Delete(_ context.Context, raw string) (oidcsession.TokenBundle, error) {
|
func (f *fakeOIDCSessions) Delete(_ context.Context, raw string) (oidcsession.TokenBundle, error) {
|
||||||
f.deleted = raw
|
f.deleted = raw
|
||||||
return oidcsession.TokenBundle{}, nil
|
return f.deleteBundle, nil
|
||||||
}
|
}
|
||||||
func (f *fakeOIDCSessions) Cleanup(context.Context) (int64, error) { return 0, nil }
|
func (f *fakeOIDCSessions) Cleanup(context.Context) (int64, error) { return 0, nil }
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"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/identity"
|
||||||
@@ -22,6 +23,10 @@ type oidcUserResolver interface {
|
|||||||
ResolveOrProvisionOIDCUser(context.Context, store.ResolveOrProvisionOIDCUserInput) (store.ResolveOrProvisionOIDCUserResult, error)
|
ResolveOrProvisionOIDCUser(context.Context, store.ResolveOrProvisionOIDCUserInput) (store.ResolveOrProvisionOIDCUserResult, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type oidcTenantBindingReader interface {
|
||||||
|
OIDCTenantBindingContext(context.Context, string, string, string) (store.OIDCTenantBindingContext, error)
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Server) requireUser(permission auth.Permission, next http.Handler) http.Handler {
|
func (s *Server) requireUser(permission auth.Permission, next http.Handler) http.Handler {
|
||||||
return s.auth.Require(permission, s.resolveGatewayUser(next))
|
return s.auth.Require(permission, s.resolveGatewayUser(next))
|
||||||
}
|
}
|
||||||
@@ -57,26 +62,119 @@ func (s *Server) resolveOIDCUserProjection(ctx context.Context, r *http.Request,
|
|||||||
if runtime == nil {
|
if runtime == nil {
|
||||||
return store.ResolveOrProvisionOIDCUserResult{}, errors.New("active identity runtime is unavailable")
|
return store.ResolveOrProvisionOIDCUserResult{}, errors.New("active identity runtime is unavailable")
|
||||||
}
|
}
|
||||||
return s.resolveOIDCUserProjectionForRevision(ctx, r, user, runtime.Revision)
|
return s.resolveOIDCUserProjectionForRuntime(ctx, r, user, runtime)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) resolveOIDCUserProjectionForRevision(ctx context.Context, r *http.Request, user *auth.User, revision identity.Revision) (store.ResolveOrProvisionOIDCUserResult, error) {
|
func (s *Server) resolveOIDCUserProjectionForRuntime(ctx context.Context, r *http.Request, user *auth.User, runtime *identityRequestRuntime) (store.ResolveOrProvisionOIDCUserResult, error) {
|
||||||
if s.oidcUserResolver == nil {
|
if s.oidcUserResolver == nil {
|
||||||
return store.ResolveOrProvisionOIDCUserResult{}, errors.New("OIDC user resolver is unavailable")
|
return store.ResolveOrProvisionOIDCUserResult{}, errors.New("OIDC user resolver is unavailable")
|
||||||
}
|
}
|
||||||
|
revision := runtime.Revision
|
||||||
|
tenantName := ""
|
||||||
|
tenantSlug := ""
|
||||||
|
tenantMetadataStatus := ""
|
||||||
|
tenantMetadataVersion := ""
|
||||||
|
tenantMetadataETag := ""
|
||||||
|
var tenantMetadataUpdatedAt time.Time
|
||||||
|
if revision.TenantMode == "multi_tenant" {
|
||||||
|
if runtime.TenantContext == nil {
|
||||||
|
return store.ResolveOrProvisionOIDCUserResult{}, errors.New("tenant context runtime is unavailable")
|
||||||
|
}
|
||||||
|
var cached store.OIDCTenantBindingContext
|
||||||
|
bindingDisabled := false
|
||||||
|
if reader, ok := s.oidcUserResolver.(oidcTenantBindingReader); ok {
|
||||||
|
cached, _ = reader.OIDCTenantBindingContext(ctx, revision.Issuer, revision.ApplicationID, user.TenantID)
|
||||||
|
bindingDisabled = cached.ID != "" && cached.AccessStatus != "active"
|
||||||
|
if !bindingDisabled && cached.MetadataStatus == "synced" && cached.NextSyncAt.After(time.Now()) {
|
||||||
|
tenantName, tenantSlug = cached.DisplayName, cached.Slug
|
||||||
|
tenantMetadataStatus, tenantMetadataVersion = "synced", cached.Version
|
||||||
|
tenantMetadataETag, tenantMetadataUpdatedAt = cached.ETag, cached.MetadataUpdatedAt
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if tenantMetadataStatus == "synced" {
|
||||||
|
return s.resolveOIDCUserProjectionWithTenantContext(ctx, r, user, revision,
|
||||||
|
tenantName, tenantSlug, tenantMetadataStatus, tenantMetadataVersion, tenantMetadataETag, tenantMetadataUpdatedAt)
|
||||||
|
}
|
||||||
|
etag := cached.ETag
|
||||||
|
if bindingDisabled {
|
||||||
|
// A revoked binding can only be reactivated after a fresh positive
|
||||||
|
// Tenant Context response. Do not let a stale 304 or cached profile
|
||||||
|
// reopen access after reassignment.
|
||||||
|
etag = ""
|
||||||
|
}
|
||||||
|
tenant, unchanged, err := runtime.TenantContext.Get(ctx, user.TenantID, etag)
|
||||||
|
switch {
|
||||||
|
case err == nil && unchanged && !bindingDisabled && cached.MetadataStatus == "synced":
|
||||||
|
tenantName, tenantSlug = cached.DisplayName, cached.Slug
|
||||||
|
tenantMetadataStatus, tenantMetadataVersion = "synced", cached.Version
|
||||||
|
tenantMetadataETag, tenantMetadataUpdatedAt = cached.ETag, cached.MetadataUpdatedAt
|
||||||
|
case err == nil && tenant.Active():
|
||||||
|
tenantName = tenant.DisplayName
|
||||||
|
tenantSlug = tenant.Slug
|
||||||
|
tenantMetadataStatus = "synced"
|
||||||
|
tenantMetadataVersion = tenant.Version
|
||||||
|
tenantMetadataETag = tenant.ETag
|
||||||
|
tenantMetadataUpdatedAt = tenant.UpdatedAt
|
||||||
|
case err == nil, errors.Is(err, identity.ErrTenantContextNotFound):
|
||||||
|
return store.ResolveOrProvisionOIDCUserResult{}, store.ErrOIDCTenantUnavailable
|
||||||
|
case errors.Is(err, identity.ErrTenantContextUnavailable):
|
||||||
|
if bindingDisabled {
|
||||||
|
return store.ResolveOrProvisionOIDCUserResult{}, store.ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
if cached.MetadataStatus == "synced" {
|
||||||
|
tenantName, tenantSlug = cached.DisplayName, cached.Slug
|
||||||
|
tenantMetadataStatus, tenantMetadataVersion = "synced", cached.Version
|
||||||
|
tenantMetadataETag, tenantMetadataUpdatedAt = cached.ETag, cached.MetadataUpdatedAt
|
||||||
|
} else {
|
||||||
|
tenantName = oidcTenantPlaceholderName(user.TenantID)
|
||||||
|
tenantMetadataStatus = "metadata_pending"
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return store.ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return s.resolveOIDCUserProjectionWithTenantContext(ctx, r, user, revision,
|
||||||
|
tenantName, tenantSlug, tenantMetadataStatus, tenantMetadataVersion, tenantMetadataETag, tenantMetadataUpdatedAt)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Server) resolveOIDCUserProjectionWithTenantContext(
|
||||||
|
ctx context.Context,
|
||||||
|
r *http.Request,
|
||||||
|
user *auth.User,
|
||||||
|
revision identity.Revision,
|
||||||
|
tenantName, tenantSlug, tenantMetadataStatus, tenantMetadataVersion, tenantMetadataETag string,
|
||||||
|
tenantMetadataUpdatedAt time.Time,
|
||||||
|
) (store.ResolveOrProvisionOIDCUserResult, error) {
|
||||||
return s.oidcUserResolver.ResolveOrProvisionOIDCUser(ctx, store.ResolveOrProvisionOIDCUserInput{
|
return s.oidcUserResolver.ResolveOrProvisionOIDCUser(ctx, store.ResolveOrProvisionOIDCUserInput{
|
||||||
Issuer: revision.Issuer,
|
Issuer: revision.Issuer,
|
||||||
Subject: user.ID,
|
ApplicationID: revision.ApplicationID,
|
||||||
Username: user.Username,
|
Subject: user.ID,
|
||||||
Roles: user.Roles,
|
Username: user.Username,
|
||||||
TenantID: user.TenantID,
|
Roles: user.Roles,
|
||||||
GatewayTenantKey: revision.LocalTenantKey,
|
TenantID: user.TenantID,
|
||||||
ProvisioningEnabled: revision.JITEnabled,
|
TenantMode: revision.TenantMode,
|
||||||
RequestIP: limitAuditText(requestIP(r), 128),
|
TenantName: tenantName,
|
||||||
UserAgent: limitAuditText(r.UserAgent(), 512),
|
TenantSlug: tenantSlug,
|
||||||
|
TenantMetadataStatus: tenantMetadataStatus,
|
||||||
|
TenantMetadataVersion: tenantMetadataVersion,
|
||||||
|
TenantMetadataETag: tenantMetadataETag,
|
||||||
|
TenantMetadataUpdatedAt: tenantMetadataUpdatedAt,
|
||||||
|
OIDCClientID: user.OIDCClientID,
|
||||||
|
GatewayTenantKey: revision.LocalTenantKey,
|
||||||
|
ProvisioningEnabled: revision.JITEnabled,
|
||||||
|
RequestIP: limitAuditText(requestIP(r), 128),
|
||||||
|
UserAgent: limitAuditText(r.UserAgent(), 512),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func oidcTenantPlaceholderName(tenantID string) string {
|
||||||
|
compact := strings.ReplaceAll(strings.TrimSpace(tenantID), "-", "")
|
||||||
|
if len(compact) > 8 {
|
||||||
|
compact = compact[:8]
|
||||||
|
}
|
||||||
|
return "认证中心租户 " + compact
|
||||||
|
}
|
||||||
|
|
||||||
func (s *Server) writeOIDCUserResolutionError(w http.ResponseWriter, r *http.Request, err error) {
|
func (s *Server) writeOIDCUserResolutionError(w http.ResponseWriter, r *http.Request, err error) {
|
||||||
switch {
|
switch {
|
||||||
case errors.Is(err, store.ErrOIDCUserNotProvisioned):
|
case errors.Is(err, store.ErrOIDCUserNotProvisioned):
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"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/identity"
|
||||||
@@ -16,10 +17,30 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type fakeOIDCUserResolver struct {
|
type fakeOIDCUserResolver struct {
|
||||||
result store.ResolveOrProvisionOIDCUserResult
|
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
|
err error
|
||||||
calls int
|
calls *int
|
||||||
input store.ResolveOrProvisionOIDCUserInput
|
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) {
|
func (f *fakeOIDCUserResolver) ResolveOrProvisionOIDCUser(_ context.Context, input store.ResolveOrProvisionOIDCUserInput) (store.ResolveOrProvisionOIDCUserResult, error) {
|
||||||
@@ -28,6 +49,14 @@ func (f *fakeOIDCUserResolver) ResolveOrProvisionOIDCUser(_ context.Context, inp
|
|||||||
return f.result, f.err
|
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) {
|
func TestResolveGatewayUserAddsLocalOIDCContext(t *testing.T) {
|
||||||
resolver := &fakeOIDCUserResolver{result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{
|
resolver := &fakeOIDCUserResolver{result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{
|
||||||
ID: "platform-user",
|
ID: "platform-user",
|
||||||
@@ -99,6 +128,203 @@ func TestResolveGatewayUserLeavesNonOIDCIdentityChainsUnchanged(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
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) {
|
func TestResolveGatewayUserReturnsStructuredErrors(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ type Server struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type oidcPublicClient interface {
|
type oidcPublicClient interface {
|
||||||
AuthorizationURL(context.Context, string, string, string) (string, error)
|
AuthorizationURL(context.Context, string, string, string, string) (string, error)
|
||||||
ExchangeCode(context.Context, string, string) (auth.OIDCTokenResponse, error)
|
ExchangeCode(context.Context, string, string) (auth.OIDCTokenResponse, error)
|
||||||
VerifyIDToken(context.Context, string, string) (string, error)
|
VerifyIDToken(context.Context, string, string) (string, error)
|
||||||
Refresh(context.Context, string) (auth.OIDCTokenResponse, error)
|
Refresh(context.Context, string) (auth.OIDCTokenResponse, error)
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type RevisionPolicy struct {
|
type RevisionPolicy struct {
|
||||||
|
TenantMode string `json:"tenantMode"`
|
||||||
LocalTenantKey string `json:"localTenantKey"`
|
LocalTenantKey string `json:"localTenantKey"`
|
||||||
RolePrefix string `json:"rolePrefix"`
|
RolePrefix string `json:"rolePrefix"`
|
||||||
JITEnabled bool `json:"jitEnabled"`
|
JITEnabled bool `json:"jitEnabled"`
|
||||||
@@ -32,7 +33,8 @@ type RevisionPolicy struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (policy RevisionPolicy) Validate() error {
|
func (policy RevisionPolicy) Validate() error {
|
||||||
if strings.TrimSpace(policy.LocalTenantKey) == "" || strings.TrimSpace(policy.RolePrefix) == "" {
|
if strings.TrimSpace(policy.RolePrefix) == "" ||
|
||||||
|
policy.TenantMode != "multi_tenant" && strings.TrimSpace(policy.LocalTenantKey) == "" {
|
||||||
return errors.New("local tenant mapping and role prefix are required")
|
return errors.New("local tenant mapping and role prefix are required")
|
||||||
}
|
}
|
||||||
if policy.SessionIdleSeconds <= 0 || policy.SessionAbsoluteSeconds <= policy.SessionIdleSeconds ||
|
if policy.SessionIdleSeconds <= 0 || policy.SessionAbsoluteSeconds <= policy.SessionIdleSeconds ||
|
||||||
@@ -56,6 +58,7 @@ type Revision struct {
|
|||||||
ID string `json:"id"`
|
ID string `json:"id"`
|
||||||
State RevisionState `json:"state"`
|
State RevisionState `json:"state"`
|
||||||
SchemaVersion int `json:"schemaVersion"`
|
SchemaVersion int `json:"schemaVersion"`
|
||||||
|
TenantMode string `json:"tenantMode,omitempty"`
|
||||||
AuthCenterURL string `json:"authCenterUrl"`
|
AuthCenterURL string `json:"authCenterUrl"`
|
||||||
Issuer string `json:"issuer,omitempty"`
|
Issuer string `json:"issuer,omitempty"`
|
||||||
TenantID string `json:"tenantId,omitempty"`
|
TenantID string `json:"tenantId,omitempty"`
|
||||||
@@ -76,6 +79,9 @@ type Revision struct {
|
|||||||
SecurityEventIssuer string `json:"securityEventIssuer,omitempty"`
|
SecurityEventIssuer string `json:"securityEventIssuer,omitempty"`
|
||||||
SecurityEventConfigURL string `json:"securityEventConfigurationUrl,omitempty"`
|
SecurityEventConfigURL string `json:"securityEventConfigurationUrl,omitempty"`
|
||||||
SecurityEventAudience string `json:"securityEventAudience,omitempty"`
|
SecurityEventAudience string `json:"securityEventAudience,omitempty"`
|
||||||
|
TenantContextEndpoint string `json:"tenantContextEndpoint,omitempty"`
|
||||||
|
TenantContextAudience string `json:"tenantContextAudience,omitempty"`
|
||||||
|
TenantContextScope string `json:"tenantContextScope,omitempty"`
|
||||||
MachineCredentialRef string `json:"-"`
|
MachineCredentialRef string `json:"-"`
|
||||||
SessionEncryptionKeyRef string `json:"-"`
|
SessionEncryptionKeyRef string `json:"-"`
|
||||||
SessionIdleSeconds int `json:"sessionIdleSeconds"`
|
SessionIdleSeconds int `json:"sessionIdleSeconds"`
|
||||||
@@ -126,7 +132,7 @@ func NewDraft(input PairingInput, appEnv string) (Revision, error) {
|
|||||||
publicBase, _ := exactBaseURL(input.PublicBaseURL, appEnv)
|
publicBase, _ := exactBaseURL(input.PublicBaseURL, appEnv)
|
||||||
webBase, _ := exactBaseURL(input.WebBaseURL, appEnv)
|
webBase, _ := exactBaseURL(input.WebBaseURL, appEnv)
|
||||||
return Revision{
|
return Revision{
|
||||||
ID: uuid.NewString(), State: RevisionDraft, SchemaVersion: 1,
|
ID: uuid.NewString(), State: RevisionDraft, SchemaVersion: 1, TenantMode: "single_tenant",
|
||||||
AuthCenterURL: authCenter, RolePrefix: "gateway.", LocalTenantKey: strings.TrimSpace(input.LocalTenantKey),
|
AuthCenterURL: authCenter, RolePrefix: "gateway.", LocalTenantKey: strings.TrimSpace(input.LocalTenantKey),
|
||||||
PublicBaseURL: publicBase, WebBaseURL: webBase, JITEnabled: true, LegacyJWTEnabled: input.LegacyJWTEnabled,
|
PublicBaseURL: publicBase, WebBaseURL: webBase, JITEnabled: true, LegacyJWTEnabled: input.LegacyJWTEnabled,
|
||||||
Scopes: []string{}, Capabilities: []string{}, SessionIdleSeconds: 1800, SessionAbsoluteSeconds: 28800,
|
Scopes: []string{}, Capabilities: []string{}, SessionIdleSeconds: 1800, SessionAbsoluteSeconds: 28800,
|
||||||
@@ -151,8 +157,13 @@ func ApplyManifest(revision Revision, input ManifestApplication) (Revision, erro
|
|||||||
if capabilities["oidc_login"] && strings.TrimSpace(input.SessionEncryptionKeyRef) == "" {
|
if capabilities["oidc_login"] && strings.TrimSpace(input.SessionEncryptionKeyRef) == "" {
|
||||||
return Revision{}, errors.New("session encryption key reference is required")
|
return Revision{}, errors.New("session encryption key reference is required")
|
||||||
}
|
}
|
||||||
|
if input.Manifest.SchemaVersion == 1 && strings.TrimSpace(revision.LocalTenantKey) == "" {
|
||||||
|
return Revision{}, ErrLocalTenantInvalid
|
||||||
|
}
|
||||||
|
revision.SchemaVersion = input.Manifest.SchemaVersion
|
||||||
revision.Issuer = strings.TrimRight(input.Manifest.Issuer, "/")
|
revision.Issuer = strings.TrimRight(input.Manifest.Issuer, "/")
|
||||||
revision.TenantID = input.Manifest.TenantID
|
revision.TenantID = input.Manifest.TenantID
|
||||||
|
revision.TenantMode = input.Manifest.TenantMode
|
||||||
revision.ApplicationID = input.Manifest.ApplicationID
|
revision.ApplicationID = input.Manifest.ApplicationID
|
||||||
revision.Audience = input.Manifest.Audience
|
revision.Audience = input.Manifest.Audience
|
||||||
revision.Scopes = append([]string(nil), input.Manifest.Scopes...)
|
revision.Scopes = append([]string(nil), input.Manifest.Scopes...)
|
||||||
@@ -165,6 +176,16 @@ func ApplyManifest(revision Revision, input ManifestApplication) (Revision, erro
|
|||||||
}
|
}
|
||||||
revision.TokenIntrospection = capabilities["token_introspection"]
|
revision.TokenIntrospection = capabilities["token_introspection"]
|
||||||
revision.SessionRevocation = capabilities["session_revocation"]
|
revision.SessionRevocation = capabilities["session_revocation"]
|
||||||
|
if input.Manifest.SchemaVersion == 2 {
|
||||||
|
revision.LocalTenantKey = ""
|
||||||
|
revision.TenantContextEndpoint = input.Manifest.TenantContext.Endpoint
|
||||||
|
revision.TenantContextAudience = input.Manifest.TenantContext.Audience
|
||||||
|
revision.TenantContextScope = input.Manifest.TenantContext.Scope
|
||||||
|
} else {
|
||||||
|
revision.TenantContextEndpoint = ""
|
||||||
|
revision.TenantContextAudience = ""
|
||||||
|
revision.TenantContextScope = ""
|
||||||
|
}
|
||||||
if input.Manifest.SecurityEvents != nil {
|
if input.Manifest.SecurityEvents != nil {
|
||||||
revision.SecurityEventIssuer = input.Manifest.SecurityEvents.TransmitterIssuer
|
revision.SecurityEventIssuer = input.Manifest.SecurityEvents.TransmitterIssuer
|
||||||
revision.SecurityEventConfigURL = input.Manifest.SecurityEvents.ConfigurationEndpoint
|
revision.SecurityEventConfigURL = input.Manifest.SecurityEvents.ConfigurationEndpoint
|
||||||
@@ -209,9 +230,6 @@ func (input PairingInput) ConsumerMetadata(sessionRevocation bool, appEnv string
|
|||||||
if !sameSiteBaseURLs(publicBase, webBase) {
|
if !sameSiteBaseURLs(publicBase, webBase) {
|
||||||
return ConsumerMetadata{}, errors.New("public base URL and web base URL must be same-site")
|
return ConsumerMetadata{}, errors.New("public base URL and web base URL must be same-site")
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(input.LocalTenantKey) == "" {
|
|
||||||
return ConsumerMetadata{}, errors.New("local tenant mapping is required")
|
|
||||||
}
|
|
||||||
metadata := ConsumerMetadata{
|
metadata := ConsumerMetadata{
|
||||||
PublicBaseURL: publicBase, WebBaseURL: webBase,
|
PublicBaseURL: publicBase, WebBaseURL: webBase,
|
||||||
RedirectURIs: []string{publicBase + "/api/v1/auth/oidc/callback"},
|
RedirectURIs: []string{publicBase + "/api/v1/auth/oidc/callback"},
|
||||||
|
|||||||
@@ -176,3 +176,53 @@ func TestNewDraftAllowsLoopbackHTTPOnlyInLocalEnvironments(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestApplyManifestV2RemovesFixedTenantMappingWhileV1StillRequiresIt(t *testing.T) {
|
||||||
|
draft, err := NewDraft(PairingInput{
|
||||||
|
AuthCenterURL: "https://auth.example.com", PublicBaseURL: "https://api.example.com",
|
||||||
|
WebBaseURL: "https://gateway.example.com",
|
||||||
|
}, "production")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("multi-tenant pairing draft rejected: %v", err)
|
||||||
|
}
|
||||||
|
applied, err := ApplyManifest(draft, ManifestApplication{
|
||||||
|
AppEnv: "production", MachineCredentialRef: "identity-machine-example",
|
||||||
|
SessionEncryptionKeyRef: "identity-session-example",
|
||||||
|
Manifest: ManifestV2{
|
||||||
|
SchemaVersion: 2, TenantMode: "multi_tenant", Issuer: "https://auth.example.com/issuer/shared",
|
||||||
|
ApplicationID: "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb",
|
||||||
|
Audience: "urn:easyai:resource:bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb",
|
||||||
|
Capabilities: []string{"oidc_login", "api_access", "machine_to_machine", "token_introspection", "session_revocation"},
|
||||||
|
Scopes: []string{"openid", "gateway.access"},
|
||||||
|
Clients: ManifestClients{
|
||||||
|
BrowserLogin: &ManifestClient{ClientID: "browser"}, MachineToMachine: &ManifestClient{ClientID: "service"},
|
||||||
|
},
|
||||||
|
TenantContext: &ManifestTenantContext{
|
||||||
|
Endpoint: "https://auth.example.com/api/v1/runtime/tenants/{tenantId}",
|
||||||
|
Audience: "urn:easyai:auth-center:tenant-context", Scope: "tenant.context.read",
|
||||||
|
},
|
||||||
|
SecurityEvents: &ManifestSecurityEvents{
|
||||||
|
TransmitterIssuer: "https://auth.example.com/ssf",
|
||||||
|
ConfigurationEndpoint: "https://auth.example.com/.well-known/ssf-configuration/ssf",
|
||||||
|
Audience: "urn:easyai:ssf:receiver:bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if applied.SchemaVersion != 2 || applied.TenantMode != "multi_tenant" || applied.TenantID != "" ||
|
||||||
|
applied.LocalTenantKey != "" || applied.TenantContextEndpoint == "" {
|
||||||
|
t.Fatalf("unexpected multi-tenant revision: %#v", applied)
|
||||||
|
}
|
||||||
|
|
||||||
|
v1Draft := draft
|
||||||
|
v1Draft.ID = "v1-draft"
|
||||||
|
if _, err := ApplyManifest(v1Draft, ManifestApplication{AppEnv: "production", Manifest: ManifestV1{
|
||||||
|
SchemaVersion: 1, Issuer: "https://auth.example.com/issuer/shared",
|
||||||
|
TenantID: "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa", ApplicationID: "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb",
|
||||||
|
Capabilities: []string{}, Scopes: []string{}, Clients: ManifestClients{},
|
||||||
|
}}); err == nil {
|
||||||
|
t.Fatal("Manifest V1 without a fixed local tenant mapping was accepted")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -61,18 +61,31 @@ type ManifestSecurityEvents struct {
|
|||||||
Audience string `json:"audience"`
|
Audience string `json:"audience"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ManifestTenantContext struct {
|
||||||
|
Endpoint string `json:"endpoint"`
|
||||||
|
Audience string `json:"audience"`
|
||||||
|
Scope string `json:"scope"`
|
||||||
|
}
|
||||||
|
|
||||||
type ManifestV1 struct {
|
type ManifestV1 struct {
|
||||||
SchemaVersion int `json:"schema_version"`
|
SchemaVersion int `json:"schema_version"`
|
||||||
|
TenantMode string `json:"tenant_mode,omitempty"`
|
||||||
Issuer string `json:"issuer"`
|
Issuer string `json:"issuer"`
|
||||||
TenantID string `json:"tenant_id"`
|
TenantID string `json:"tenant_id,omitempty"`
|
||||||
ApplicationID string `json:"application_id"`
|
ApplicationID string `json:"application_id"`
|
||||||
Capabilities []string `json:"capabilities"`
|
Capabilities []string `json:"capabilities"`
|
||||||
Audience string `json:"audience,omitempty"`
|
Audience string `json:"audience,omitempty"`
|
||||||
Scopes []string `json:"scopes"`
|
Scopes []string `json:"scopes"`
|
||||||
Clients ManifestClients `json:"clients"`
|
Clients ManifestClients `json:"clients"`
|
||||||
|
TenantContext *ManifestTenantContext `json:"tenant_context,omitempty"`
|
||||||
SecurityEvents *ManifestSecurityEvents `json:"security_events,omitempty"`
|
SecurityEvents *ManifestSecurityEvents `json:"security_events,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ManifestV2 intentionally shares the wire structure with ManifestV1 so the
|
||||||
|
// onboarding delivery can be decoded without guessing a union variant. Validate
|
||||||
|
// always branches on schema_version and rejects fields from the other mode.
|
||||||
|
type ManifestV2 = ManifestV1
|
||||||
|
|
||||||
type MachineCredential struct {
|
type MachineCredential struct {
|
||||||
ClientID string `json:"client_id"`
|
ClientID string `json:"client_id"`
|
||||||
ClientSecret string `json:"client_secret"`
|
ClientSecret string `json:"client_secret"`
|
||||||
@@ -86,11 +99,30 @@ type CredentialDelivery struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (manifest ManifestV1) Validate(appEnv string) error {
|
func (manifest ManifestV1) Validate(appEnv string) error {
|
||||||
if manifest.SchemaVersion != 1 || validatePublicIdentityURL(manifest.Issuer, appEnv) != nil {
|
if validatePublicIdentityURL(manifest.Issuer, appEnv) != nil {
|
||||||
return errors.New("application manifest identity metadata is invalid")
|
return errors.New("application manifest identity metadata is invalid")
|
||||||
}
|
}
|
||||||
if _, err := uuid.Parse(manifest.TenantID); err != nil {
|
switch manifest.SchemaVersion {
|
||||||
return errors.New("application manifest tenant is invalid")
|
case 1:
|
||||||
|
if strings.TrimSpace(manifest.TenantMode) != "" || manifest.TenantContext != nil {
|
||||||
|
return errors.New("application manifest tenant mode is invalid")
|
||||||
|
}
|
||||||
|
if _, err := uuid.Parse(manifest.TenantID); err != nil {
|
||||||
|
return errors.New("application manifest tenant is invalid")
|
||||||
|
}
|
||||||
|
case 2:
|
||||||
|
if manifest.TenantMode != "multi_tenant" || strings.TrimSpace(manifest.TenantID) != "" {
|
||||||
|
return errors.New("application manifest tenant mode is invalid")
|
||||||
|
}
|
||||||
|
if manifest.TenantContext == nil ||
|
||||||
|
validatePublicIdentityURL(strings.ReplaceAll(manifest.TenantContext.Endpoint, "{tenantId}", uuid.Nil.String()), appEnv) != nil ||
|
||||||
|
!strings.Contains(manifest.TenantContext.Endpoint, "{tenantId}") ||
|
||||||
|
manifest.TenantContext.Audience != "urn:easyai:auth-center:tenant-context" ||
|
||||||
|
manifest.TenantContext.Scope != "tenant.context.read" {
|
||||||
|
return errors.New("application manifest tenant context is invalid")
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
return errors.New("application manifest identity metadata is invalid")
|
||||||
}
|
}
|
||||||
if _, err := uuid.Parse(manifest.ApplicationID); err != nil {
|
if _, err := uuid.Parse(manifest.ApplicationID); err != nil {
|
||||||
return errors.New("application manifest application is invalid")
|
return errors.New("application manifest application is invalid")
|
||||||
@@ -107,6 +139,11 @@ func (manifest ManifestV1) Validate(appEnv string) error {
|
|||||||
capabilities["token_introspection"] && !capabilities["machine_to_machine"] {
|
capabilities["token_introspection"] && !capabilities["machine_to_machine"] {
|
||||||
return errors.New("application manifest capability dependencies are invalid")
|
return errors.New("application manifest capability dependencies are invalid")
|
||||||
}
|
}
|
||||||
|
if manifest.SchemaVersion == 2 &&
|
||||||
|
(!capabilities["oidc_login"] || !capabilities["machine_to_machine"] ||
|
||||||
|
!capabilities["token_introspection"] || !capabilities["session_revocation"]) {
|
||||||
|
return errors.New("application manifest multi-tenant capabilities are invalid")
|
||||||
|
}
|
||||||
if capabilities["oidc_login"] && (manifest.Clients.BrowserLogin == nil || strings.TrimSpace(manifest.Clients.BrowserLogin.ClientID) == "") {
|
if capabilities["oidc_login"] && (manifest.Clients.BrowserLogin == nil || strings.TrimSpace(manifest.Clients.BrowserLogin.ClientID) == "") {
|
||||||
return errors.New("application manifest browser client is missing")
|
return errors.New("application manifest browser client is missing")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -78,6 +78,47 @@ func TestManifestV1ValidationRequiresStableFieldsAndCapabilityDependencies(t *te
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestManifestV2ValidationRequiresExplicitMultiTenantRuntimeContract(t *testing.T) {
|
||||||
|
valid := ManifestV2{
|
||||||
|
SchemaVersion: 2, TenantMode: "multi_tenant",
|
||||||
|
Issuer: "https://auth.example.com/issuer/shared", ApplicationID: "bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb",
|
||||||
|
Capabilities: []string{"oidc_login", "api_access", "machine_to_machine", "token_introspection", "session_revocation"},
|
||||||
|
Audience: "urn:easyai:resource:bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb", Scopes: []string{"openid", "gateway.access"},
|
||||||
|
Clients: ManifestClients{
|
||||||
|
BrowserLogin: &ManifestClient{ClientID: "browser"},
|
||||||
|
MachineToMachine: &ManifestClient{ClientID: "service"},
|
||||||
|
},
|
||||||
|
TenantContext: &ManifestTenantContext{
|
||||||
|
Endpoint: "https://auth.example.com/api/v1/runtime/tenants/{tenantId}",
|
||||||
|
Audience: "urn:easyai:auth-center:tenant-context",
|
||||||
|
Scope: "tenant.context.read",
|
||||||
|
},
|
||||||
|
SecurityEvents: &ManifestSecurityEvents{
|
||||||
|
TransmitterIssuer: "https://auth.example.com/ssf",
|
||||||
|
ConfigurationEndpoint: "https://auth.example.com/.well-known/ssf-configuration/ssf",
|
||||||
|
Audience: "urn:easyai:ssf:receiver:bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := valid.Validate("production"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for _, mutate := range []func(*ManifestV2){
|
||||||
|
func(value *ManifestV2) { value.TenantMode = "" },
|
||||||
|
func(value *ManifestV2) { value.TenantID = "aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa" },
|
||||||
|
func(value *ManifestV2) { value.TenantContext = nil },
|
||||||
|
func(value *ManifestV2) { value.TenantContext.Scope = "" },
|
||||||
|
func(value *ManifestV2) { value.TenantContext.Audience = "urn:wrong" },
|
||||||
|
} {
|
||||||
|
candidate := valid
|
||||||
|
contextCopy := *valid.TenantContext
|
||||||
|
candidate.TenantContext = &contextCopy
|
||||||
|
mutate(&candidate)
|
||||||
|
if err := candidate.Validate("production"); err == nil {
|
||||||
|
t.Fatalf("invalid Manifest V2 was accepted: %#v", candidate)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestOnboardingClientRejectsRedirects(t *testing.T) {
|
func TestOnboardingClientRejectsRedirects(t *testing.T) {
|
||||||
target := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
target := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
http.Error(w, "should not be reached", http.StatusTeapot)
|
http.Error(w, "should not be reached", http.StatusTeapot)
|
||||||
|
|||||||
@@ -0,0 +1,188 @@
|
|||||||
|
package identity
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
var (
|
||||||
|
ErrTenantContextNotFound = errors.New("tenant context was not found")
|
||||||
|
ErrTenantContextUnavailable = errors.New("tenant context is unavailable")
|
||||||
|
)
|
||||||
|
|
||||||
|
type TenantContext struct {
|
||||||
|
ApplicationID string `json:"application_id"`
|
||||||
|
TenantID string `json:"tenant_id"`
|
||||||
|
DisplayName string `json:"display_name"`
|
||||||
|
Slug string `json:"slug"`
|
||||||
|
TenantStatus string `json:"tenant_status"`
|
||||||
|
TenantApplicationStatus string `json:"tenant_application_status"`
|
||||||
|
Version string `json:"version"`
|
||||||
|
UpdatedAt time.Time `json:"updated_at"`
|
||||||
|
ETag string `json:"-"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (context TenantContext) Active() bool {
|
||||||
|
return context.TenantStatus == "active" && context.TenantApplicationStatus == "active"
|
||||||
|
}
|
||||||
|
|
||||||
|
type TenantContextCredentialProvider func(context.Context) (string, []byte, error)
|
||||||
|
|
||||||
|
type TenantContextClient struct {
|
||||||
|
issuer string
|
||||||
|
appEnv string
|
||||||
|
application string
|
||||||
|
endpoint string
|
||||||
|
audience string
|
||||||
|
scope string
|
||||||
|
credentials TenantContextCredentialProvider
|
||||||
|
client *http.Client
|
||||||
|
mutex sync.Mutex
|
||||||
|
token string
|
||||||
|
tokenExpiry time.Time
|
||||||
|
tokenURL string
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewTenantContextClient(appEnv, issuer, applicationID, endpoint, audience, scope string, credentials TenantContextCredentialProvider, base *http.Client) (*TenantContextClient, error) {
|
||||||
|
issuer = strings.TrimRight(strings.TrimSpace(issuer), "/")
|
||||||
|
endpoint = strings.TrimSpace(endpoint)
|
||||||
|
applicationID = strings.TrimSpace(applicationID)
|
||||||
|
if validatePublicIdentityURL(issuer, appEnv) != nil ||
|
||||||
|
validatePublicIdentityURL(strings.ReplaceAll(endpoint, "{tenantId}", uuid.Nil.String()), appEnv) != nil ||
|
||||||
|
!strings.Contains(endpoint, "{tenantId}") || uuid.Validate(applicationID) != nil ||
|
||||||
|
audience != "urn:easyai:auth-center:tenant-context" || scope != "tenant.context.read" || credentials == nil {
|
||||||
|
return nil, errors.New("tenant context configuration is invalid")
|
||||||
|
}
|
||||||
|
if base == nil {
|
||||||
|
base = &http.Client{Timeout: 10 * time.Second}
|
||||||
|
}
|
||||||
|
client := *base
|
||||||
|
if client.Timeout <= 0 {
|
||||||
|
client.Timeout = 10 * time.Second
|
||||||
|
}
|
||||||
|
client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }
|
||||||
|
return &TenantContextClient{
|
||||||
|
issuer: issuer, appEnv: appEnv, application: applicationID, endpoint: endpoint,
|
||||||
|
audience: audience, scope: scope, credentials: credentials, client: &client,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (client *TenantContextClient) Get(ctx context.Context, tenantID, etag string) (TenantContext, bool, error) {
|
||||||
|
if uuid.Validate(strings.TrimSpace(tenantID)) != nil {
|
||||||
|
return TenantContext{}, false, ErrTenantContextNotFound
|
||||||
|
}
|
||||||
|
token, err := client.accessToken(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return TenantContext{}, false, ErrTenantContextUnavailable
|
||||||
|
}
|
||||||
|
request, err := http.NewRequestWithContext(ctx, http.MethodGet, strings.ReplaceAll(client.endpoint, "{tenantId}", tenantID), nil)
|
||||||
|
if err != nil {
|
||||||
|
return TenantContext{}, false, ErrTenantContextUnavailable
|
||||||
|
}
|
||||||
|
request.Header.Set("Accept", "application/json")
|
||||||
|
request.Header.Set("Authorization", "Bearer "+token)
|
||||||
|
if strings.TrimSpace(etag) != "" {
|
||||||
|
request.Header.Set("If-None-Match", strings.TrimSpace(etag))
|
||||||
|
}
|
||||||
|
response, err := client.client.Do(request)
|
||||||
|
if err != nil {
|
||||||
|
return TenantContext{}, false, ErrTenantContextUnavailable
|
||||||
|
}
|
||||||
|
defer response.Body.Close()
|
||||||
|
if response.StatusCode == http.StatusNotModified {
|
||||||
|
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, maxOnboardingResponseBytes))
|
||||||
|
return TenantContext{}, true, nil
|
||||||
|
}
|
||||||
|
if response.StatusCode == http.StatusNotFound {
|
||||||
|
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, maxOnboardingResponseBytes))
|
||||||
|
return TenantContext{}, false, ErrTenantContextNotFound
|
||||||
|
}
|
||||||
|
if response.StatusCode != http.StatusOK {
|
||||||
|
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, maxOnboardingResponseBytes))
|
||||||
|
return TenantContext{}, false, ErrTenantContextUnavailable
|
||||||
|
}
|
||||||
|
var output TenantContext
|
||||||
|
decoder := json.NewDecoder(io.LimitReader(response.Body, maxOnboardingResponseBytes))
|
||||||
|
if decoder.Decode(&output) != nil || output.ApplicationID != client.application || output.TenantID != tenantID ||
|
||||||
|
strings.TrimSpace(output.DisplayName) == "" || strings.TrimSpace(output.Slug) == "" || strings.TrimSpace(output.Version) == "" ||
|
||||||
|
output.UpdatedAt.IsZero() {
|
||||||
|
return TenantContext{}, false, ErrTenantContextUnavailable
|
||||||
|
}
|
||||||
|
output.ETag = strings.TrimSpace(response.Header.Get("ETag"))
|
||||||
|
return output, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (client *TenantContextClient) accessToken(ctx context.Context) (string, error) {
|
||||||
|
client.mutex.Lock()
|
||||||
|
defer client.mutex.Unlock()
|
||||||
|
if client.token != "" && time.Now().Add(30*time.Second).Before(client.tokenExpiry) {
|
||||||
|
return client.token, nil
|
||||||
|
}
|
||||||
|
if client.tokenURL == "" {
|
||||||
|
request, err := http.NewRequestWithContext(ctx, http.MethodGet, client.issuer+"/.well-known/openid-configuration", nil)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
request.Header.Set("Accept", "application/json")
|
||||||
|
response, err := client.client.Do(request)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer response.Body.Close()
|
||||||
|
var discovery struct {
|
||||||
|
Issuer string `json:"issuer"`
|
||||||
|
TokenEndpoint string `json:"token_endpoint"`
|
||||||
|
}
|
||||||
|
if response.StatusCode != http.StatusOK ||
|
||||||
|
json.NewDecoder(io.LimitReader(response.Body, maxOnboardingResponseBytes)).Decode(&discovery) != nil ||
|
||||||
|
strings.TrimRight(discovery.Issuer, "/") != client.issuer ||
|
||||||
|
validatePublicIdentityURL(discovery.TokenEndpoint, client.appEnv) != nil {
|
||||||
|
return "", errors.New("tenant context discovery is invalid")
|
||||||
|
}
|
||||||
|
client.tokenURL = discovery.TokenEndpoint
|
||||||
|
}
|
||||||
|
clientID, secret, err := client.credentials(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer clear(secret)
|
||||||
|
if strings.TrimSpace(clientID) == "" || len(secret) < 16 {
|
||||||
|
return "", errors.New("tenant context credential is unavailable")
|
||||||
|
}
|
||||||
|
form := url.Values{"grant_type": {"client_credentials"}, "scope": {client.scope}}
|
||||||
|
request, err := http.NewRequestWithContext(ctx, http.MethodPost, client.tokenURL, strings.NewReader(form.Encode()))
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||||
|
request.Header.Set("Accept", "application/json")
|
||||||
|
request.SetBasicAuth(clientID, string(secret))
|
||||||
|
response, err := client.client.Do(request)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
defer response.Body.Close()
|
||||||
|
var tokenResponse struct {
|
||||||
|
AccessToken string `json:"access_token"`
|
||||||
|
TokenType string `json:"token_type"`
|
||||||
|
ExpiresIn int `json:"expires_in"`
|
||||||
|
}
|
||||||
|
if response.StatusCode != http.StatusOK ||
|
||||||
|
json.NewDecoder(io.LimitReader(response.Body, maxOnboardingResponseBytes)).Decode(&tokenResponse) != nil ||
|
||||||
|
tokenResponse.AccessToken == "" || !strings.EqualFold(tokenResponse.TokenType, "Bearer") ||
|
||||||
|
tokenResponse.ExpiresIn <= 0 {
|
||||||
|
return "", errors.New("tenant context token response is invalid")
|
||||||
|
}
|
||||||
|
client.token = tokenResponse.AccessToken
|
||||||
|
client.tokenExpiry = time.Now().Add(time.Duration(tokenResponse.ExpiresIn) * time.Second)
|
||||||
|
return client.token, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,118 @@
|
|||||||
|
package identity
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestTenantContextClientUsesApplicationMachineTokenAndETag(t *testing.T) {
|
||||||
|
tenantID, applicationID := uuid.NewString(), uuid.NewString()
|
||||||
|
var issuer string
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
switch r.URL.Path {
|
||||||
|
case "/issuer/shared/.well-known/openid-configuration":
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"issuer": issuer, "token_endpoint": serverURL(r) + "/issuer/shared/token",
|
||||||
|
})
|
||||||
|
case "/issuer/shared/token":
|
||||||
|
clientID, secret, ok := r.BasicAuth()
|
||||||
|
if !ok || clientID != "gateway-machine" || secret != "machine-secret-material" ||
|
||||||
|
r.FormValue("grant_type") != "client_credentials" || r.FormValue("scope") != "tenant.context.read" {
|
||||||
|
t.Fatal("tenant context token request was not least-privilege client credentials")
|
||||||
|
}
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"access_token": "opaque-machine-token", "token_type": "Bearer", "expires_in": 300,
|
||||||
|
})
|
||||||
|
case "/api/v1/runtime/tenants/" + tenantID:
|
||||||
|
if r.Header.Get("Authorization") != "Bearer opaque-machine-token" {
|
||||||
|
t.Fatal("tenant context request omitted the machine token")
|
||||||
|
}
|
||||||
|
if r.Header.Get("If-None-Match") == `"tenant-v7"` {
|
||||||
|
w.WriteHeader(http.StatusNotModified)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.Header().Set("ETag", `"tenant-v7"`)
|
||||||
|
_ = json.NewEncoder(w).Encode(TenantContext{
|
||||||
|
ApplicationID: applicationID, TenantID: tenantID, DisplayName: "Tenant A", Slug: "tenant-a",
|
||||||
|
TenantStatus: "active", TenantApplicationStatus: "active", Version: "7",
|
||||||
|
UpdatedAt: time.Unix(1_780_000_000, 0).UTC(),
|
||||||
|
})
|
||||||
|
default:
|
||||||
|
http.NotFound(w, r)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
issuer = server.URL + "/issuer/shared"
|
||||||
|
client, err := NewTenantContextClient(
|
||||||
|
"test", issuer, applicationID, server.URL+"/api/v1/runtime/tenants/{tenantId}",
|
||||||
|
"urn:easyai:auth-center:tenant-context", "tenant.context.read",
|
||||||
|
func(context.Context) (string, []byte, error) {
|
||||||
|
return "gateway-machine", []byte("machine-secret-material"), nil
|
||||||
|
}, server.Client(),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
tenant, unchanged, err := client.Get(context.Background(), tenantID, "")
|
||||||
|
if err != nil || unchanged || tenant.ETag != `"tenant-v7"` || !tenant.Active() {
|
||||||
|
t.Fatalf("tenant=%#v unchanged=%t error=%v", tenant, unchanged, err)
|
||||||
|
}
|
||||||
|
_, unchanged, err = client.Get(context.Background(), tenantID, tenant.ETag)
|
||||||
|
if err != nil || !unchanged {
|
||||||
|
t.Fatalf("etag request unchanged=%t error=%v", unchanged, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTenantContextClientDistinguishesNotFoundFromTemporaryFailure(t *testing.T) {
|
||||||
|
tenantID, applicationID := uuid.NewString(), uuid.NewString()
|
||||||
|
for _, test := range []struct {
|
||||||
|
status int
|
||||||
|
want error
|
||||||
|
}{
|
||||||
|
{status: http.StatusNotFound, want: ErrTenantContextNotFound},
|
||||||
|
{status: http.StatusServiceUnavailable, want: ErrTenantContextUnavailable},
|
||||||
|
} {
|
||||||
|
t.Run(http.StatusText(test.status), func(t *testing.T) {
|
||||||
|
var issuer string
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
switch {
|
||||||
|
case r.URL.Path == "/issuer/shared/.well-known/openid-configuration":
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"issuer": issuer, "token_endpoint": serverURL(r) + "/issuer/shared/token",
|
||||||
|
})
|
||||||
|
case r.URL.Path == "/issuer/shared/token":
|
||||||
|
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||||
|
"access_token": "opaque-machine-token", "token_type": "Bearer", "expires_in": 300,
|
||||||
|
})
|
||||||
|
default:
|
||||||
|
w.WriteHeader(test.status)
|
||||||
|
}
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
issuer = server.URL + "/issuer/shared"
|
||||||
|
client, err := NewTenantContextClient(
|
||||||
|
"test", issuer, applicationID, server.URL+"/api/v1/runtime/tenants/{tenantId}",
|
||||||
|
"urn:easyai:auth-center:tenant-context", "tenant.context.read",
|
||||||
|
func(context.Context) (string, []byte, error) {
|
||||||
|
return "gateway-machine", []byte("machine-secret-material"), nil
|
||||||
|
}, server.Client(),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, _, err := client.Get(context.Background(), tenantID, ""); err != test.want {
|
||||||
|
t.Fatalf("error=%v want=%v", err, test.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func serverURL(request *http.Request) string {
|
||||||
|
return "http://" + request.Host
|
||||||
|
}
|
||||||
@@ -99,14 +99,18 @@ func (builder *RuntimeBuilder) Build(ctx context.Context, revision identity.Revi
|
|||||||
if err := identity.ValidateRevisionURLs(revision, builder.config.AppEnv); err != nil {
|
if err := identity.ValidateRevisionURLs(revision, builder.config.AppEnv); err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if builder.store == nil || builder.secrets == nil || revision.Issuer == "" || revision.TenantID == "" ||
|
if builder.store == nil || builder.secrets == nil || revision.Issuer == "" ||
|
||||||
revision.Audience == "" || revision.RolePrefix == "" {
|
revision.Audience == "" || revision.RolePrefix == "" ||
|
||||||
|
revision.TenantMode != "multi_tenant" && revision.TenantID == "" ||
|
||||||
|
revision.TenantMode == "multi_tenant" && (revision.ApplicationID == "" || revision.TenantContextEndpoint == "") {
|
||||||
return nil, errors.New("identity runtime configuration is incomplete")
|
return nil, errors.New("identity runtime configuration is incomplete")
|
||||||
}
|
}
|
||||||
if exists, err := builder.store.HasActiveTenantKey(ctx, revision.LocalTenantKey); err != nil {
|
if revision.TenantMode != "multi_tenant" {
|
||||||
return nil, err
|
if exists, err := builder.store.HasActiveTenantKey(ctx, revision.LocalTenantKey); err != nil {
|
||||||
} else if !exists {
|
return nil, err
|
||||||
return nil, identity.ErrLocalTenantInvalid
|
} else if !exists {
|
||||||
|
return nil, identity.ErrLocalTenantInvalid
|
||||||
|
}
|
||||||
}
|
}
|
||||||
runtimeCtx, cancel := context.WithCancel(builder.ctx)
|
runtimeCtx, cancel := context.WithCancel(builder.ctx)
|
||||||
runtime := &Runtime{Revision: revision, CookieSecure: secureCookieFor(revision.PublicBaseURL), close: cancel}
|
runtime := &Runtime{Revision: revision, CookieSecure: secureCookieFor(revision.PublicBaseURL), close: cancel}
|
||||||
@@ -151,6 +155,7 @@ func (builder *RuntimeBuilder) Build(ctx context.Context, revision identity.Revi
|
|||||||
verifier, err := auth.NewOIDCVerifier(auth.OIDCConfig{
|
verifier, err := auth.NewOIDCVerifier(auth.OIDCConfig{
|
||||||
AppEnv: builder.config.AppEnv,
|
AppEnv: builder.config.AppEnv,
|
||||||
Issuer: revision.Issuer, Audience: revision.Audience, TenantID: revision.TenantID,
|
Issuer: revision.Issuer, Audience: revision.Audience, TenantID: revision.TenantID,
|
||||||
|
TenantMode: revision.TenantMode, ApplicationID: revision.ApplicationID, ClientID: revision.BrowserClientID,
|
||||||
RolePrefix: revision.RolePrefix, RequiredScopes: append([]string(nil), revision.Scopes...),
|
RolePrefix: revision.RolePrefix, RequiredScopes: append([]string(nil), revision.Scopes...),
|
||||||
JWKSCacheTTL: builder.config.JWKSCacheTTL, IntrospectionEnabled: revision.TokenIntrospection,
|
JWKSCacheTTL: builder.config.JWKSCacheTTL, IntrospectionEnabled: revision.TokenIntrospection,
|
||||||
IntrospectionCredentialProvider: credentialProvider, SecurityEventEvaluator: evaluator,
|
IntrospectionCredentialProvider: credentialProvider, SecurityEventEvaluator: evaluator,
|
||||||
@@ -166,6 +171,23 @@ func (builder *RuntimeBuilder) Build(ctx context.Context, revision identity.Revi
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
runtime.Verifier = verifier
|
runtime.Verifier = verifier
|
||||||
|
if revision.TenantMode == "multi_tenant" {
|
||||||
|
tenantContext, err := identity.NewTenantContextClient(
|
||||||
|
builder.config.AppEnv, revision.Issuer, revision.ApplicationID,
|
||||||
|
revision.TenantContextEndpoint, revision.TenantContextAudience, revision.TenantContextScope,
|
||||||
|
credentialProvider, nil,
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
cancel()
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
runtime.TenantContext = tenantContext
|
||||||
|
runtime.start = func() {
|
||||||
|
go runTenantContextSynchronizer(
|
||||||
|
runtimeCtx, builder.store, tenantContext, revision.Issuer, revision.ApplicationID,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if slices.Contains(revision.Capabilities, "oidc_login") {
|
if slices.Contains(revision.Capabilities, "oidc_login") {
|
||||||
if revision.BrowserClientID == "" || revision.SessionEncryptionKeyRef == "" {
|
if revision.BrowserClientID == "" || revision.SessionEncryptionKeyRef == "" {
|
||||||
@@ -396,6 +418,7 @@ func (builder *RuntimeBuilder) newSecurityEventManager(ctx context.Context, revi
|
|||||||
func (builder *RuntimeBuilder) securityEventManagerConfig(revision identity.Revision, strictRevisionBinding bool) securityevents.ConnectionManagerConfig {
|
func (builder *RuntimeBuilder) securityEventManagerConfig(revision identity.Revision, strictRevisionBinding bool) securityevents.ConnectionManagerConfig {
|
||||||
return securityevents.ConnectionManagerConfig{
|
return securityevents.ConnectionManagerConfig{
|
||||||
AppEnv: builder.config.AppEnv, OIDCEnabled: true, OIDCIssuer: revision.Issuer, OIDCTenantID: revision.TenantID,
|
AppEnv: builder.config.AppEnv, OIDCEnabled: true, OIDCIssuer: revision.Issuer, OIDCTenantID: revision.TenantID,
|
||||||
|
OIDCApplicationID: revision.ApplicationID, OIDCTenantMode: revision.TenantMode,
|
||||||
ManagementClientID: revision.MachineClientID, PublicBaseURL: revision.PublicBaseURL,
|
ManagementClientID: revision.MachineClientID, PublicBaseURL: revision.PublicBaseURL,
|
||||||
ExpectedTransmitterIssuer: strings.TrimRight(strings.TrimSpace(revision.SecurityEventIssuer), "/"),
|
ExpectedTransmitterIssuer: strings.TrimRight(strings.TrimSpace(revision.SecurityEventIssuer), "/"),
|
||||||
ExpectedAudience: revision.SecurityEventAudience,
|
ExpectedAudience: revision.SecurityEventAudience,
|
||||||
|
|||||||
@@ -40,11 +40,20 @@ type Runtime struct {
|
|||||||
Sessions *oidcsession.Service
|
Sessions *oidcsession.Service
|
||||||
SessionCipher *oidcsession.Cipher
|
SessionCipher *oidcsession.Cipher
|
||||||
SecurityEvents *securityevents.ConnectionManager
|
SecurityEvents *securityevents.ConnectionManager
|
||||||
|
TenantContext *identity.TenantContextClient
|
||||||
CookieSecure bool
|
CookieSecure bool
|
||||||
close func()
|
close func()
|
||||||
|
start func()
|
||||||
|
startOnce sync.Once
|
||||||
securityEventDisconnector SecurityEventDisconnector
|
securityEventDisconnector SecurityEventDisconnector
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (runtime *Runtime) Start() {
|
||||||
|
if runtime != nil && runtime.start != nil {
|
||||||
|
runtime.startOnce.Do(runtime.start)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (runtime *Runtime) Close() {
|
func (runtime *Runtime) Close() {
|
||||||
if runtime != nil && runtime.close != nil {
|
if runtime != nil && runtime.close != nil {
|
||||||
runtime.close()
|
runtime.close()
|
||||||
@@ -378,6 +387,7 @@ func (manager *Manager) publishRuntime(candidate *Runtime, revision identity.Rev
|
|||||||
manager.rememberTrustedWebBaseURL(revision.WebBaseURL)
|
manager.rememberTrustedWebBaseURL(revision.WebBaseURL)
|
||||||
manager.legacyJWTAllowed.Store(revision.LegacyJWTEnabled)
|
manager.legacyJWTAllowed.Store(revision.LegacyJWTEnabled)
|
||||||
old := manager.current.Swap(candidate)
|
old := manager.current.Swap(candidate)
|
||||||
|
candidate.Start()
|
||||||
manager.adoptPreparedSecurityEvents(candidate, revision.ID)
|
manager.adoptPreparedSecurityEvents(candidate, revision.ID)
|
||||||
manager.reconciliationRequired.Store(false)
|
manager.reconciliationRequired.Store(false)
|
||||||
retireRuntime(old)
|
retireRuntime(old)
|
||||||
|
|||||||
@@ -0,0 +1,76 @@
|
|||||||
|
package identityruntime
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identity"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
type tenantContextSyncRepository interface {
|
||||||
|
DueOIDCTenantBindingSyncs(context.Context, string, string, int) ([]store.OIDCTenantBindingSyncTarget, error)
|
||||||
|
ApplyOIDCTenantBindingSync(context.Context, string, identity.TenantContext, bool, time.Time) error
|
||||||
|
RejectOIDCTenantBinding(context.Context, string, string, time.Time) error
|
||||||
|
FailOIDCTenantBindingSync(context.Context, string, string, time.Time) error
|
||||||
|
}
|
||||||
|
|
||||||
|
type tenantContextSyncReader interface {
|
||||||
|
Get(context.Context, string, string) (identity.TenantContext, bool, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runTenantContextSynchronizer(ctx context.Context, repository tenantContextSyncRepository, reader tenantContextSyncReader, issuer, applicationID string) {
|
||||||
|
ticker := time.NewTicker(15 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case now := <-ticker.C:
|
||||||
|
_ = synchronizeDueTenantContexts(ctx, repository, reader, issuer, applicationID, now.UTC())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func synchronizeDueTenantContexts(ctx context.Context, repository tenantContextSyncRepository, reader tenantContextSyncReader, issuer, applicationID string, now time.Time) error {
|
||||||
|
targets, err := repository.DueOIDCTenantBindingSyncs(ctx, issuer, applicationID, 50)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for _, target := range targets {
|
||||||
|
tenant, unchanged, readErr := reader.Get(ctx, target.ExternalTenantID, target.ETag)
|
||||||
|
switch {
|
||||||
|
case readErr == nil && unchanged:
|
||||||
|
err = repository.ApplyOIDCTenantBindingSync(ctx, target.ID, identity.TenantContext{}, true, now)
|
||||||
|
case readErr == nil && tenant.Active():
|
||||||
|
err = repository.ApplyOIDCTenantBindingSync(ctx, target.ID, tenant, false, now)
|
||||||
|
case readErr == nil:
|
||||||
|
err = repository.RejectOIDCTenantBinding(ctx, target.ID, "tenant_inactive", now)
|
||||||
|
case errors.Is(readErr, identity.ErrTenantContextNotFound):
|
||||||
|
err = repository.RejectOIDCTenantBinding(ctx, target.ID, "tenant_not_found", now)
|
||||||
|
default:
|
||||||
|
err = repository.FailOIDCTenantBindingSync(
|
||||||
|
ctx, target.ID, "tenant_context_unavailable", now.Add(tenantContextRetryDelay(target.FailureCount)),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if err != nil && !errors.Is(err, store.ErrOIDCTenantUnavailable) {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func tenantContextRetryDelay(failureCount int) time.Duration {
|
||||||
|
if failureCount < 0 {
|
||||||
|
failureCount = 0
|
||||||
|
}
|
||||||
|
delay := 30 * time.Second
|
||||||
|
for index := 0; index < failureCount && delay < 5*time.Minute; index++ {
|
||||||
|
delay *= 2
|
||||||
|
}
|
||||||
|
if delay > 5*time.Minute {
|
||||||
|
return 5 * time.Minute
|
||||||
|
}
|
||||||
|
return delay
|
||||||
|
}
|
||||||
@@ -0,0 +1,113 @@
|
|||||||
|
package identityruntime
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identity"
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
|
||||||
|
)
|
||||||
|
|
||||||
|
type tenantSyncRepository struct {
|
||||||
|
targets []store.OIDCTenantBindingSyncTarget
|
||||||
|
synced []string
|
||||||
|
unchanged []string
|
||||||
|
rejected []string
|
||||||
|
failed []string
|
||||||
|
failureNext time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func (repository *tenantSyncRepository) DueOIDCTenantBindingSyncs(context.Context, string, string, int) ([]store.OIDCTenantBindingSyncTarget, error) {
|
||||||
|
return repository.targets, nil
|
||||||
|
}
|
||||||
|
func (repository *tenantSyncRepository) ApplyOIDCTenantBindingSync(_ context.Context, id string, _ identity.TenantContext, unchanged bool, _ time.Time) error {
|
||||||
|
if unchanged {
|
||||||
|
repository.unchanged = append(repository.unchanged, id)
|
||||||
|
} else {
|
||||||
|
repository.synced = append(repository.synced, id)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (repository *tenantSyncRepository) RejectOIDCTenantBinding(_ context.Context, id, _ string, _ time.Time) error {
|
||||||
|
repository.rejected = append(repository.rejected, id)
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func (repository *tenantSyncRepository) FailOIDCTenantBindingSync(_ context.Context, id, _ string, next time.Time) error {
|
||||||
|
repository.failed = append(repository.failed, id)
|
||||||
|
repository.failureNext = next
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type tenantSyncReader struct {
|
||||||
|
tenants map[string]identity.TenantContext
|
||||||
|
unchanged map[string]bool
|
||||||
|
errors map[string]error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (reader tenantSyncReader) Get(_ context.Context, tenantID, _ string) (identity.TenantContext, bool, error) {
|
||||||
|
return reader.tenants[tenantID], reader.unchanged[tenantID], reader.errors[tenantID]
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSynchronizeDueTenantContextsAppliesETagRejectsAndBacksOff(t *testing.T) {
|
||||||
|
now := time.Unix(1_780_000_000, 0).UTC()
|
||||||
|
repository := &tenantSyncRepository{targets: []store.OIDCTenantBindingSyncTarget{
|
||||||
|
{ID: "binding-a", ExternalTenantID: "tenant-a", ETag: `"a"`, FailureCount: 0},
|
||||||
|
{ID: "binding-b", ExternalTenantID: "tenant-b", ETag: `"b"`, FailureCount: 0},
|
||||||
|
{ID: "binding-c", ExternalTenantID: "tenant-c", FailureCount: 3},
|
||||||
|
{ID: "binding-d", ExternalTenantID: "tenant-d", FailureCount: 0},
|
||||||
|
}}
|
||||||
|
reader := tenantSyncReader{
|
||||||
|
tenants: map[string]identity.TenantContext{
|
||||||
|
"tenant-a": {TenantStatus: "active", TenantApplicationStatus: "active"},
|
||||||
|
"tenant-d": {TenantStatus: "suspended", TenantApplicationStatus: "active"},
|
||||||
|
},
|
||||||
|
unchanged: map[string]bool{"tenant-b": true},
|
||||||
|
errors: map[string]error{
|
||||||
|
"tenant-c": identity.ErrTenantContextUnavailable,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if err := synchronizeDueTenantContexts(context.Background(), repository, reader, "issuer", "application", now); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(repository.synced) != 1 || repository.synced[0] != "binding-a" ||
|
||||||
|
len(repository.unchanged) != 1 || repository.unchanged[0] != "binding-b" ||
|
||||||
|
len(repository.rejected) != 1 || repository.rejected[0] != "binding-d" ||
|
||||||
|
len(repository.failed) != 1 || repository.failed[0] != "binding-c" {
|
||||||
|
t.Fatalf("sync outcomes=%#v", repository)
|
||||||
|
}
|
||||||
|
if repository.failureNext.Sub(now) != 4*time.Minute {
|
||||||
|
t.Fatalf("retry delay=%s", repository.failureNext.Sub(now))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSynchronizeDueTenantContextsTreatsNotFoundAsRevocation(t *testing.T) {
|
||||||
|
repository := &tenantSyncRepository{targets: []store.OIDCTenantBindingSyncTarget{
|
||||||
|
{ID: "binding-a", ExternalTenantID: "tenant-a"},
|
||||||
|
}}
|
||||||
|
reader := tenantSyncReader{
|
||||||
|
tenants: map[string]identity.TenantContext{}, unchanged: map[string]bool{},
|
||||||
|
errors: map[string]error{"tenant-a": identity.ErrTenantContextNotFound},
|
||||||
|
}
|
||||||
|
if err := synchronizeDueTenantContexts(
|
||||||
|
context.Background(), repository, reader, "issuer", "application", time.Now(),
|
||||||
|
); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(repository.rejected) != 1 || len(repository.failed) != 0 {
|
||||||
|
t.Fatalf("not-found outcome=%#v", repository)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTenantContextRetryDelayIsBounded(t *testing.T) {
|
||||||
|
if got := tenantContextRetryDelay(0); got != 30*time.Second {
|
||||||
|
t.Fatalf("first retry=%s", got)
|
||||||
|
}
|
||||||
|
if got := tenantContextRetryDelay(100); got != 5*time.Minute {
|
||||||
|
t.Fatalf("bounded retry=%s", got)
|
||||||
|
}
|
||||||
|
if !errors.Is(identity.ErrTenantContextUnavailable, identity.ErrTenantContextUnavailable) {
|
||||||
|
t.Fatal("sentinel error changed unexpectedly")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -7,6 +7,8 @@ import (
|
|||||||
"net/url"
|
"net/url"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
)
|
)
|
||||||
|
|
||||||
const LoginTransactionCookieName = "easyai_gateway_oidc_login"
|
const LoginTransactionCookieName = "easyai_gateway_oidc_login"
|
||||||
@@ -18,10 +20,15 @@ type LoginTransaction struct {
|
|||||||
Nonce string `json:"nonce"`
|
Nonce string `json:"nonce"`
|
||||||
PKCEVerifier string `json:"pkceVerifier"`
|
PKCEVerifier string `json:"pkceVerifier"`
|
||||||
ReturnTo string `json:"returnTo"`
|
ReturnTo string `json:"returnTo"`
|
||||||
|
TenantHint string `json:"tenantHint,omitempty"`
|
||||||
CreatedAt time.Time `json:"createdAt"`
|
CreatedAt time.Time `json:"createdAt"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewLoginTransaction(returnTo string, now time.Time) (LoginTransaction, error) {
|
func NewLoginTransaction(returnTo string, now time.Time) (LoginTransaction, error) {
|
||||||
|
return NewLoginTransactionWithTenantHint(returnTo, "", now)
|
||||||
|
}
|
||||||
|
|
||||||
|
func NewLoginTransactionWithTenantHint(returnTo, tenantHint string, now time.Time) (LoginTransaction, error) {
|
||||||
if !ValidReturnTo(returnTo) {
|
if !ValidReturnTo(returnTo) {
|
||||||
return LoginTransaction{}, errors.New("returnTo must be a same-origin relative path")
|
return LoginTransaction{}, errors.New("returnTo must be a same-origin relative path")
|
||||||
}
|
}
|
||||||
@@ -37,7 +44,10 @@ func NewLoginTransaction(returnTo string, now time.Time) (LoginTransaction, erro
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return LoginTransaction{}, err
|
return LoginTransaction{}, err
|
||||||
}
|
}
|
||||||
return LoginTransaction{State: state, Nonce: nonce, PKCEVerifier: verifier, ReturnTo: returnTo, CreatedAt: now.UTC()}, nil
|
return LoginTransaction{
|
||||||
|
State: state, Nonce: nonce, PKCEVerifier: verifier,
|
||||||
|
ReturnTo: returnTo, TenantHint: strings.TrimSpace(tenantHint), CreatedAt: now.UTC(),
|
||||||
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Cipher) EncodeLoginTransaction(transaction LoginTransaction) (string, error) {
|
func (c *Cipher) EncodeLoginTransaction(transaction LoginTransaction) (string, error) {
|
||||||
@@ -58,6 +68,7 @@ func (c *Cipher) DecodeLoginTransaction(encoded string, now time.Time) (LoginTra
|
|||||||
return LoginTransaction{}, err
|
return LoginTransaction{}, err
|
||||||
}
|
}
|
||||||
if transaction.State == "" || transaction.Nonce == "" || transaction.PKCEVerifier == "" || !ValidReturnTo(transaction.ReturnTo) ||
|
if transaction.State == "" || transaction.Nonce == "" || transaction.PKCEVerifier == "" || !ValidReturnTo(transaction.ReturnTo) ||
|
||||||
|
transaction.TenantHint != "" && uuid.Validate(transaction.TenantHint) != nil ||
|
||||||
transaction.CreatedAt.IsZero() || now.Before(transaction.CreatedAt.Add(-time.Minute)) || !now.Before(transaction.CreatedAt.Add(10*time.Minute)) {
|
transaction.CreatedAt.IsZero() || now.Before(transaction.CreatedAt.Add(-time.Minute)) || !now.Before(transaction.CreatedAt.Add(10*time.Minute)) {
|
||||||
return LoginTransaction{}, errors.New("OIDC login transaction has expired or is invalid")
|
return LoginTransaction{}, errors.New("OIDC login transaction has expired or is invalid")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -33,6 +33,27 @@ func TestLoginTransactionIsEncryptedAndBounded(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLoginTransactionEncryptsTenantHint(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 28, 12, 0, 0, 0, time.UTC)
|
||||||
|
cipher, _ := NewCipher(bytes.Repeat([]byte{9}, 32))
|
||||||
|
tenantHint := "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa"
|
||||||
|
transaction, err := NewLoginTransactionWithTenantHint("/", tenantHint, now)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
encoded, err := cipher.EncodeLoginTransaction(transaction)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if strings.Contains(encoded, tenantHint) {
|
||||||
|
t.Fatal("login transaction cookie contains plaintext tenant hint")
|
||||||
|
}
|
||||||
|
decoded, err := cipher.DecodeLoginTransaction(encoded, now)
|
||||||
|
if err != nil || decoded.TenantHint != tenantHint {
|
||||||
|
t.Fatalf("decoded transaction=%+v err=%v", decoded, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestValidReturnToRejectsOpenRedirects(t *testing.T) {
|
func TestValidReturnToRejectsOpenRedirects(t *testing.T) {
|
||||||
for _, value := range []string{"https://evil.example", "//evil.example", "/\\evil", "", "workspace"} {
|
for _, value := range []string{"https://evil.example", "//evil.example", "/\\evil", "", "workspace"} {
|
||||||
if ValidReturnTo(value) {
|
if ValidReturnTo(value) {
|
||||||
|
|||||||
@@ -85,6 +85,12 @@ func (s *Service) Create(ctx context.Context, bundle TokenBundle, localUser *aut
|
|||||||
if err != nil || verified == nil || verified.Source != "oidc" || verified.ID == "" || verified.ID != localUser.ID {
|
if err != nil || verified == nil || verified.Source != "oidc" || verified.ID == "" || verified.ID != localUser.ID {
|
||||||
return "", ErrSessionInvalid
|
return "", ErrSessionInvalid
|
||||||
}
|
}
|
||||||
|
if localUser.OIDCUserBindingID != "" &&
|
||||||
|
(localUser.Issuer == "" || localUser.ApplicationID == "" || localUser.TenantID == "" || localUser.OIDCClientID == "" ||
|
||||||
|
verified.Issuer != localUser.Issuer || verified.ApplicationID != localUser.ApplicationID ||
|
||||||
|
verified.TenantID != localUser.TenantID || verified.OIDCClientID != localUser.OIDCClientID) {
|
||||||
|
return "", ErrSessionInvalid
|
||||||
|
}
|
||||||
now := s.now()
|
now := s.now()
|
||||||
if !verified.TokenExpiresAt.After(now) {
|
if !verified.TokenExpiresAt.After(now) {
|
||||||
return "", ErrSessionExpired
|
return "", ErrSessionExpired
|
||||||
@@ -100,6 +106,8 @@ func (s *Service) Create(ctx context.Context, bundle TokenBundle, localUser *aut
|
|||||||
}
|
}
|
||||||
_, err = s.repository.CreateOIDCSession(ctx, store.CreateOIDCSessionInput{
|
_, err = s.repository.CreateOIDCSession(ctx, store.CreateOIDCSessionInput{
|
||||||
SessionTokenHash: hash, GatewayUserID: localUser.GatewayUserID, GatewayTenantID: localUser.GatewayTenantID,
|
SessionTokenHash: hash, GatewayUserID: localUser.GatewayUserID, GatewayTenantID: localUser.GatewayTenantID,
|
||||||
|
OIDCUserBindingID: localUser.OIDCUserBindingID, OIDCClientID: verified.OIDCClientID,
|
||||||
|
Issuer: verified.Issuer, ApplicationID: verified.ApplicationID, TenantID: verified.TenantID,
|
||||||
TokenCiphertext: ciphertext, AccessTokenExpiresAt: verified.TokenExpiresAt,
|
TokenCiphertext: ciphertext, AccessTokenExpiresAt: verified.TokenExpiresAt,
|
||||||
LastSeenAt: now, IdleExpiresAt: now.Add(s.config.IdleTTL), AbsoluteExpiresAt: now.Add(s.config.AbsoluteTTL),
|
LastSeenAt: now, IdleExpiresAt: now.Add(s.config.IdleTTL), AbsoluteExpiresAt: now.Add(s.config.AbsoluteTTL),
|
||||||
})
|
})
|
||||||
@@ -142,6 +150,9 @@ func (s *Service) resolveRecord(ctx context.Context, hash []byte, record store.O
|
|||||||
}
|
}
|
||||||
user, err := s.verifySessionUser(ctx, record, bundle.AccessToken)
|
user, err := s.verifySessionUser(ctx, record, bundle.AccessToken)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
if errors.Is(err, ErrSessionInvalid) {
|
||||||
|
_ = s.repository.DeleteOIDCSessionByID(ctx, record.ID)
|
||||||
|
}
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if err := s.touch(ctx, record, now); err != nil {
|
if err := s.touch(ctx, record, now); err != nil {
|
||||||
@@ -264,6 +275,11 @@ func (s *Service) verifySessionUser(ctx context.Context, record store.OIDCSessio
|
|||||||
if err != nil || user == nil || user.Source != "oidc" || user.ID != record.ExternalUserID {
|
if err != nil || user == nil || user.Source != "oidc" || user.ID != record.ExternalUserID {
|
||||||
return nil, ErrSessionInvalid
|
return nil, ErrSessionInvalid
|
||||||
}
|
}
|
||||||
|
if record.OIDCUserBindingID != "" &&
|
||||||
|
(user.Issuer != record.Issuer || user.ApplicationID != record.ApplicationID ||
|
||||||
|
user.TenantID != record.TenantID || user.OIDCClientID != record.OIDCClientID) {
|
||||||
|
return nil, ErrSessionInvalid
|
||||||
|
}
|
||||||
return user, nil
|
return user, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -255,6 +255,42 @@ func TestServiceDeletesSessionEvenWhenCiphertextCannotBeDecryptedDuringLogout(t
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestServiceBindsMultiTenantSessionToIssuerApplicationTenantAndClient(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 28, 12, 0, 0, 0, time.UTC)
|
||||||
|
tenantA := "aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa"
|
||||||
|
tenantB := "bbbbbbbb-bbbb-4bbb-8bbb-bbbbbbbbbbbb"
|
||||||
|
applicationID := "cccccccc-cccc-4ccc-8ccc-cccccccccccc"
|
||||||
|
verified := &auth.User{
|
||||||
|
ID: "shared-subject", Source: "oidc", Issuer: "https://auth.example.test",
|
||||||
|
ApplicationID: applicationID, TenantID: tenantA, OIDCClientID: "gateway-browser",
|
||||||
|
TokenExpiresAt: now.Add(5 * time.Minute),
|
||||||
|
}
|
||||||
|
verifier := fakeVerifier{users: map[string]*auth.User{"access": verified}}
|
||||||
|
repository := newFakeRepository("shared-subject")
|
||||||
|
service := newTestService(t, repository, verifier, &fakePublicClient{})
|
||||||
|
service.now = func() time.Time { return now }
|
||||||
|
local := &auth.User{
|
||||||
|
ID: "shared-subject", GatewayUserID: "11111111-1111-4111-8111-111111111111",
|
||||||
|
GatewayTenantID: "22222222-2222-4222-8222-222222222222",
|
||||||
|
OIDCUserBindingID: "44444444-4444-4444-8444-444444444444",
|
||||||
|
Issuer: "https://auth.example.test", ApplicationID: applicationID,
|
||||||
|
TenantID: tenantA, OIDCClientID: "gateway-browser",
|
||||||
|
}
|
||||||
|
raw, err := service.Create(context.Background(), TokenBundle{
|
||||||
|
AccessToken: "access", RefreshToken: "refresh",
|
||||||
|
}, local)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Create() error = %v", err)
|
||||||
|
}
|
||||||
|
verified.TenantID = tenantB
|
||||||
|
if _, err := service.Resolve(context.Background(), raw); !errors.Is(err, ErrSessionInvalid) {
|
||||||
|
t.Fatalf("Resolve() error = %v, want ErrSessionInvalid", err)
|
||||||
|
}
|
||||||
|
if !repository.isDeleted() {
|
||||||
|
t.Fatal("identity-binding mismatch did not destroy the session")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func testLocalUser() *auth.User {
|
func testLocalUser() *auth.User {
|
||||||
return &auth.User{
|
return &auth.User{
|
||||||
ID: "subject-1", GatewayUserID: "11111111-1111-4111-8111-111111111111",
|
ID: "subject-1", GatewayUserID: "11111111-1111-4111-8111-111111111111",
|
||||||
@@ -323,6 +359,8 @@ func (f *fakeRepository) CreateOIDCSession(_ context.Context, input store.Create
|
|||||||
f.record = store.OIDCSession{
|
f.record = store.OIDCSession{
|
||||||
ID: "33333333-3333-4333-8333-333333333333", SessionTokenHash: append([]byte(nil), input.SessionTokenHash...),
|
ID: "33333333-3333-4333-8333-333333333333", SessionTokenHash: append([]byte(nil), input.SessionTokenHash...),
|
||||||
GatewayUserID: input.GatewayUserID, GatewayTenantID: input.GatewayTenantID, ExternalUserID: f.external,
|
GatewayUserID: input.GatewayUserID, GatewayTenantID: input.GatewayTenantID, ExternalUserID: f.external,
|
||||||
|
OIDCUserBindingID: input.OIDCUserBindingID, OIDCClientID: input.OIDCClientID,
|
||||||
|
Issuer: input.Issuer, ApplicationID: input.ApplicationID, TenantID: input.TenantID,
|
||||||
UserStatus: "active", TokenCiphertext: append([]byte(nil), input.TokenCiphertext...), AccessTokenExpiresAt: input.AccessTokenExpiresAt,
|
UserStatus: "active", TokenCiphertext: append([]byte(nil), input.TokenCiphertext...), AccessTokenExpiresAt: input.AccessTokenExpiresAt,
|
||||||
LastSeenAt: input.LastSeenAt, IdleExpiresAt: input.IdleExpiresAt, AbsoluteExpiresAt: input.AbsoluteExpiresAt, RefreshVersion: 1,
|
LastSeenAt: input.LastSeenAt, IdleExpiresAt: input.IdleExpiresAt, AbsoluteExpiresAt: input.AbsoluteExpiresAt, RefreshVersion: 1,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -50,6 +50,8 @@ type ConnectionManagerConfig struct {
|
|||||||
OIDCEnabled bool
|
OIDCEnabled bool
|
||||||
OIDCIssuer string
|
OIDCIssuer string
|
||||||
OIDCTenantID string
|
OIDCTenantID string
|
||||||
|
OIDCApplicationID string
|
||||||
|
OIDCTenantMode string
|
||||||
ManagementClientID string
|
ManagementClientID string
|
||||||
ManagementClientSecret string
|
ManagementClientSecret string
|
||||||
ExpectedTransmitterIssuer string
|
ExpectedTransmitterIssuer string
|
||||||
@@ -145,7 +147,9 @@ func (err safeConnectionError) Unwrap() error { return err.cause }
|
|||||||
func (err safeConnectionError) SafeErrorCategory() string { return err.category }
|
func (err safeConnectionError) SafeErrorCategory() string { return err.category }
|
||||||
|
|
||||||
func NewConnectionManager(ctx context.Context, repository ConnectionRepository, secrets SecretStore, config ConnectionManagerConfig, metrics *Metrics) (*ConnectionManager, error) {
|
func NewConnectionManager(ctx context.Context, repository ConnectionRepository, secrets SecretStore, config ConnectionManagerConfig, metrics *Metrics) (*ConnectionManager, error) {
|
||||||
if repository == nil || secrets == nil || !config.OIDCEnabled || config.OIDCIssuer == "" || config.OIDCTenantID == "" {
|
validIdentityBinding := config.OIDCTenantMode == "multi_tenant" && uuid.Validate(config.OIDCApplicationID) == nil ||
|
||||||
|
config.OIDCTenantMode != "multi_tenant" && uuid.Validate(config.OIDCTenantID) == nil
|
||||||
|
if repository == nil || secrets == nil || !config.OIDCEnabled || config.OIDCIssuer == "" || !validIdentityBinding {
|
||||||
return nil, errors.New("security event connection prerequisites are incomplete")
|
return nil, errors.New("security event connection prerequisites are incomplete")
|
||||||
}
|
}
|
||||||
if config.HeartbeatInterval <= 0 {
|
if config.HeartbeatInterval <= 0 {
|
||||||
@@ -773,7 +777,8 @@ func (m *ConnectionManager) Evaluate(ctx context.Context, identity auth.OIDCSecu
|
|||||||
return auth.OIDCSecurityEventEvaluation{Enabled: true, RequireIntrospection: true}, nil
|
return auth.OIDCSecurityEventEvaluation{Enabled: true, RequireIntrospection: true}, nil
|
||||||
}
|
}
|
||||||
result, evaluateErr := m.repository.EvaluateOIDCSecurityEvent(
|
result, evaluateErr := m.repository.EvaluateOIDCSecurityEvent(
|
||||||
ctx, connection.TransmitterIssuer, *connection.Audience, identity.Issuer, identity.TenantID, identity.Subject,
|
ctx, connection.TransmitterIssuer, *connection.Audience, identity.Issuer, identity.ApplicationID,
|
||||||
|
identity.TenantID, identity.Subject,
|
||||||
identity.IssuedAt, time.Now().UTC(), m.config.StaleAfter,
|
identity.IssuedAt, time.Now().UTC(), m.config.StaleAfter,
|
||||||
)
|
)
|
||||||
if result.Revoked && m.metrics != nil {
|
if result.Revoked && m.metrics != nil {
|
||||||
@@ -821,7 +826,9 @@ func (m *ConnectionManager) activate(ctx context.Context, connection store.Secur
|
|||||||
}
|
}
|
||||||
verifier, err := NewVerifier(VerifierConfig{
|
verifier, err := NewVerifier(VerifierConfig{
|
||||||
TransmitterIssuer: connection.TransmitterIssuer, Audience: *connection.Audience,
|
TransmitterIssuer: connection.TransmitterIssuer, Audience: *connection.Audience,
|
||||||
SubjectIssuer: m.config.OIDCIssuer, TenantID: m.config.OIDCTenantID, StreamID: *connection.StreamID,
|
SubjectIssuer: m.config.OIDCIssuer, TenantID: m.config.OIDCTenantID,
|
||||||
|
ApplicationID: m.config.OIDCApplicationID, TenantMode: m.config.OIDCTenantMode,
|
||||||
|
StreamID: *connection.StreamID,
|
||||||
ClockSkew: m.config.ClockSkew, HTTPClient: safeClient,
|
ClockSkew: m.config.ClockSkew, HTTPClient: safeClient,
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -1038,7 +1045,11 @@ func (m *ConnectionManager) validatePrerequisites(issuer, managementClientID str
|
|||||||
if !m.config.OIDCEnabled || strings.TrimSpace(managementClientID) == "" || len(managementSecret) < 16 {
|
if !m.config.OIDCEnabled || strings.TrimSpace(managementClientID) == "" || len(managementSecret) < 16 {
|
||||||
return errors.New("OIDC and RFC 7662 machine client must be configured")
|
return errors.New("OIDC and RFC 7662 machine client must be configured")
|
||||||
}
|
}
|
||||||
if _, err := uuid.Parse(m.config.OIDCTenantID); err != nil {
|
if m.config.OIDCTenantMode == "multi_tenant" {
|
||||||
|
if uuid.Validate(m.config.OIDCApplicationID) != nil {
|
||||||
|
return errors.New("OIDC application id is invalid")
|
||||||
|
}
|
||||||
|
} else if uuid.Validate(m.config.OIDCTenantID) != nil {
|
||||||
return errors.New("OIDC tenant id is invalid")
|
return errors.New("OIDC tenant id is invalid")
|
||||||
}
|
}
|
||||||
if _, err := validatedIssuerURL(issuer, m.config.AppEnv); err != nil {
|
if _, err := validatedIssuerURL(issuer, m.config.AppEnv); err != nil {
|
||||||
|
|||||||
@@ -123,7 +123,7 @@ func (*memoryConnectionRepository) RecordSecurityEventHeartbeatFailure(context.C
|
|||||||
func (*memoryConnectionRepository) AdvanceSecurityEventStreamState(context.Context, string, string, time.Time, time.Duration) error {
|
func (*memoryConnectionRepository) AdvanceSecurityEventStreamState(context.Context, string, string, time.Time, time.Duration) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
func (r *memoryConnectionRepository) EvaluateOIDCSecurityEvent(context.Context, string, string, string, string, string, time.Time, time.Time, time.Duration) (store.SecurityEventEvaluation, error) {
|
func (r *memoryConnectionRepository) EvaluateOIDCSecurityEvent(context.Context, string, string, string, string, string, string, time.Time, time.Time, time.Duration) (store.SecurityEventEvaluation, error) {
|
||||||
if r.evaluation.Mode == "" {
|
if r.evaluation.Mode == "" {
|
||||||
return store.SecurityEventEvaluation{Mode: "bootstrap", RequireIntrospection: true}, nil
|
return store.SecurityEventEvaluation{Mode: "bootstrap", RequireIntrospection: true}, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -88,7 +88,8 @@ func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|||||||
var result store.ApplySecurityEventResult
|
var result store.ApplySecurityEventResult
|
||||||
result, err = h.repository.ApplySessionRevoked(r.Context(), store.ApplySessionRevokedInput{
|
result, err = h.repository.ApplySessionRevoked(r.Context(), store.ApplySessionRevokedInput{
|
||||||
Issuer: event.Issuer, Audience: event.Audience, JTI: event.JTI, TransactionID: event.TransactionID,
|
Issuer: event.Issuer, Audience: event.Audience, JTI: event.JTI, TransactionID: event.TransactionID,
|
||||||
SubjectIssuer: event.SubjectIssuer, TenantID: event.TenantID, Subject: event.Subject, EventTimestamp: event.EventTimestamp,
|
SubjectIssuer: event.SubjectIssuer, ApplicationID: event.ApplicationID, SubjectType: event.SubjectType,
|
||||||
|
TenantID: event.TenantID, Subject: event.Subject, EventTimestamp: event.EventTimestamp,
|
||||||
InitiatingEntity: event.InitiatingEntity,
|
InitiatingEntity: event.InitiatingEntity,
|
||||||
})
|
})
|
||||||
if err == nil {
|
if err == nil {
|
||||||
|
|||||||
@@ -109,6 +109,54 @@ func TestReceiverAcceptsValidSessionRevokedAndDuplicate(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestVerifierAcceptsApplicationScopedPrincipalAndTenantEvents(t *testing.T) {
|
||||||
|
privateKey, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||||
|
applicationID, tenantID, subjectID := uuid.NewString(), uuid.NewString(), uuid.NewString()
|
||||||
|
verifier := &Verifier{config: VerifierConfig{
|
||||||
|
TransmitterIssuer: "https://auth.example/ssf", Audience: "urn:easyai:ssf:receiver:" + applicationID,
|
||||||
|
SubjectIssuer: "https://auth.example/issuer/shared", ApplicationID: applicationID,
|
||||||
|
TenantMode: "multi_tenant", StreamID: uuid.NewString(), ClockSkew: time.Minute,
|
||||||
|
}, keys: map[string]*ecdsa.PublicKey{"current": &privateKey.PublicKey}, expiresAt: time.Now().Add(time.Hour), client: http.DefaultClient}
|
||||||
|
now := time.Now().UTC().Truncate(time.Second)
|
||||||
|
claims := func(subjectType, subject, eventApplicationID string) jwt.MapClaims {
|
||||||
|
return jwt.MapClaims{
|
||||||
|
"iss": verifier.config.TransmitterIssuer, "aud": verifier.config.Audience, "iat": now.Unix(),
|
||||||
|
"jti": uuid.NewString(), "txn": uuid.NewString(),
|
||||||
|
"sub_id": map[string]any{
|
||||||
|
"format": "complex",
|
||||||
|
"user": map[string]any{"format": "iss_sub", "iss": verifier.config.SubjectIssuer, "sub": subject},
|
||||||
|
"tenant": map[string]any{"format": "opaque", "id": tenantID},
|
||||||
|
},
|
||||||
|
"events": map[string]any{SessionRevokedEventType: map[string]any{
|
||||||
|
"event_timestamp": now.Unix(), "initiating_entity": "admin",
|
||||||
|
"application_id": eventApplicationID, "subject_type": subjectType,
|
||||||
|
}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, test := range []struct {
|
||||||
|
name, subjectType, subject string
|
||||||
|
}{
|
||||||
|
{name: "principal", subjectType: "principal", subject: subjectID},
|
||||||
|
{name: "tenant", subjectType: "tenant", subject: tenantID},
|
||||||
|
} {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
event, err := verifier.Verify(context.Background(), signSET(t, privateKey, claims(test.subjectType, test.subject, applicationID)))
|
||||||
|
if err != nil || event.ApplicationID != applicationID || event.SubjectType != test.subjectType ||
|
||||||
|
event.Subject != test.subject || event.TenantID != tenantID {
|
||||||
|
t.Fatalf("event=%#v error=%v", event, err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
for _, invalid := range []jwt.MapClaims{
|
||||||
|
claims("principal", subjectID, uuid.NewString()),
|
||||||
|
claims("tenant", subjectID, applicationID),
|
||||||
|
} {
|
||||||
|
if _, err := verifier.Verify(context.Background(), signSET(t, privateKey, invalid)); err == nil {
|
||||||
|
t.Fatal("invalid application-scoped event was accepted")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestReceiverRejectsAuthenticationMediaTypeAndForbiddenClaims(t *testing.T) {
|
func TestReceiverRejectsAuthenticationMediaTypeAndForbiddenClaims(t *testing.T) {
|
||||||
privateKey, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
privateKey, _ := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
||||||
verifier := &Verifier{config: VerifierConfig{
|
verifier := &Verifier{config: VerifierConfig{
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ type StateRepository interface {
|
|||||||
BeginSecurityEventVerification(context.Context, string, string, []byte, time.Time) error
|
BeginSecurityEventVerification(context.Context, string, string, []byte, time.Time) error
|
||||||
RecordSecurityEventHeartbeatFailure(context.Context, string, string, string) error
|
RecordSecurityEventHeartbeatFailure(context.Context, string, string, string) error
|
||||||
AdvanceSecurityEventStreamState(context.Context, string, string, time.Time, time.Duration) error
|
AdvanceSecurityEventStreamState(context.Context, string, string, time.Time, time.Duration) error
|
||||||
EvaluateOIDCSecurityEvent(context.Context, string, string, string, string, string, time.Time, time.Time, time.Duration) (store.SecurityEventEvaluation, error)
|
EvaluateOIDCSecurityEvent(context.Context, string, string, string, string, string, string, time.Time, time.Time, time.Duration) (store.SecurityEventEvaluation, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type ServiceConfig struct {
|
type ServiceConfig struct {
|
||||||
@@ -95,7 +95,7 @@ func (s *Service) RequestVerification(ctx context.Context) { s.sendVerification(
|
|||||||
|
|
||||||
func (s *Service) Evaluate(ctx context.Context, identity auth.OIDCSecurityEventIdentity) (auth.OIDCSecurityEventEvaluation, error) {
|
func (s *Service) Evaluate(ctx context.Context, identity auth.OIDCSecurityEventIdentity) (auth.OIDCSecurityEventEvaluation, error) {
|
||||||
result, err := s.repository.EvaluateOIDCSecurityEvent(
|
result, err := s.repository.EvaluateOIDCSecurityEvent(
|
||||||
ctx, s.config.Issuer, s.config.Audience, identity.Issuer, identity.TenantID, identity.Subject,
|
ctx, s.config.Issuer, s.config.Audience, identity.Issuer, identity.ApplicationID, identity.TenantID, identity.Subject,
|
||||||
identity.IssuedAt, s.clock().UTC(), s.config.StaleAfter,
|
identity.IssuedAt, s.clock().UTC(), s.config.StaleAfter,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ func (*fakeStateRepository) RecordSecurityEventHeartbeatFailure(context.Context,
|
|||||||
func (*fakeStateRepository) AdvanceSecurityEventStreamState(context.Context, string, string, time.Time, time.Duration) error {
|
func (*fakeStateRepository) AdvanceSecurityEventStreamState(context.Context, string, string, time.Time, time.Duration) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
func (f *fakeStateRepository) EvaluateOIDCSecurityEvent(context.Context, string, string, string, string, string, time.Time, time.Time, time.Duration) (store.SecurityEventEvaluation, error) {
|
func (f *fakeStateRepository) EvaluateOIDCSecurityEvent(context.Context, string, string, string, string, string, string, time.Time, time.Time, time.Duration) (store.SecurityEventEvaluation, error) {
|
||||||
return f.evaluation, nil
|
return f.evaluation, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -36,6 +36,8 @@ type Event struct {
|
|||||||
EventType string
|
EventType string
|
||||||
EventTimestamp time.Time
|
EventTimestamp time.Time
|
||||||
SubjectIssuer string
|
SubjectIssuer string
|
||||||
|
ApplicationID string
|
||||||
|
SubjectType string
|
||||||
Subject string
|
Subject string
|
||||||
TenantID string
|
TenantID string
|
||||||
InitiatingEntity string
|
InitiatingEntity string
|
||||||
@@ -50,6 +52,8 @@ type VerifierConfig struct {
|
|||||||
Audience string
|
Audience string
|
||||||
SubjectIssuer string
|
SubjectIssuer string
|
||||||
TenantID string
|
TenantID string
|
||||||
|
ApplicationID string
|
||||||
|
TenantMode string
|
||||||
StreamID string
|
StreamID string
|
||||||
ClockSkew time.Duration
|
ClockSkew time.Duration
|
||||||
HTTPClient *http.Client
|
HTTPClient *http.Client
|
||||||
@@ -93,8 +97,11 @@ func (e *protocolError) Error() string { return e.Code }
|
|||||||
func NewVerifier(config VerifierConfig) (*Verifier, error) {
|
func NewVerifier(config VerifierConfig) (*Verifier, error) {
|
||||||
config.TransmitterIssuer = strings.TrimRight(strings.TrimSpace(config.TransmitterIssuer), "/")
|
config.TransmitterIssuer = strings.TrimRight(strings.TrimSpace(config.TransmitterIssuer), "/")
|
||||||
config.SubjectIssuer = strings.TrimRight(strings.TrimSpace(config.SubjectIssuer), "/")
|
config.SubjectIssuer = strings.TrimRight(strings.TrimSpace(config.SubjectIssuer), "/")
|
||||||
if !absoluteHTTPURL(config.TransmitterIssuer) || !absoluteHTTPURL(config.SubjectIssuer) || config.Audience == "" || config.TenantID == "" {
|
multiTenant := config.TenantMode == "multi_tenant"
|
||||||
return nil, errors.New("SSF transmitter, audience, subject issuer, and tenant are required")
|
validBinding := multiTenant && uuid.Validate(config.ApplicationID) == nil ||
|
||||||
|
!multiTenant && uuid.Validate(config.TenantID) == nil
|
||||||
|
if !absoluteHTTPURL(config.TransmitterIssuer) || !absoluteHTTPURL(config.SubjectIssuer) || config.Audience == "" || !validBinding {
|
||||||
|
return nil, errors.New("SSF transmitter, audience, subject issuer, and identity binding are required")
|
||||||
}
|
}
|
||||||
if _, err := uuid.Parse(config.StreamID); err != nil {
|
if _, err := uuid.Parse(config.StreamID); err != nil {
|
||||||
return nil, errors.New("SSF stream id must be a UUID")
|
return nil, errors.New("SSF stream id must be a UUID")
|
||||||
@@ -179,6 +186,7 @@ func (v *Verifier) Verify(ctx context.Context, raw string) (Event, error) {
|
|||||||
return Event{}, &protocolError{Code: "invalid_request"}
|
return Event{}, &protocolError{Code: "invalid_request"}
|
||||||
}
|
}
|
||||||
event.EventType, event.EventTimestamp, event.InitiatingEntity = SessionRevokedEventType, unixClaim(value["event_timestamp"]), stringValue(value["initiating_entity"])
|
event.EventType, event.EventTimestamp, event.InitiatingEntity = SessionRevokedEventType, unixClaim(value["event_timestamp"]), stringValue(value["initiating_entity"])
|
||||||
|
event.ApplicationID, event.SubjectType = stringValue(value["application_id"]), stringValue(value["subject_type"])
|
||||||
if event.EventTimestamp.IsZero() || event.EventTimestamp.After(v.now().Add(v.config.ClockSkew)) ||
|
if event.EventTimestamp.IsZero() || event.EventTimestamp.After(v.now().Add(v.config.ClockSkew)) ||
|
||||||
!validInitiator(event.InitiatingEntity) || !v.readSessionSubject(claims["sub_id"], &event) {
|
!validInitiator(event.InitiatingEntity) || !v.readSessionSubject(claims["sub_id"], &event) {
|
||||||
return Event{}, &protocolError{Code: "invalid_request"}
|
return Event{}, &protocolError{Code: "invalid_request"}
|
||||||
@@ -259,11 +267,19 @@ func (v *Verifier) readSessionSubject(raw any, event *Event) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
event.SubjectIssuer, event.Subject, event.TenantID = strings.TrimRight(stringValue(user["iss"]), "/"), stringValue(user["sub"]), stringValue(tenant["id"])
|
event.SubjectIssuer, event.Subject, event.TenantID = strings.TrimRight(stringValue(user["iss"]), "/"), stringValue(user["sub"]), stringValue(tenant["id"])
|
||||||
if event.SubjectIssuer != v.config.SubjectIssuer || event.TenantID != v.config.TenantID {
|
if event.SubjectIssuer != v.config.SubjectIssuer || uuid.Validate(event.TenantID) != nil {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
_, err := uuid.Parse(event.Subject)
|
if v.config.TenantMode == "multi_tenant" {
|
||||||
return err == nil
|
if event.ApplicationID != v.config.ApplicationID ||
|
||||||
|
event.SubjectType != "principal" && event.SubjectType != "tenant" ||
|
||||||
|
event.SubjectType == "tenant" && event.Subject != event.TenantID {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
} else if event.ApplicationID != "" || event.SubjectType != "" || event.TenantID != v.config.TenantID {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return uuid.Validate(event.Subject) == nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (v *Verifier) readStreamSubject(raw any, event *Event) bool {
|
func (v *Verifier) readStreamSubject(raw any, event *Event) bool {
|
||||||
|
|||||||
@@ -21,10 +21,11 @@ type IdentityManagementRequest struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
const identityRevisionColumns = `
|
const identityRevisionColumns = `
|
||||||
id::text,state,schema_version,auth_center_url,COALESCE(issuer,''),COALESCE(tenant_id,''),COALESCE(application_id,''),
|
id::text,state,schema_version,tenant_mode,auth_center_url,COALESCE(issuer,''),COALESCE(tenant_id,''),COALESCE(application_id,''),
|
||||||
COALESCE(audience,''),COALESCE(browser_client_id,''),COALESCE(machine_client_id,''),scopes,capabilities,role_prefix,
|
COALESCE(audience,''),COALESCE(browser_client_id,''),COALESCE(machine_client_id,''),scopes,capabilities,role_prefix,
|
||||||
local_tenant_key,public_base_url,web_base_url,jit_enabled,legacy_jwt_enabled,token_introspection,session_revocation,
|
local_tenant_key,public_base_url,web_base_url,jit_enabled,legacy_jwt_enabled,token_introspection,session_revocation,
|
||||||
COALESCE(security_event_issuer,''),COALESCE(security_event_configuration_url,''),COALESCE(security_event_audience,''),
|
COALESCE(security_event_issuer,''),COALESCE(security_event_configuration_url,''),COALESCE(security_event_audience,''),
|
||||||
|
COALESCE(tenant_context_endpoint,''),COALESCE(tenant_context_audience,''),COALESCE(tenant_context_scope,''),
|
||||||
COALESCE(machine_credential_ref,''),COALESCE(session_encryption_key_ref,''),session_idle_seconds,session_absolute_seconds,
|
COALESCE(machine_credential_ref,''),COALESCE(session_encryption_key_ref,''),session_idle_seconds,session_absolute_seconds,
|
||||||
session_refresh_seconds,version,COALESCE(last_error_category,''),COALESCE(last_trace_id,''),COALESCE(last_audit_id,''),
|
session_refresh_seconds,version,COALESCE(last_error_category,''),COALESCE(last_trace_id,''),COALESCE(last_audit_id,''),
|
||||||
validated_at,activated_at,superseded_at,created_at,updated_at`
|
validated_at,activated_at,superseded_at,created_at,updated_at`
|
||||||
@@ -83,11 +84,11 @@ func (s *Store) CreateIdentityConfigurationRevision(ctx context.Context, revisio
|
|||||||
capabilities, _ := json.Marshal(revision.Capabilities)
|
capabilities, _ := json.Marshal(revision.Capabilities)
|
||||||
return scanIdentityRevision(s.pool.QueryRow(ctx, `
|
return scanIdentityRevision(s.pool.QueryRow(ctx, `
|
||||||
INSERT INTO gateway_identity_configuration_revisions (
|
INSERT INTO gateway_identity_configuration_revisions (
|
||||||
id,state,schema_version,auth_center_url,role_prefix,local_tenant_key,public_base_url,web_base_url,
|
id,state,schema_version,tenant_mode,auth_center_url,role_prefix,local_tenant_key,public_base_url,web_base_url,
|
||||||
jit_enabled,legacy_jwt_enabled,scopes,capabilities,session_idle_seconds,session_absolute_seconds,session_refresh_seconds
|
jit_enabled,legacy_jwt_enabled,scopes,capabilities,session_idle_seconds,session_absolute_seconds,session_refresh_seconds
|
||||||
) VALUES ($1::uuid,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11::jsonb,$12::jsonb,$13,$14,$15)
|
) VALUES ($1::uuid,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12::jsonb,$13::jsonb,$14,$15,$16)
|
||||||
RETURNING `+identityRevisionColumns,
|
RETURNING `+identityRevisionColumns,
|
||||||
revision.ID, revision.State, revision.SchemaVersion, revision.AuthCenterURL, revision.RolePrefix,
|
revision.ID, revision.State, revision.SchemaVersion, revision.TenantMode, revision.AuthCenterURL, revision.RolePrefix,
|
||||||
revision.LocalTenantKey, revision.PublicBaseURL, revision.WebBaseURL, revision.JITEnabled, revision.LegacyJWTEnabled,
|
revision.LocalTenantKey, revision.PublicBaseURL, revision.WebBaseURL, revision.JITEnabled, revision.LegacyJWTEnabled,
|
||||||
string(scopes), string(capabilities), revision.SessionIdleSeconds, revision.SessionAbsoluteSeconds, revision.SessionRefreshSeconds,
|
string(scopes), string(capabilities), revision.SessionIdleSeconds, revision.SessionAbsoluteSeconds, revision.SessionRefreshSeconds,
|
||||||
))
|
))
|
||||||
@@ -198,16 +199,19 @@ FROM gateway_identity_configuration_revisions WHERE id=$1::uuid FOR UPDATE`, id)
|
|||||||
capabilities, _ := json.Marshal(updated.Capabilities)
|
capabilities, _ := json.Marshal(updated.Capabilities)
|
||||||
revision, err := scanIdentityRevision(tx.QueryRow(ctx, `
|
revision, err := scanIdentityRevision(tx.QueryRow(ctx, `
|
||||||
UPDATE gateway_identity_configuration_revisions SET
|
UPDATE gateway_identity_configuration_revisions SET
|
||||||
issuer=$3,tenant_id=$4,application_id=$5,audience=NULLIF($6,''),browser_client_id=NULLIF($7,''),machine_client_id=NULLIF($8,''),
|
schema_version=$3,tenant_mode=$4,issuer=$5,tenant_id=NULLIF($6,''),application_id=$7,audience=NULLIF($8,''),
|
||||||
scopes=$9::jsonb,capabilities=$10::jsonb,token_introspection=$11,session_revocation=$12,
|
browser_client_id=NULLIF($9,''),machine_client_id=NULLIF($10,''),local_tenant_key=$11,
|
||||||
security_event_issuer=NULLIF($13,''),security_event_configuration_url=NULLIF($14,''),security_event_audience=NULLIF($15,''),
|
scopes=$12::jsonb,capabilities=$13::jsonb,token_introspection=$14,session_revocation=$15,
|
||||||
machine_credential_ref=NULLIF($16,''),session_encryption_key_ref=NULLIF($17,''),last_trace_id=NULLIF($18,''),
|
security_event_issuer=NULLIF($16,''),security_event_configuration_url=NULLIF($17,''),security_event_audience=NULLIF($18,''),
|
||||||
last_audit_id=NULLIF($19,''),last_error_category=NULL,version=version+1,updated_at=now()
|
tenant_context_endpoint=NULLIF($19,''),tenant_context_audience=NULLIF($20,''),tenant_context_scope=NULLIF($21,''),
|
||||||
|
machine_credential_ref=NULLIF($22,''),session_encryption_key_ref=NULLIF($23,''),last_trace_id=NULLIF($24,''),
|
||||||
|
last_audit_id=NULLIF($25,''),last_error_category=NULL,version=version+1,updated_at=now()
|
||||||
WHERE id=$1::uuid AND version=$2 AND state='draft'
|
WHERE id=$1::uuid AND version=$2 AND state='draft'
|
||||||
RETURNING `+identityRevisionColumns,
|
RETURNING `+identityRevisionColumns,
|
||||||
id, expectedVersion, updated.Issuer, updated.TenantID, updated.ApplicationID, updated.Audience,
|
id, expectedVersion, updated.SchemaVersion, updated.TenantMode, updated.Issuer, updated.TenantID, updated.ApplicationID, updated.Audience,
|
||||||
updated.BrowserClientID, updated.MachineClientID, string(scopes), string(capabilities), updated.TokenIntrospection,
|
updated.BrowserClientID, updated.MachineClientID, updated.LocalTenantKey, string(scopes), string(capabilities), updated.TokenIntrospection,
|
||||||
updated.SessionRevocation, updated.SecurityEventIssuer, updated.SecurityEventConfigURL, updated.SecurityEventAudience,
|
updated.SessionRevocation, updated.SecurityEventIssuer, updated.SecurityEventConfigURL, updated.SecurityEventAudience,
|
||||||
|
updated.TenantContextEndpoint, updated.TenantContextAudience, updated.TenantContextScope,
|
||||||
updated.MachineCredentialRef, updated.SessionEncryptionKeyRef, updated.LastTraceID, updated.LastAuditID,
|
updated.MachineCredentialRef, updated.SessionEncryptionKeyRef, updated.LastTraceID, updated.LastAuditID,
|
||||||
))
|
))
|
||||||
if errors.Is(err, pgx.ErrNoRows) {
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
@@ -296,8 +300,12 @@ func (s *Store) ActivateIdentityRevision(ctx context.Context, id string, expecte
|
|||||||
return identity.Revision{}, false, identity.ErrBreakGlassRequired
|
return identity.Revision{}, false, identity.ErrBreakGlassRequired
|
||||||
}
|
}
|
||||||
var tenantExists bool
|
var tenantExists bool
|
||||||
if err := tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM gateway_tenants WHERE tenant_key=(
|
if err := tx.QueryRow(ctx, `SELECT revision.tenant_mode='multi_tenant' OR EXISTS(
|
||||||
SELECT local_tenant_key FROM gateway_identity_configuration_revisions WHERE id=$1::uuid) AND status='active')`, id).Scan(&tenantExists); err != nil {
|
SELECT 1 FROM gateway_tenants tenant
|
||||||
|
WHERE tenant.tenant_key=revision.local_tenant_key AND tenant.status='active'
|
||||||
|
)
|
||||||
|
FROM gateway_identity_configuration_revisions revision
|
||||||
|
WHERE revision.id=$1::uuid`, id).Scan(&tenantExists); err != nil {
|
||||||
return identity.Revision{}, false, err
|
return identity.Revision{}, false, err
|
||||||
}
|
}
|
||||||
if !tenantExists {
|
if !tenantExists {
|
||||||
@@ -428,11 +436,12 @@ func scanIdentityRevision(row scanner) (identity.Revision, error) {
|
|||||||
var state string
|
var state string
|
||||||
var scopes, capabilities []byte
|
var scopes, capabilities []byte
|
||||||
if err := row.Scan(
|
if err := row.Scan(
|
||||||
&revision.ID, &state, &revision.SchemaVersion, &revision.AuthCenterURL, &revision.Issuer, &revision.TenantID,
|
&revision.ID, &state, &revision.SchemaVersion, &revision.TenantMode, &revision.AuthCenterURL, &revision.Issuer, &revision.TenantID,
|
||||||
&revision.ApplicationID, &revision.Audience, &revision.BrowserClientID, &revision.MachineClientID, &scopes,
|
&revision.ApplicationID, &revision.Audience, &revision.BrowserClientID, &revision.MachineClientID, &scopes,
|
||||||
&capabilities, &revision.RolePrefix, &revision.LocalTenantKey, &revision.PublicBaseURL, &revision.WebBaseURL,
|
&capabilities, &revision.RolePrefix, &revision.LocalTenantKey, &revision.PublicBaseURL, &revision.WebBaseURL,
|
||||||
&revision.JITEnabled, &revision.LegacyJWTEnabled, &revision.TokenIntrospection, &revision.SessionRevocation,
|
&revision.JITEnabled, &revision.LegacyJWTEnabled, &revision.TokenIntrospection, &revision.SessionRevocation,
|
||||||
&revision.SecurityEventIssuer, &revision.SecurityEventConfigURL, &revision.SecurityEventAudience,
|
&revision.SecurityEventIssuer, &revision.SecurityEventConfigURL, &revision.SecurityEventAudience,
|
||||||
|
&revision.TenantContextEndpoint, &revision.TenantContextAudience, &revision.TenantContextScope,
|
||||||
&revision.MachineCredentialRef, &revision.SessionEncryptionKeyRef, &revision.SessionIdleSeconds,
|
&revision.MachineCredentialRef, &revision.SessionEncryptionKeyRef, &revision.SessionIdleSeconds,
|
||||||
&revision.SessionAbsoluteSeconds, &revision.SessionRefreshSeconds, &revision.Version, &revision.LastErrorCategory,
|
&revision.SessionAbsoluteSeconds, &revision.SessionRefreshSeconds, &revision.Version, &revision.LastErrorCategory,
|
||||||
&revision.LastTraceID, &revision.LastAuditID, &revision.ValidatedAt, &revision.ActivatedAt, &revision.SupersededAt,
|
&revision.LastTraceID, &revision.LastAuditID, &revision.ValidatedAt, &revision.ActivatedAt, &revision.SupersededAt,
|
||||||
|
|||||||
@@ -0,0 +1,261 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
)
|
||||||
|
|
||||||
|
func (s *Store) resolveOrProvisionOIDCMultiTenantUser(ctx context.Context, input ResolveOrProvisionOIDCUserInput) (ResolveOrProvisionOIDCUserResult, error) {
|
||||||
|
if input.Issuer == "" || uuid.Validate(input.ApplicationID) != nil || uuid.Validate(input.TenantID) != nil ||
|
||||||
|
input.Subject == "" || input.TenantMetadataStatus != "synced" && input.TenantMetadataStatus != "metadata_pending" {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, errors.New("invalid multi-tenant OIDC user projection input")
|
||||||
|
}
|
||||||
|
tx, err := s.pool.Begin(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
defer func() { _ = tx.Rollback(ctx) }()
|
||||||
|
|
||||||
|
bindingID, gatewayTenantID, tenantKey, groupID, groupKey, accessStatus, err :=
|
||||||
|
loadOIDCTenantBinding(ctx, tx, input)
|
||||||
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
|
if !input.ProvisioningEnabled {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCUserNotProvisioned
|
||||||
|
}
|
||||||
|
bindingID, gatewayTenantID, tenantKey, groupID, groupKey, err =
|
||||||
|
createOIDCTenantBinding(ctx, tx, input)
|
||||||
|
accessStatus = "active"
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
if accessStatus != "active" {
|
||||||
|
if input.TenantMetadataStatus != "synced" {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
tag, reactivateErr := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
||||||
|
SET access_status='active',metadata_status='synced',last_error_category=NULL,updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND access_status='disabled'`, bindingID)
|
||||||
|
if reactivateErr != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, reactivateErr
|
||||||
|
}
|
||||||
|
if tag.RowsAffected() != 1 {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
accessStatus = "active"
|
||||||
|
}
|
||||||
|
if input.TenantMetadataStatus == "synced" {
|
||||||
|
if _, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings SET
|
||||||
|
metadata_status='synced',display_name=$2,slug=$3,metadata_version=NULLIF($4,''),
|
||||||
|
metadata_etag=NULLIF($5,''),metadata_updated_at=$6,
|
||||||
|
last_sync_at=now(),next_sync_at=now()+interval '15 minutes',sync_failure_count=0,last_error_category=NULL,updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND access_status='active'`,
|
||||||
|
bindingID, input.TenantName, input.TenantSlug, input.TenantMetadataVersion,
|
||||||
|
input.TenantMetadataETag, nullableOIDCTenantMetadataTime(input.TenantMetadataUpdatedAt),
|
||||||
|
); err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
if _, err := tx.Exec(ctx, `UPDATE gateway_tenants SET name=$2,synced_at=now(),source_updated_at=$3,updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND source='oidc_v2' AND status='active' AND deleted_at IS NULL`,
|
||||||
|
gatewayTenantID, input.TenantName, input.TenantMetadataUpdatedAt,
|
||||||
|
); err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
user, userBindingID, err := loadOIDCMultiTenantUser(ctx, tx, bindingID, input.Subject)
|
||||||
|
created := false
|
||||||
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
|
if !input.ProvisioningEnabled {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCUserNotProvisioned
|
||||||
|
}
|
||||||
|
user, userBindingID, created, err = s.createOIDCMultiTenantUser(
|
||||||
|
ctx, tx, bindingID, gatewayTenantID, tenantKey, groupID, input,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
if user.Status != "active" {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCUserDisabled
|
||||||
|
}
|
||||||
|
rolesJSON, _ := json.Marshal(input.Roles)
|
||||||
|
user, err = scanUser(tx.QueryRow(ctx, `UPDATE gateway_users SET
|
||||||
|
username=COALESCE(NULLIF($2,''),username),roles=$3::jsonb,last_login_at=now(),synced_at=now(),source_updated_at=now(),updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND source='oidc_v2' AND status='active' AND deleted_at IS NULL
|
||||||
|
RETURNING `+userColumns, user.ID, input.Username, string(rolesJSON)))
|
||||||
|
if err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
var auditID string
|
||||||
|
if created {
|
||||||
|
subjectHash := sha256.Sum256([]byte(input.Subject))
|
||||||
|
audit, err := s.RecordAuditLogTx(ctx, tx, AuditLogInput{
|
||||||
|
Category: "identity", Action: "identity.oidc_user.provisioned",
|
||||||
|
ActorGatewayUserID: user.ID, ActorUsername: user.Username, ActorSource: "oidc", ActorRoles: user.Roles,
|
||||||
|
TargetType: "gateway_user", TargetID: user.ID, TargetGatewayUserID: user.ID,
|
||||||
|
TargetGatewayTenantID: user.GatewayTenantID, RequestIP: input.RequestIP, UserAgent: input.UserAgent,
|
||||||
|
AfterState: map[string]any{"source": "oidc_v2", "tenantKey": user.TenantKey, "userGroupId": user.DefaultUserGroupID},
|
||||||
|
Metadata: map[string]any{
|
||||||
|
"provisioningMode": "oidc-multi-tenant-jit",
|
||||||
|
"externalSubjectHash": hex.EncodeToString(subjectHash[:])[:16],
|
||||||
|
},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
auditID = audit.ID
|
||||||
|
}
|
||||||
|
if err := tx.Commit(ctx); err != nil {
|
||||||
|
return ResolveOrProvisionOIDCUserResult{}, err
|
||||||
|
}
|
||||||
|
return ResolveOrProvisionOIDCUserResult{
|
||||||
|
User: multiTenantAuthUser(user, groupKey, input, userBindingID),
|
||||||
|
Created: created, AuditID: auditID,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadOIDCTenantBinding(ctx context.Context, tx pgx.Tx, input ResolveOrProvisionOIDCUserInput) (string, string, string, string, string, string, error) {
|
||||||
|
var bindingID, gatewayTenantID, tenantKey, groupID, groupKey, accessStatus string
|
||||||
|
err := tx.QueryRow(ctx, `SELECT binding.id::text,tenant.id::text,tenant.tenant_key,
|
||||||
|
tenant.default_user_group_id::text,group_record.group_key,binding.access_status
|
||||||
|
FROM gateway_oidc_tenant_bindings binding
|
||||||
|
JOIN gateway_tenants tenant ON tenant.id=binding.gateway_tenant_id
|
||||||
|
JOIN gateway_user_groups group_record ON group_record.id=tenant.default_user_group_id
|
||||||
|
WHERE binding.issuer=$1 AND binding.application_id=$2 AND binding.external_tenant_id=$3
|
||||||
|
AND tenant.status='active' AND tenant.deleted_at IS NULL AND group_record.status='active'
|
||||||
|
FOR UPDATE OF binding,tenant`,
|
||||||
|
input.Issuer, input.ApplicationID, input.TenantID,
|
||||||
|
).Scan(&bindingID, &gatewayTenantID, &tenantKey, &groupID, &groupKey, &accessStatus)
|
||||||
|
return bindingID, gatewayTenantID, tenantKey, groupID, groupKey, accessStatus, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func createOIDCTenantBinding(ctx context.Context, tx pgx.Tx, input ResolveOrProvisionOIDCUserInput) (string, string, string, string, string, error) {
|
||||||
|
var groupID, groupKey string
|
||||||
|
if err := tx.QueryRow(ctx, `SELECT id::text,group_key FROM gateway_user_groups
|
||||||
|
WHERE group_key='default' AND status='active'`).Scan(&groupID, &groupKey); err != nil {
|
||||||
|
return "", "", "", "", "", ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
tenantKey := deriveOIDCMultiTenantKey(input.Issuer, input.ApplicationID, input.TenantID)
|
||||||
|
name := input.TenantName
|
||||||
|
if name == "" {
|
||||||
|
name = "认证中心租户 " + strings.ReplaceAll(input.TenantID, "-", "")[:8]
|
||||||
|
}
|
||||||
|
metadata, _ := json.Marshal(map[string]any{"provisioningMode": "oidc-multi-tenant-jit", "metadataStatus": input.TenantMetadataStatus})
|
||||||
|
var gatewayTenantID string
|
||||||
|
err := tx.QueryRow(ctx, `INSERT INTO gateway_tenants(
|
||||||
|
tenant_key,source,external_tenant_id,name,default_user_group_id,metadata,status,synced_at,source_updated_at
|
||||||
|
) VALUES($1,'oidc_v2',NULL,$2,$3::uuid,$4::jsonb,'active',
|
||||||
|
CASE WHEN $5='synced' THEN now() ELSE NULL END,$6)
|
||||||
|
ON CONFLICT(tenant_key) DO UPDATE SET updated_at=gateway_tenants.updated_at
|
||||||
|
RETURNING id::text`, tenantKey, name, groupID, string(metadata), input.TenantMetadataStatus, nullableOIDCTenantMetadataTime(input.TenantMetadataUpdatedAt)).Scan(&gatewayTenantID)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", "", "", "", err
|
||||||
|
}
|
||||||
|
var bindingID string
|
||||||
|
err = tx.QueryRow(ctx, `INSERT INTO gateway_oidc_tenant_bindings(
|
||||||
|
issuer,application_id,external_tenant_id,gateway_tenant_id,access_status,metadata_status,
|
||||||
|
display_name,slug,metadata_version,metadata_etag,metadata_updated_at,last_sync_at,next_sync_at
|
||||||
|
) VALUES($1,$2,$3,$4::uuid,'active',$5,NULLIF($6,''),NULLIF($7,''),NULLIF($8,''),NULLIF($9,''),
|
||||||
|
$10,
|
||||||
|
CASE WHEN $5='synced' THEN now() ELSE NULL END,
|
||||||
|
CASE WHEN $5='synced' THEN now()+interval '15 minutes' ELSE now()+interval '30 seconds' END)
|
||||||
|
ON CONFLICT(issuer,application_id,external_tenant_id) DO UPDATE SET updated_at=gateway_oidc_tenant_bindings.updated_at
|
||||||
|
RETURNING id::text`,
|
||||||
|
input.Issuer, input.ApplicationID, input.TenantID, gatewayTenantID, input.TenantMetadataStatus,
|
||||||
|
input.TenantName, input.TenantSlug, input.TenantMetadataVersion, input.TenantMetadataETag,
|
||||||
|
nullableOIDCTenantMetadataTime(input.TenantMetadataUpdatedAt),
|
||||||
|
).Scan(&bindingID)
|
||||||
|
return bindingID, gatewayTenantID, tenantKey, groupID, groupKey, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func loadOIDCMultiTenantUser(ctx context.Context, tx pgx.Tx, bindingID, subject string) (GatewayUser, string, error) {
|
||||||
|
var userBindingID, gatewayUserID string
|
||||||
|
if err := tx.QueryRow(ctx, `SELECT binding.id::text,binding.gateway_user_id::text
|
||||||
|
FROM gateway_oidc_user_bindings binding
|
||||||
|
JOIN gateway_users ON gateway_users.id=binding.gateway_user_id
|
||||||
|
WHERE binding.tenant_binding_id=$1::uuid AND binding.subject=$2
|
||||||
|
AND gateway_users.deleted_at IS NULL
|
||||||
|
FOR UPDATE OF binding,gateway_users`, bindingID, subject).Scan(&userBindingID, &gatewayUserID); err != nil {
|
||||||
|
return GatewayUser{}, "", err
|
||||||
|
}
|
||||||
|
user, err := scanUser(tx.QueryRow(ctx, `SELECT `+userColumns+`
|
||||||
|
FROM gateway_users WHERE id=$1::uuid`, gatewayUserID))
|
||||||
|
return user, userBindingID, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) createOIDCMultiTenantUser(ctx context.Context, tx pgx.Tx, bindingID, gatewayTenantID, tenantKey, groupID string, input ResolveOrProvisionOIDCUserInput) (GatewayUser, string, bool, error) {
|
||||||
|
userKey := deriveOIDCMultiTenantUserKey(input.Issuer, input.ApplicationID, input.TenantID, input.Subject)
|
||||||
|
username := input.Username
|
||||||
|
if username == "" {
|
||||||
|
username = "oidc-" + strings.TrimPrefix(userKey, "oidc2:")[:12]
|
||||||
|
}
|
||||||
|
rolesJSON, _ := json.Marshal(input.Roles)
|
||||||
|
metadata, _ := json.Marshal(map[string]any{"provisioningMode": "oidc-multi-tenant-jit"})
|
||||||
|
user, err := scanUser(tx.QueryRow(ctx, `INSERT INTO gateway_users(
|
||||||
|
user_key,source,external_user_id,username,gateway_tenant_id,tenant_id,tenant_key,
|
||||||
|
default_user_group_id,roles,auth_profile,metadata,status,last_login_at,synced_at,source_updated_at
|
||||||
|
) VALUES($1,'oidc_v2',NULL,$2,$3::uuid,$4,$5,$6::uuid,$7::jsonb,'{}'::jsonb,$8::jsonb,
|
||||||
|
'active',now(),now(),now())
|
||||||
|
ON CONFLICT(user_key) DO UPDATE SET updated_at=gateway_users.updated_at
|
||||||
|
RETURNING `+userColumns,
|
||||||
|
userKey, username, gatewayTenantID, input.TenantID, tenantKey, groupID, string(rolesJSON), string(metadata),
|
||||||
|
))
|
||||||
|
if err != nil {
|
||||||
|
return GatewayUser{}, "", false, err
|
||||||
|
}
|
||||||
|
var userBindingID string
|
||||||
|
tag, err := tx.Exec(ctx, `INSERT INTO gateway_oidc_user_bindings(tenant_binding_id,subject,gateway_user_id)
|
||||||
|
VALUES($1::uuid,$2,$3::uuid) ON CONFLICT(tenant_binding_id,subject) DO NOTHING`,
|
||||||
|
bindingID, input.Subject, user.ID)
|
||||||
|
if err != nil {
|
||||||
|
return GatewayUser{}, "", false, err
|
||||||
|
}
|
||||||
|
created := tag.RowsAffected() == 1
|
||||||
|
if err := tx.QueryRow(ctx, `SELECT id::text FROM gateway_oidc_user_bindings
|
||||||
|
WHERE tenant_binding_id=$1::uuid AND subject=$2`, bindingID, input.Subject).Scan(&userBindingID); err != nil {
|
||||||
|
return GatewayUser{}, "", false, err
|
||||||
|
}
|
||||||
|
if _, err := s.ensureWalletAccount(ctx, tx, user.ID, "resource"); err != nil {
|
||||||
|
return GatewayUser{}, "", false, err
|
||||||
|
}
|
||||||
|
return user, userBindingID, created, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func deriveOIDCMultiTenantKey(issuer, applicationID, tenantID string) string {
|
||||||
|
sum := sha256.Sum256([]byte(strings.TrimRight(issuer, "/") + "\x00" + applicationID + "\x00" + tenantID))
|
||||||
|
return fmt.Sprintf("oidc2-tenant-%x", sum[:16])
|
||||||
|
}
|
||||||
|
|
||||||
|
func nullableOIDCTenantMetadataTime(value time.Time) any {
|
||||||
|
if value.IsZero() {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
func deriveOIDCMultiTenantUserKey(issuer, applicationID, tenantID, subject string) string {
|
||||||
|
sum := sha256.Sum256([]byte(strings.TrimRight(issuer, "/") + "\x00" + applicationID + "\x00" + tenantID + "\x00" + subject))
|
||||||
|
return fmt.Sprintf("oidc2:%x", sum)
|
||||||
|
}
|
||||||
|
|
||||||
|
func multiTenantAuthUser(user GatewayUser, groupKey string, input ResolveOrProvisionOIDCUserInput, userBindingID string) *auth.User {
|
||||||
|
result := authUserFromOIDCProjection(user, groupKey)
|
||||||
|
result.ID = input.Subject
|
||||||
|
result.TenantName = input.TenantName
|
||||||
|
result.Issuer = input.Issuer
|
||||||
|
result.ApplicationID = input.ApplicationID
|
||||||
|
result.OIDCClientID = input.OIDCClientID
|
||||||
|
result.OIDCUserBindingID = userBindingID
|
||||||
|
return result
|
||||||
|
}
|
||||||
@@ -15,7 +15,12 @@ type OIDCSession struct {
|
|||||||
SessionTokenHash []byte
|
SessionTokenHash []byte
|
||||||
GatewayUserID string
|
GatewayUserID string
|
||||||
GatewayTenantID string
|
GatewayTenantID string
|
||||||
|
OIDCUserBindingID string
|
||||||
ExternalUserID string
|
ExternalUserID string
|
||||||
|
Issuer string
|
||||||
|
ApplicationID string
|
||||||
|
TenantID string
|
||||||
|
OIDCClientID string
|
||||||
UserStatus string
|
UserStatus string
|
||||||
UserDeleted bool
|
UserDeleted bool
|
||||||
TokenCiphertext []byte
|
TokenCiphertext []byte
|
||||||
@@ -34,6 +39,11 @@ type CreateOIDCSessionInput struct {
|
|||||||
SessionTokenHash []byte
|
SessionTokenHash []byte
|
||||||
GatewayUserID string
|
GatewayUserID string
|
||||||
GatewayTenantID string
|
GatewayTenantID string
|
||||||
|
OIDCUserBindingID string
|
||||||
|
OIDCClientID string
|
||||||
|
Issuer string
|
||||||
|
ApplicationID string
|
||||||
|
TenantID string
|
||||||
TokenCiphertext []byte
|
TokenCiphertext []byte
|
||||||
AccessTokenExpiresAt time.Time
|
AccessTokenExpiresAt time.Time
|
||||||
LastSeenAt time.Time
|
LastSeenAt time.Time
|
||||||
@@ -46,12 +56,27 @@ func (s *Store) CreateOIDCSession(ctx context.Context, input CreateOIDCSessionIn
|
|||||||
err := s.pool.QueryRow(ctx, `
|
err := s.pool.QueryRow(ctx, `
|
||||||
INSERT INTO gateway_oidc_sessions (
|
INSERT INTO gateway_oidc_sessions (
|
||||||
session_token_hash, gateway_user_id, gateway_tenant_id, token_ciphertext,
|
session_token_hash, gateway_user_id, gateway_tenant_id, token_ciphertext,
|
||||||
access_token_expires_at, last_seen_at, idle_expires_at, absolute_expires_at
|
access_token_expires_at, last_seen_at, idle_expires_at, absolute_expires_at,
|
||||||
|
oidc_user_binding_id, oidc_client_id
|
||||||
)
|
)
|
||||||
VALUES ($1, $2::uuid, $3::uuid, $4, $5, $6, $7, $8)
|
SELECT $1, u.id, $3::uuid, $4, $5, $6, $7, $8, NULLIF($9, '')::uuid, NULLIF($10, '')
|
||||||
|
FROM gateway_users u
|
||||||
|
LEFT JOIN gateway_oidc_user_bindings ub ON ub.id = NULLIF($9, '')::uuid
|
||||||
|
LEFT JOIN gateway_oidc_tenant_bindings tb ON tb.id = ub.tenant_binding_id
|
||||||
|
WHERE u.id = $2::uuid
|
||||||
|
AND u.gateway_tenant_id = $3::uuid
|
||||||
|
AND ($9 = '' OR (
|
||||||
|
ub.gateway_user_id = u.id
|
||||||
|
AND tb.gateway_tenant_id = u.gateway_tenant_id
|
||||||
|
AND tb.issuer = $11
|
||||||
|
AND tb.application_id = $12
|
||||||
|
AND tb.external_tenant_id = $13
|
||||||
|
))
|
||||||
RETURNING id::text`,
|
RETURNING id::text`,
|
||||||
input.SessionTokenHash, input.GatewayUserID, input.GatewayTenantID, input.TokenCiphertext,
|
input.SessionTokenHash, input.GatewayUserID, input.GatewayTenantID, input.TokenCiphertext,
|
||||||
input.AccessTokenExpiresAt, input.LastSeenAt, input.IdleExpiresAt, input.AbsoluteExpiresAt,
|
input.AccessTokenExpiresAt, input.LastSeenAt, input.IdleExpiresAt, input.AbsoluteExpiresAt,
|
||||||
|
input.OIDCUserBindingID, input.OIDCClientID,
|
||||||
|
input.Issuer, input.ApplicationID, input.TenantID,
|
||||||
).Scan(&id)
|
).Scan(&id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return OIDCSession{}, err
|
return OIDCSession{}, err
|
||||||
@@ -64,6 +89,8 @@ func (s *Store) FindOIDCSessionByHash(ctx context.Context, hash []byte) (OIDCSes
|
|||||||
SELECT `+oidcSessionColumns+`
|
SELECT `+oidcSessionColumns+`
|
||||||
FROM gateway_oidc_sessions s
|
FROM gateway_oidc_sessions s
|
||||||
JOIN gateway_users u ON u.id = s.gateway_user_id
|
JOIN gateway_users u ON u.id = s.gateway_user_id
|
||||||
|
LEFT JOIN gateway_oidc_user_bindings ub ON ub.id = s.oidc_user_binding_id
|
||||||
|
LEFT JOIN gateway_oidc_tenant_bindings tb ON tb.id = ub.tenant_binding_id
|
||||||
WHERE s.session_token_hash = $1`, hash))
|
WHERE s.session_token_hash = $1`, hash))
|
||||||
if errors.Is(err, pgx.ErrNoRows) {
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
return OIDCSession{}, ErrOIDCSessionNotFound
|
return OIDCSession{}, ErrOIDCSessionNotFound
|
||||||
@@ -138,7 +165,9 @@ WHERE idle_expires_at <= $1 OR absolute_expires_at <= $1`, now)
|
|||||||
|
|
||||||
const oidcSessionColumns = `
|
const oidcSessionColumns = `
|
||||||
s.id::text, s.session_token_hash, s.gateway_user_id::text, s.gateway_tenant_id::text,
|
s.id::text, s.session_token_hash, s.gateway_user_id::text, s.gateway_tenant_id::text,
|
||||||
COALESCE(u.external_user_id, ''), u.status, u.deleted_at IS NOT NULL,
|
COALESCE(s.oidc_user_binding_id::text, ''), COALESCE(ub.subject, u.external_user_id, ''),
|
||||||
|
COALESCE(tb.issuer, ''), COALESCE(tb.application_id, ''), COALESCE(tb.external_tenant_id, ''),
|
||||||
|
COALESCE(s.oidc_client_id, ''), u.status, u.deleted_at IS NOT NULL,
|
||||||
s.token_ciphertext, s.access_token_expires_at, s.last_seen_at, s.idle_expires_at,
|
s.token_ciphertext, s.access_token_expires_at, s.last_seen_at, s.idle_expires_at,
|
||||||
s.absolute_expires_at, s.refresh_version, COALESCE(s.refresh_lock_id::text, ''),
|
s.absolute_expires_at, s.refresh_version, COALESCE(s.refresh_lock_id::text, ''),
|
||||||
COALESCE(s.refresh_lock_until, 'epoch'::timestamptz), s.created_at, s.updated_at`
|
COALESCE(s.refresh_lock_until, 'epoch'::timestamptz), s.created_at, s.updated_at`
|
||||||
@@ -147,7 +176,8 @@ func scanOIDCSession(row pgx.Row) (OIDCSession, error) {
|
|||||||
var item OIDCSession
|
var item OIDCSession
|
||||||
err := row.Scan(
|
err := row.Scan(
|
||||||
&item.ID, &item.SessionTokenHash, &item.GatewayUserID, &item.GatewayTenantID,
|
&item.ID, &item.SessionTokenHash, &item.GatewayUserID, &item.GatewayTenantID,
|
||||||
&item.ExternalUserID, &item.UserStatus, &item.UserDeleted, &item.TokenCiphertext,
|
&item.OIDCUserBindingID, &item.ExternalUserID, &item.Issuer, &item.ApplicationID,
|
||||||
|
&item.TenantID, &item.OIDCClientID, &item.UserStatus, &item.UserDeleted, &item.TokenCiphertext,
|
||||||
&item.AccessTokenExpiresAt, &item.LastSeenAt, &item.IdleExpiresAt, &item.AbsoluteExpiresAt,
|
&item.AccessTokenExpiresAt, &item.LastSeenAt, &item.IdleExpiresAt, &item.AbsoluteExpiresAt,
|
||||||
&item.RefreshVersion, &item.RefreshLockID, &item.RefreshLockUntil, &item.CreatedAt, &item.UpdatedAt,
|
&item.RefreshVersion, &item.RefreshLockID, &item.RefreshLockUntil, &item.CreatedAt, &item.UpdatedAt,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,176 @@
|
|||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/identity"
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"github.com/jackc/pgx/v5"
|
||||||
|
)
|
||||||
|
|
||||||
|
type OIDCTenantBindingSyncTarget struct {
|
||||||
|
ID string
|
||||||
|
ExternalTenantID string
|
||||||
|
ETag string
|
||||||
|
FailureCount int
|
||||||
|
}
|
||||||
|
|
||||||
|
type OIDCTenantBindingContext struct {
|
||||||
|
ID string
|
||||||
|
AccessStatus string
|
||||||
|
MetadataStatus string
|
||||||
|
DisplayName string
|
||||||
|
Slug string
|
||||||
|
Version string
|
||||||
|
ETag string
|
||||||
|
MetadataUpdatedAt time.Time
|
||||||
|
NextSyncAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) DueOIDCTenantBindingSyncs(ctx context.Context, issuer, applicationID string, limit int) ([]OIDCTenantBindingSyncTarget, error) {
|
||||||
|
if limit <= 0 || limit > 100 {
|
||||||
|
limit = 50
|
||||||
|
}
|
||||||
|
rows, err := s.pool.Query(ctx, `SELECT id::text,external_tenant_id,COALESCE(metadata_etag,''),sync_failure_count
|
||||||
|
FROM gateway_oidc_tenant_bindings
|
||||||
|
WHERE issuer=$1 AND application_id=$2 AND access_status='active' AND next_sync_at <= now()
|
||||||
|
ORDER BY next_sync_at,id
|
||||||
|
LIMIT $3`, strings.TrimRight(strings.TrimSpace(issuer), "/"), strings.TrimSpace(applicationID), limit)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
result := make([]OIDCTenantBindingSyncTarget, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var item OIDCTenantBindingSyncTarget
|
||||||
|
if err := rows.Scan(&item.ID, &item.ExternalTenantID, &item.ETag, &item.FailureCount); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
result = append(result, item)
|
||||||
|
}
|
||||||
|
return result, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) OIDCTenantBindingContext(ctx context.Context, issuer, applicationID, tenantID string) (OIDCTenantBindingContext, error) {
|
||||||
|
var item OIDCTenantBindingContext
|
||||||
|
err := s.pool.QueryRow(ctx, `SELECT id::text,access_status,metadata_status,COALESCE(display_name,''),
|
||||||
|
COALESCE(slug,''),COALESCE(metadata_version,''),COALESCE(metadata_etag,''),
|
||||||
|
COALESCE(metadata_updated_at,'epoch'::timestamptz),next_sync_at
|
||||||
|
FROM gateway_oidc_tenant_bindings
|
||||||
|
WHERE issuer=$1 AND application_id=$2 AND external_tenant_id=$3`,
|
||||||
|
strings.TrimRight(strings.TrimSpace(issuer), "/"), strings.TrimSpace(applicationID), strings.TrimSpace(tenantID),
|
||||||
|
).Scan(&item.ID, &item.AccessStatus, &item.MetadataStatus, &item.DisplayName, &item.Slug,
|
||||||
|
&item.Version, &item.ETag, &item.MetadataUpdatedAt, &item.NextSyncAt)
|
||||||
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
|
return OIDCTenantBindingContext{}, ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
return item, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) ApplyOIDCTenantBindingSync(ctx context.Context, bindingID string, tenant identity.TenantContext, unchanged bool, now time.Time) error {
|
||||||
|
if uuid.Validate(bindingID) != nil {
|
||||||
|
return errors.New("OIDC tenant binding id is invalid")
|
||||||
|
}
|
||||||
|
tx, err := s.pool.Begin(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer func() { _ = tx.Rollback(ctx) }()
|
||||||
|
if unchanged {
|
||||||
|
tag, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
||||||
|
SET last_sync_at=$2::timestamptz,next_sync_at=$2::timestamptz+interval '15 minutes',sync_failure_count=0,
|
||||||
|
last_error_category=NULL,updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND access_status='active'`, bindingID, now)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if tag.RowsAffected() == 0 {
|
||||||
|
return ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
return tx.Commit(ctx)
|
||||||
|
}
|
||||||
|
if !tenant.Active() || strings.TrimSpace(tenant.DisplayName) == "" || strings.TrimSpace(tenant.Slug) == "" {
|
||||||
|
return errors.New("OIDC tenant context is inactive or incomplete")
|
||||||
|
}
|
||||||
|
var gatewayTenantID string
|
||||||
|
err = tx.QueryRow(ctx, `UPDATE gateway_oidc_tenant_bindings
|
||||||
|
SET metadata_status='synced',display_name=$2,slug=$3,metadata_version=$4,metadata_etag=NULLIF($5,''),
|
||||||
|
metadata_updated_at=$6,last_sync_at=$7::timestamptz,
|
||||||
|
next_sync_at=$7::timestamptz+interval '15 minutes',
|
||||||
|
sync_failure_count=0,last_error_category=NULL,updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND access_status='active'
|
||||||
|
RETURNING gateway_tenant_id::text`,
|
||||||
|
bindingID, tenant.DisplayName, tenant.Slug, tenant.Version, tenant.ETag, tenant.UpdatedAt, now,
|
||||||
|
).Scan(&gatewayTenantID)
|
||||||
|
if errors.Is(err, pgx.ErrNoRows) {
|
||||||
|
return ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if _, err := tx.Exec(ctx, `UPDATE gateway_tenants
|
||||||
|
SET name=$2,synced_at=$3,source_updated_at=$4,updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND source='oidc_v2' AND status='active' AND deleted_at IS NULL`,
|
||||||
|
gatewayTenantID, tenant.DisplayName, now, tenant.UpdatedAt,
|
||||||
|
); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return tx.Commit(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) RejectOIDCTenantBinding(ctx context.Context, bindingID, category string, now time.Time) error {
|
||||||
|
if uuid.Validate(bindingID) != nil {
|
||||||
|
return errors.New("OIDC tenant binding id is invalid")
|
||||||
|
}
|
||||||
|
tx, err := s.pool.Begin(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer func() { _ = tx.Rollback(ctx) }()
|
||||||
|
tag, err := tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
||||||
|
SET access_status='disabled',metadata_status='rejected',last_sync_at=$2::timestamptz,
|
||||||
|
next_sync_at=$2::timestamptz+interval '15 minutes',
|
||||||
|
last_error_category=$3,updated_at=now()
|
||||||
|
WHERE id=$1::uuid`, bindingID, now, limitOIDCTenantContextCategory(category))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if tag.RowsAffected() == 0 {
|
||||||
|
return ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
if _, err := tx.Exec(ctx, `DELETE FROM gateway_oidc_sessions session
|
||||||
|
USING gateway_oidc_user_bindings user_binding
|
||||||
|
WHERE session.oidc_user_binding_id=user_binding.id AND user_binding.tenant_binding_id=$1::uuid`,
|
||||||
|
bindingID,
|
||||||
|
); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return tx.Commit(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Store) FailOIDCTenantBindingSync(ctx context.Context, bindingID, category string, next time.Time) error {
|
||||||
|
if uuid.Validate(bindingID) != nil {
|
||||||
|
return errors.New("OIDC tenant binding id is invalid")
|
||||||
|
}
|
||||||
|
tag, err := s.pool.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
||||||
|
SET sync_failure_count=sync_failure_count+1,next_sync_at=$2,last_error_category=$3,updated_at=now()
|
||||||
|
WHERE id=$1::uuid AND access_status='active'`, bindingID, next, limitOIDCTenantContextCategory(category))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if tag.RowsAffected() == 0 {
|
||||||
|
return ErrOIDCTenantUnavailable
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func limitOIDCTenantContextCategory(value string) string {
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
if len(value) > 64 {
|
||||||
|
return value[:64]
|
||||||
|
}
|
||||||
|
return value
|
||||||
|
}
|
||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
"github.com/jackc/pgx/v5"
|
"github.com/jackc/pgx/v5"
|
||||||
@@ -20,15 +21,24 @@ var (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type ResolveOrProvisionOIDCUserInput struct {
|
type ResolveOrProvisionOIDCUserInput struct {
|
||||||
Issuer string
|
Issuer string
|
||||||
Subject string
|
ApplicationID string
|
||||||
Username string
|
Subject string
|
||||||
Roles []string
|
Username string
|
||||||
TenantID string
|
Roles []string
|
||||||
GatewayTenantKey string
|
TenantID string
|
||||||
ProvisioningEnabled bool
|
TenantMode string
|
||||||
RequestIP string
|
TenantName string
|
||||||
UserAgent string
|
TenantSlug string
|
||||||
|
TenantMetadataStatus string
|
||||||
|
TenantMetadataVersion string
|
||||||
|
TenantMetadataETag string
|
||||||
|
TenantMetadataUpdatedAt time.Time
|
||||||
|
OIDCClientID string
|
||||||
|
GatewayTenantKey string
|
||||||
|
ProvisioningEnabled bool
|
||||||
|
RequestIP string
|
||||||
|
UserAgent string
|
||||||
}
|
}
|
||||||
|
|
||||||
type ResolveOrProvisionOIDCUserResult struct {
|
type ResolveOrProvisionOIDCUserResult struct {
|
||||||
@@ -48,6 +58,9 @@ type oidcUserProjection struct {
|
|||||||
|
|
||||||
func (s *Store) ResolveOrProvisionOIDCUser(ctx context.Context, input ResolveOrProvisionOIDCUserInput) (ResolveOrProvisionOIDCUserResult, error) {
|
func (s *Store) ResolveOrProvisionOIDCUser(ctx context.Context, input ResolveOrProvisionOIDCUserInput) (ResolveOrProvisionOIDCUserResult, error) {
|
||||||
input = normalizeOIDCUserInput(input)
|
input = normalizeOIDCUserInput(input)
|
||||||
|
if input.TenantMode == "multi_tenant" {
|
||||||
|
return s.resolveOrProvisionOIDCMultiTenantUser(ctx, input)
|
||||||
|
}
|
||||||
if input.Issuer == "" || input.Subject == "" || input.TenantID == "" {
|
if input.Issuer == "" || input.Subject == "" || input.TenantID == "" {
|
||||||
return ResolveOrProvisionOIDCUserResult{}, errors.New("invalid OIDC user projection input")
|
return ResolveOrProvisionOIDCUserResult{}, errors.New("invalid OIDC user projection input")
|
||||||
}
|
}
|
||||||
@@ -296,6 +309,13 @@ func normalizeOIDCUserInput(input ResolveOrProvisionOIDCUserInput) ResolveOrProv
|
|||||||
input.Subject = strings.TrimSpace(input.Subject)
|
input.Subject = strings.TrimSpace(input.Subject)
|
||||||
input.Username = strings.TrimSpace(input.Username)
|
input.Username = strings.TrimSpace(input.Username)
|
||||||
input.TenantID = strings.TrimSpace(input.TenantID)
|
input.TenantID = strings.TrimSpace(input.TenantID)
|
||||||
|
input.ApplicationID = strings.TrimSpace(input.ApplicationID)
|
||||||
|
input.TenantMode = strings.TrimSpace(input.TenantMode)
|
||||||
|
input.TenantName = strings.TrimSpace(input.TenantName)
|
||||||
|
input.TenantSlug = strings.TrimSpace(input.TenantSlug)
|
||||||
|
input.TenantMetadataStatus = strings.TrimSpace(input.TenantMetadataStatus)
|
||||||
|
input.TenantMetadataVersion = strings.TrimSpace(input.TenantMetadataVersion)
|
||||||
|
input.TenantMetadataETag = strings.TrimSpace(input.TenantMetadataETag)
|
||||||
input.GatewayTenantKey = strings.TrimSpace(input.GatewayTenantKey)
|
input.GatewayTenantKey = strings.TrimSpace(input.GatewayTenantKey)
|
||||||
input.RequestIP = strings.TrimSpace(input.RequestIP)
|
input.RequestIP = strings.TrimSpace(input.RequestIP)
|
||||||
input.UserAgent = strings.TrimSpace(input.UserAgent)
|
input.UserAgent = strings.TrimSpace(input.UserAgent)
|
||||||
|
|||||||
@@ -11,9 +11,147 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"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"
|
"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) {
|
func TestResolveOrProvisionOIDCUserLifecycleAndConcurrency(t *testing.T) {
|
||||||
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
|
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
|
||||||
if databaseURL == "" {
|
if databaseURL == "" {
|
||||||
|
|||||||
@@ -16,6 +16,8 @@ type ApplySessionRevokedInput struct {
|
|||||||
JTI string
|
JTI string
|
||||||
TransactionID string
|
TransactionID string
|
||||||
SubjectIssuer string
|
SubjectIssuer string
|
||||||
|
ApplicationID string
|
||||||
|
SubjectType string
|
||||||
TenantID string
|
TenantID string
|
||||||
Subject string
|
Subject string
|
||||||
EventTimestamp time.Time
|
EventTimestamp time.Time
|
||||||
@@ -36,6 +38,9 @@ type SecurityEventEvaluation struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (s *Store) ApplySessionRevoked(ctx context.Context, input ApplySessionRevokedInput) (ApplySecurityEventResult, error) {
|
func (s *Store) ApplySessionRevoked(ctx context.Context, input ApplySessionRevokedInput) (ApplySecurityEventResult, error) {
|
||||||
|
if input.SubjectType == "" {
|
||||||
|
input.SubjectType = "principal"
|
||||||
|
}
|
||||||
tx, err := s.pool.Begin(ctx)
|
tx, err := s.pool.Begin(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ApplySecurityEventResult{}, err
|
return ApplySecurityEventResult{}, err
|
||||||
@@ -44,13 +49,15 @@ func (s *Store) ApplySessionRevoked(ctx context.Context, input ApplySessionRevok
|
|||||||
subjectHash := shortSecurityEventHash(input.Subject)
|
subjectHash := shortSecurityEventHash(input.Subject)
|
||||||
tag, err := tx.Exec(ctx, `
|
tag, err := tx.Exec(ctx, `
|
||||||
INSERT INTO gateway_security_event_receipts (
|
INSERT INTO gateway_security_event_receipts (
|
||||||
issuer, jti, audience, event_type, transaction_id, tenant_id, subject_hash, event_timestamp
|
issuer, jti, audience, event_type, transaction_id, tenant_id, subject_hash, event_timestamp,
|
||||||
|
application_id, subject_type
|
||||||
)
|
)
|
||||||
VALUES ($1, $2::uuid, $3, $4, NULLIF($5, ''), $6, $7, $8)
|
VALUES ($1, $2::uuid, $3, $4, NULLIF($5, ''), $6, $7, $8, NULLIF($9, ''), $10)
|
||||||
ON CONFLICT (issuer, jti) DO NOTHING`,
|
ON CONFLICT (issuer, jti) DO NOTHING`,
|
||||||
input.Issuer, input.JTI, input.Audience,
|
input.Issuer, input.JTI, input.Audience,
|
||||||
"https://schemas.openid.net/secevent/caep/event-type/session-revoked",
|
"https://schemas.openid.net/secevent/caep/event-type/session-revoked",
|
||||||
input.TransactionID, input.TenantID, subjectHash, input.EventTimestamp,
|
input.TransactionID, input.TenantID, subjectHash, input.EventTimestamp,
|
||||||
|
input.ApplicationID, input.SubjectType,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ApplySecurityEventResult{}, err
|
return ApplySecurityEventResult{}, err
|
||||||
@@ -64,25 +71,56 @@ ON CONFLICT (issuer, jti) DO NOTHING`,
|
|||||||
|
|
||||||
var revokedAt time.Time
|
var revokedAt time.Time
|
||||||
err = tx.QueryRow(ctx, `
|
err = tx.QueryRow(ctx, `
|
||||||
INSERT INTO gateway_oidc_revocation_watermarks (issuer, tenant_id, subject, revoked_at, source_jti)
|
INSERT INTO gateway_oidc_revocation_watermarks (
|
||||||
VALUES ($1, $2, $3, $4, $5::uuid)
|
issuer, application_id, tenant_id, subject_type, subject, revoked_at, source_jti
|
||||||
ON CONFLICT (issuer, tenant_id, subject) DO UPDATE
|
)
|
||||||
|
VALUES ($1, $2, $3, $4, $5, $6, $7::uuid)
|
||||||
|
ON CONFLICT (issuer, application_id, tenant_id, subject_type, subject) DO UPDATE
|
||||||
SET revoked_at = EXCLUDED.revoked_at, source_jti = EXCLUDED.source_jti, updated_at = now()
|
SET revoked_at = EXCLUDED.revoked_at, source_jti = EXCLUDED.source_jti, updated_at = now()
|
||||||
WHERE EXCLUDED.revoked_at > gateway_oidc_revocation_watermarks.revoked_at
|
WHERE EXCLUDED.revoked_at > gateway_oidc_revocation_watermarks.revoked_at
|
||||||
RETURNING revoked_at`, input.SubjectIssuer, input.TenantID, input.Subject, input.EventTimestamp, input.JTI).Scan(&revokedAt)
|
RETURNING revoked_at`, input.SubjectIssuer, input.ApplicationID, input.TenantID, input.SubjectType,
|
||||||
|
input.Subject, input.EventTimestamp, input.JTI).Scan(&revokedAt)
|
||||||
watermarkMoved := err == nil
|
watermarkMoved := err == nil
|
||||||
if err != nil && !errors.Is(err, pgx.ErrNoRows) {
|
if err != nil && !errors.Is(err, pgx.ErrNoRows) {
|
||||||
return ApplySecurityEventResult{}, err
|
return ApplySecurityEventResult{}, err
|
||||||
}
|
}
|
||||||
var sessionsDeleted int64
|
var sessionsDeleted int64
|
||||||
if watermarkMoved {
|
if watermarkMoved {
|
||||||
tag, err = tx.Exec(ctx, `
|
switch {
|
||||||
|
case input.ApplicationID != "" && input.SubjectType == "principal":
|
||||||
|
tag, err = tx.Exec(ctx, `
|
||||||
|
DELETE FROM gateway_oidc_sessions session
|
||||||
|
USING gateway_oidc_user_bindings user_binding,gateway_oidc_tenant_bindings tenant_binding
|
||||||
|
WHERE session.oidc_user_binding_id=user_binding.id
|
||||||
|
AND user_binding.tenant_binding_id=tenant_binding.id
|
||||||
|
AND tenant_binding.issuer=$1
|
||||||
|
AND tenant_binding.application_id=$2
|
||||||
|
AND tenant_binding.external_tenant_id=$3
|
||||||
|
AND user_binding.subject=$4`, input.SubjectIssuer, input.ApplicationID, input.TenantID, input.Subject)
|
||||||
|
case input.ApplicationID != "" && input.SubjectType == "tenant":
|
||||||
|
tag, err = tx.Exec(ctx, `
|
||||||
|
DELETE FROM gateway_oidc_sessions session
|
||||||
|
USING gateway_oidc_user_bindings user_binding,gateway_oidc_tenant_bindings tenant_binding
|
||||||
|
WHERE session.oidc_user_binding_id=user_binding.id
|
||||||
|
AND user_binding.tenant_binding_id=tenant_binding.id
|
||||||
|
AND tenant_binding.issuer=$1
|
||||||
|
AND tenant_binding.application_id=$2
|
||||||
|
AND tenant_binding.external_tenant_id=$3`, input.SubjectIssuer, input.ApplicationID, input.TenantID)
|
||||||
|
if err == nil {
|
||||||
|
_, err = tx.Exec(ctx, `UPDATE gateway_oidc_tenant_bindings
|
||||||
|
SET access_status='disabled',updated_at=now()
|
||||||
|
WHERE issuer=$1 AND application_id=$2 AND external_tenant_id=$3`,
|
||||||
|
input.SubjectIssuer, input.ApplicationID, input.TenantID)
|
||||||
|
}
|
||||||
|
default:
|
||||||
|
tag, err = tx.Exec(ctx, `
|
||||||
DELETE FROM gateway_oidc_sessions session
|
DELETE FROM gateway_oidc_sessions session
|
||||||
USING gateway_users gateway_user
|
USING gateway_users gateway_user
|
||||||
WHERE session.gateway_user_id = gateway_user.id
|
WHERE session.gateway_user_id = gateway_user.id
|
||||||
AND gateway_user.source = 'oidc'
|
AND gateway_user.source = 'oidc'
|
||||||
AND gateway_user.external_user_id = $1
|
AND gateway_user.external_user_id = $1
|
||||||
AND gateway_user.tenant_id = $2`, input.Subject, input.TenantID)
|
AND gateway_user.tenant_id = $2`, input.Subject, input.TenantID)
|
||||||
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ApplySecurityEventResult{}, err
|
return ApplySecurityEventResult{}, err
|
||||||
}
|
}
|
||||||
@@ -95,7 +133,7 @@ WHERE session.gateway_user_id = gateway_user.id
|
|||||||
AfterState: map[string]any{"watermarkMoved": watermarkMoved, "sessionsDeleted": sessionsDeleted},
|
AfterState: map[string]any{"watermarkMoved": watermarkMoved, "sessionsDeleted": sessionsDeleted},
|
||||||
Metadata: map[string]any{
|
Metadata: map[string]any{
|
||||||
"issuer": input.SubjectIssuer, "transmitterIssuer": input.Issuer, "tenantId": input.TenantID, "jtiHash": jtiHash,
|
"issuer": input.SubjectIssuer, "transmitterIssuer": input.Issuer, "tenantId": input.TenantID, "jtiHash": jtiHash,
|
||||||
"initiatingEntity": input.InitiatingEntity,
|
"applicationId": input.ApplicationID, "subjectType": input.SubjectType, "initiatingEntity": input.InitiatingEntity,
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -314,7 +352,7 @@ WHERE issuer=$1 AND audience=$2 AND mode <> 'introspection_fallback'
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Store) EvaluateOIDCSecurityEvent(ctx context.Context, transmitterIssuer, audience, subjectIssuer, tenantID, subject string, issuedAt, now time.Time, staleAfter time.Duration) (SecurityEventEvaluation, error) {
|
func (s *Store) EvaluateOIDCSecurityEvent(ctx context.Context, transmitterIssuer, audience, subjectIssuer, applicationID, tenantID, subject string, issuedAt, now time.Time, staleAfter time.Duration) (SecurityEventEvaluation, error) {
|
||||||
var mode, streamStatus string
|
var mode, streamStatus string
|
||||||
var lastVerification, bootstrapUntil *time.Time
|
var lastVerification, bootstrapUntil *time.Time
|
||||||
var createdAt time.Time
|
var createdAt time.Time
|
||||||
@@ -322,9 +360,17 @@ func (s *Store) EvaluateOIDCSecurityEvent(ctx context.Context, transmitterIssuer
|
|||||||
err := s.pool.QueryRow(ctx, `
|
err := s.pool.QueryRow(ctx, `
|
||||||
SELECT state.mode,state.stream_status,state.last_verification_at,state.bootstrap_until,state.created_at,watermark.revoked_at
|
SELECT state.mode,state.stream_status,state.last_verification_at,state.bootstrap_until,state.created_at,watermark.revoked_at
|
||||||
FROM gateway_security_event_stream_state state
|
FROM gateway_security_event_stream_state state
|
||||||
LEFT JOIN gateway_oidc_revocation_watermarks watermark
|
LEFT JOIN LATERAL (
|
||||||
ON watermark.issuer=$3 AND watermark.tenant_id=$4 AND watermark.subject=$5
|
SELECT max(revoked_at) AS revoked_at
|
||||||
WHERE state.issuer=$1 AND state.audience=$2`, transmitterIssuer, audience, subjectIssuer, tenantID, subject).Scan(
|
FROM gateway_oidc_revocation_watermarks
|
||||||
|
WHERE issuer=$3 AND application_id=$4 AND tenant_id=$5
|
||||||
|
AND (
|
||||||
|
(subject_type='principal' AND subject=$6)
|
||||||
|
OR (subject_type='tenant' AND subject=$5)
|
||||||
|
)
|
||||||
|
) watermark ON true
|
||||||
|
WHERE state.issuer=$1 AND state.audience=$2`,
|
||||||
|
transmitterIssuer, audience, subjectIssuer, applicationID, tenantID, subject).Scan(
|
||||||
&mode, &streamStatus, &lastVerification, &bootstrapUntil, &createdAt, &revokedAt,
|
&mode, &streamStatus, &lastVerification, &bootstrapUntil, &createdAt, &revokedAt,
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
"github.com/jackc/pgx/v5/pgxpool"
|
"github.com/jackc/pgx/v5/pgxpool"
|
||||||
)
|
)
|
||||||
@@ -55,7 +56,7 @@ func TestSecurityEventWatermarkVerificationAndFallbackLifecycle(t *testing.T) {
|
|||||||
if err := db.AdvanceSecurityEventStreamState(ctx, issuer, audience, now, 180*time.Second); err != nil {
|
if err := db.AdvanceSecurityEventStreamState(ctx, issuer, audience, now, 180*time.Second); err != nil {
|
||||||
t.Fatalf("advance fresh stream state: %v", err)
|
t.Fatalf("advance fresh stream state: %v", err)
|
||||||
}
|
}
|
||||||
evaluation, err := db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, tenantID, subject, now, now, 180*time.Second)
|
evaluation, err := db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, "", tenantID, subject, now, now, 180*time.Second)
|
||||||
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "bootstrap" {
|
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "bootstrap" {
|
||||||
t.Fatalf("bootstrap evaluation=%#v error=%v", evaluation, err)
|
t.Fatalf("bootstrap evaluation=%#v error=%v", evaluation, err)
|
||||||
}
|
}
|
||||||
@@ -77,12 +78,12 @@ func TestSecurityEventWatermarkVerificationAndFallbackLifecycle(t *testing.T) {
|
|||||||
}
|
}
|
||||||
confirmAt("first-verification-state", now)
|
confirmAt("first-verification-state", now)
|
||||||
confirmAt("second-verification-state", now.Add(time.Minute))
|
confirmAt("second-verification-state", now.Add(time.Minute))
|
||||||
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, tenantID, subject, now, now.Add(time.Minute), 180*time.Second)
|
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, "", tenantID, subject, now, now.Add(time.Minute), 180*time.Second)
|
||||||
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "introspection_fallback" {
|
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "introspection_fallback" {
|
||||||
t.Fatalf("bootstrap overlap evaluation=%#v error=%v", evaluation, err)
|
t.Fatalf("bootstrap overlap evaluation=%#v error=%v", evaluation, err)
|
||||||
}
|
}
|
||||||
confirmAt("post-bootstrap-verification-state", now.Add(361*time.Second))
|
confirmAt("post-bootstrap-verification-state", now.Add(361*time.Second))
|
||||||
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, tenantID, subject, now, now.Add(361*time.Second), 180*time.Second)
|
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, "", tenantID, subject, now, now.Add(361*time.Second), 180*time.Second)
|
||||||
if err != nil || evaluation.RequireIntrospection || evaluation.Mode != "push_healthy" {
|
if err != nil || evaluation.RequireIntrospection || evaluation.Mode != "push_healthy" {
|
||||||
t.Fatalf("healthy evaluation=%#v error=%v", evaluation, err)
|
t.Fatalf("healthy evaluation=%#v error=%v", evaluation, err)
|
||||||
}
|
}
|
||||||
@@ -97,7 +98,7 @@ func TestSecurityEventWatermarkVerificationAndFallbackLifecycle(t *testing.T) {
|
|||||||
}
|
}
|
||||||
confirmAt("paused-verification-one", now.Add(363*time.Second))
|
confirmAt("paused-verification-one", now.Add(363*time.Second))
|
||||||
confirmAt("paused-verification-two", now.Add(364*time.Second))
|
confirmAt("paused-verification-two", now.Add(364*time.Second))
|
||||||
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, tenantID, subject, now, now.Add(364*time.Second), 180*time.Second)
|
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, "", tenantID, subject, now, now.Add(364*time.Second), 180*time.Second)
|
||||||
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "introspection_fallback" {
|
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "introspection_fallback" {
|
||||||
t.Fatalf("stream update fallback=%#v error=%v", evaluation, err)
|
t.Fatalf("stream update fallback=%#v error=%v", evaluation, err)
|
||||||
}
|
}
|
||||||
@@ -114,22 +115,22 @@ func TestSecurityEventWatermarkVerificationAndFallbackLifecycle(t *testing.T) {
|
|||||||
if err != nil || !duplicate.Duplicate {
|
if err != nil || !duplicate.Duplicate {
|
||||||
t.Fatalf("duplicate result=%#v error=%v", duplicate, err)
|
t.Fatalf("duplicate result=%#v error=%v", duplicate, err)
|
||||||
}
|
}
|
||||||
evaluation, _ = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, tenantID, subject, revokedAt, now, 180*time.Second)
|
evaluation, _ = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, "", tenantID, subject, revokedAt, now, 180*time.Second)
|
||||||
if !evaluation.Revoked {
|
if !evaluation.Revoked {
|
||||||
t.Fatal("token at watermark was accepted")
|
t.Fatal("token at watermark was accepted")
|
||||||
}
|
}
|
||||||
evaluation, _ = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, "https://other.test/issuer", tenantID, subject, revokedAt, now, 180*time.Second)
|
evaluation, _ = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, "https://other.test/issuer", "", tenantID, subject, revokedAt, now, 180*time.Second)
|
||||||
if evaluation.Revoked {
|
if evaluation.Revoked {
|
||||||
t.Fatal("watermark crossed OIDC issuer boundary")
|
t.Fatal("watermark crossed OIDC issuer boundary")
|
||||||
}
|
}
|
||||||
evaluation, _ = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, tenantID, subject, now, now, 180*time.Second)
|
evaluation, _ = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, "", tenantID, subject, now, now, 180*time.Second)
|
||||||
if evaluation.Revoked {
|
if evaluation.Revoked {
|
||||||
t.Fatal("token after watermark was rejected")
|
t.Fatal("token after watermark was rejected")
|
||||||
}
|
}
|
||||||
|
|
||||||
_, _ = db.pool.Exec(ctx, `UPDATE gateway_security_event_stream_state SET last_verification_at=$3
|
_, _ = db.pool.Exec(ctx, `UPDATE gateway_security_event_stream_state SET last_verification_at=$3
|
||||||
WHERE issuer=$1 AND audience=$2`, issuer, audience, now.Add(-181*time.Second))
|
WHERE issuer=$1 AND audience=$2`, issuer, audience, now.Add(-181*time.Second))
|
||||||
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, tenantID, subject, now, now, 180*time.Second)
|
evaluation, err = db.EvaluateOIDCSecurityEvent(ctx, issuer, audience, subjectIssuer, "", tenantID, subject, now, now, 180*time.Second)
|
||||||
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "introspection_fallback" {
|
if err != nil || !evaluation.RequireIntrospection || evaluation.Mode != "introspection_fallback" {
|
||||||
t.Fatalf("fallback evaluation=%#v error=%v", evaluation, err)
|
t.Fatalf("fallback evaluation=%#v error=%v", evaluation, err)
|
||||||
}
|
}
|
||||||
@@ -183,3 +184,136 @@ func TestSecurityEventWatermarkVerificationAndFallbackLifecycle(t *testing.T) {
|
|||||||
t.Fatalf("rebind discarded security history receipts=%d watermarks=%d", retainedReceipts, retainedWatermarks)
|
t.Fatalf("rebind discarded security history receipts=%d watermarks=%d", retainedReceipts, retainedWatermarks)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestApplicationScopedSecurityEventsRevokeOnlyMatchingTenantSessions(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 application-scoped security event integration tests")
|
||||||
|
}
|
||||||
|
ctx := context.Background()
|
||||||
|
applyOIDCJITTestMigrations(t, ctx, databaseURL)
|
||||||
|
db, err := Connect(ctx, databaseURL)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
transmitterIssuer := "https://auth.test.example/ssf/" + uuid.NewString()
|
||||||
|
subjectIssuer := "https://auth.test.example/issuer/shared"
|
||||||
|
applicationID, tenantA, tenantB, subject := uuid.NewString(), uuid.NewString(), uuid.NewString(), uuid.NewString()
|
||||||
|
projection := func(tenantID string) *auth.User {
|
||||||
|
result, err := db.ResolveOrProvisionOIDCUser(ctx, ResolveOrProvisionOIDCUserInput{
|
||||||
|
Issuer: subjectIssuer, ApplicationID: applicationID, Subject: subject, Username: "shared-user",
|
||||||
|
Roles: []string{"basic"}, TenantID: tenantID, TenantMode: "multi_tenant",
|
||||||
|
TenantName: "Tenant " + tenantID[:8], TenantSlug: "tenant-" + tenantID[:8],
|
||||||
|
TenantMetadataStatus: "synced", TenantMetadataVersion: "1",
|
||||||
|
TenantMetadataUpdatedAt: time.Now().UTC(), OIDCClientID: "gateway-browser", ProvisioningEnabled: true,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return result.User
|
||||||
|
}
|
||||||
|
userA, userB := projection(tenantA), projection(tenantB)
|
||||||
|
if _, err := db.CreateAPIKey(ctx, CreateAPIKeyInput{Name: "must survive OIDC revocation"}, userA); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
createSession := func(user *auth.User, marker byte) {
|
||||||
|
now := time.Now().UTC()
|
||||||
|
sessionHash := sha256.Sum256([]byte(uuid.NewString()))
|
||||||
|
if _, err := db.CreateOIDCSession(ctx, CreateOIDCSessionInput{
|
||||||
|
SessionTokenHash: sessionHash[:], GatewayUserID: user.GatewayUserID,
|
||||||
|
GatewayTenantID: user.GatewayTenantID, OIDCUserBindingID: user.OIDCUserBindingID,
|
||||||
|
OIDCClientID: "gateway-browser", Issuer: subjectIssuer, ApplicationID: applicationID,
|
||||||
|
TenantID: user.TenantID, TokenCiphertext: []byte{marker},
|
||||||
|
AccessTokenExpiresAt: now.Add(time.Hour), LastSeenAt: now,
|
||||||
|
IdleExpiresAt: now.Add(time.Hour), AbsoluteExpiresAt: now.Add(2 * time.Hour),
|
||||||
|
}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
createSession(userA, 0xa1)
|
||||||
|
createSession(userB, 0xb1)
|
||||||
|
|
||||||
|
revokedAt := time.Now().UTC().Truncate(time.Second)
|
||||||
|
principalEvent := ApplySessionRevokedInput{
|
||||||
|
Issuer: transmitterIssuer, Audience: "urn:easyai:ssf:receiver:" + applicationID,
|
||||||
|
JTI: uuid.NewString(), TransactionID: uuid.NewString(), SubjectIssuer: subjectIssuer,
|
||||||
|
ApplicationID: applicationID, SubjectType: "principal", TenantID: tenantA, Subject: subject,
|
||||||
|
EventTimestamp: revokedAt, InitiatingEntity: "admin",
|
||||||
|
}
|
||||||
|
result, err := db.ApplySessionRevoked(ctx, principalEvent)
|
||||||
|
if err != nil || result.SessionsDeleted != 1 || !result.WatermarkMoved {
|
||||||
|
t.Fatalf("principal result=%#v error=%v", result, err)
|
||||||
|
}
|
||||||
|
var tenantASessions, tenantBSessions int
|
||||||
|
countSessions := func() {
|
||||||
|
t.Helper()
|
||||||
|
if err := db.pool.QueryRow(ctx, `SELECT
|
||||||
|
(SELECT count(*) FROM gateway_oidc_sessions WHERE gateway_tenant_id=$1::uuid),
|
||||||
|
(SELECT count(*) FROM gateway_oidc_sessions WHERE gateway_tenant_id=$2::uuid)`,
|
||||||
|
userA.GatewayTenantID, userB.GatewayTenantID,
|
||||||
|
).Scan(&tenantASessions, &tenantBSessions); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
countSessions()
|
||||||
|
if tenantASessions != 0 || tenantBSessions != 1 {
|
||||||
|
t.Fatalf("principal revocation sessions A=%d B=%d", tenantASessions, tenantBSessions)
|
||||||
|
}
|
||||||
|
older := principalEvent
|
||||||
|
older.JTI = uuid.NewString()
|
||||||
|
older.EventTimestamp = revokedAt.Add(-time.Minute)
|
||||||
|
result, err = db.ApplySessionRevoked(ctx, older)
|
||||||
|
if err != nil || result.WatermarkMoved || result.SessionsDeleted != 0 {
|
||||||
|
t.Fatalf("older result=%#v error=%v", result, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
createSession(userA, 0xa2)
|
||||||
|
tenantEvent := principalEvent
|
||||||
|
tenantEvent.JTI, tenantEvent.TransactionID = uuid.NewString(), uuid.NewString()
|
||||||
|
tenantEvent.SubjectType, tenantEvent.Subject = "tenant", tenantA
|
||||||
|
tenantEvent.EventTimestamp = revokedAt.Add(time.Minute)
|
||||||
|
result, err = db.ApplySessionRevoked(ctx, tenantEvent)
|
||||||
|
if err != nil || result.SessionsDeleted != 1 || !result.WatermarkMoved {
|
||||||
|
t.Fatalf("tenant result=%#v error=%v", result, err)
|
||||||
|
}
|
||||||
|
countSessions()
|
||||||
|
if tenantASessions != 0 || tenantBSessions != 1 {
|
||||||
|
t.Fatalf("tenant revocation sessions A=%d B=%d", tenantASessions, tenantBSessions)
|
||||||
|
}
|
||||||
|
keysA, err := db.ListAPIKeys(ctx, userA)
|
||||||
|
if err != nil || len(keysA) != 1 {
|
||||||
|
t.Fatalf("tenant revocation changed API keys: keys=%#v error=%v", keysA, err)
|
||||||
|
}
|
||||||
|
var accessStatus string
|
||||||
|
if err := db.pool.QueryRow(ctx, `SELECT access_status FROM gateway_oidc_tenant_bindings
|
||||||
|
WHERE issuer=$1 AND application_id=$2 AND external_tenant_id=$3`,
|
||||||
|
subjectIssuer, applicationID, tenantA,
|
||||||
|
).Scan(&accessStatus); err != nil || accessStatus != "disabled" {
|
||||||
|
t.Fatalf("tenant A access status=%q error=%v", accessStatus, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
streamID := uuid.NewString()
|
||||||
|
audience := tenantEvent.Audience
|
||||||
|
if err := db.EnsureSecurityEventStreamState(ctx, transmitterIssuer, audience, streamID); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_, _ = db.pool.Exec(ctx, `UPDATE gateway_security_event_stream_state
|
||||||
|
SET stream_status='enabled',mode='push_healthy',last_verification_at=$3,bootstrap_until=$3
|
||||||
|
WHERE issuer=$1 AND audience=$2`, transmitterIssuer, audience, revokedAt.Add(2*time.Minute))
|
||||||
|
evaluationA, err := db.EvaluateOIDCSecurityEvent(
|
||||||
|
ctx, transmitterIssuer, audience, subjectIssuer, applicationID, tenantA, subject,
|
||||||
|
tenantEvent.EventTimestamp, revokedAt.Add(2*time.Minute), 5*time.Minute,
|
||||||
|
)
|
||||||
|
if err != nil || !evaluationA.Revoked {
|
||||||
|
t.Fatalf("tenant A evaluation=%#v error=%v", evaluationA, err)
|
||||||
|
}
|
||||||
|
evaluationB, err := db.EvaluateOIDCSecurityEvent(
|
||||||
|
ctx, transmitterIssuer, audience, subjectIssuer, applicationID, tenantB, subject,
|
||||||
|
tenantEvent.EventTimestamp, revokedAt.Add(2*time.Minute), 5*time.Minute,
|
||||||
|
)
|
||||||
|
if err != nil || evaluationB.Revoked {
|
||||||
|
t.Fatalf("tenant B evaluation=%#v error=%v", evaluationB, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,92 @@
|
|||||||
|
ALTER TABLE gateway_identity_configuration_revisions
|
||||||
|
DROP CONSTRAINT IF EXISTS gateway_identity_configuration_revisions_schema_version_check;
|
||||||
|
|
||||||
|
ALTER TABLE gateway_identity_configuration_revisions
|
||||||
|
ADD COLUMN IF NOT EXISTS tenant_mode text NOT NULL DEFAULT 'single_tenant',
|
||||||
|
ADD COLUMN IF NOT EXISTS tenant_context_endpoint text,
|
||||||
|
ADD COLUMN IF NOT EXISTS tenant_context_audience text,
|
||||||
|
ADD COLUMN IF NOT EXISTS tenant_context_scope text,
|
||||||
|
ADD CONSTRAINT gateway_identity_configuration_revisions_schema_version_check
|
||||||
|
CHECK (schema_version IN (1,2)),
|
||||||
|
ADD CONSTRAINT gateway_identity_configuration_revisions_tenant_mode_check
|
||||||
|
CHECK (tenant_mode IN ('single_tenant','multi_tenant')),
|
||||||
|
ADD CONSTRAINT gateway_identity_configuration_revisions_manifest_mode_check
|
||||||
|
CHECK (
|
||||||
|
(state='draft' AND issuer IS NULL)
|
||||||
|
OR
|
||||||
|
(schema_version=1 AND tenant_mode='single_tenant' AND tenant_id IS NOT NULL AND local_tenant_key <> ''
|
||||||
|
AND tenant_context_endpoint IS NULL AND tenant_context_audience IS NULL AND tenant_context_scope IS NULL)
|
||||||
|
OR
|
||||||
|
(schema_version=2 AND tenant_mode='multi_tenant' AND tenant_id IS NULL AND local_tenant_key=''
|
||||||
|
AND tenant_context_endpoint IS NOT NULL AND tenant_context_audience='urn:easyai:auth-center:tenant-context'
|
||||||
|
AND tenant_context_scope='tenant.context.read')
|
||||||
|
) NOT VALID;
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS gateway_oidc_tenant_bindings (
|
||||||
|
id uuid PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
issuer text NOT NULL,
|
||||||
|
application_id text NOT NULL,
|
||||||
|
external_tenant_id text NOT NULL,
|
||||||
|
gateway_tenant_id uuid NOT NULL REFERENCES gateway_tenants(id) ON DELETE CASCADE,
|
||||||
|
access_status text NOT NULL DEFAULT 'active'
|
||||||
|
CHECK (access_status IN ('active','disabled')),
|
||||||
|
metadata_status text NOT NULL DEFAULT 'metadata_pending'
|
||||||
|
CHECK (metadata_status IN ('metadata_pending','synced','rejected')),
|
||||||
|
display_name text,
|
||||||
|
slug text,
|
||||||
|
metadata_version text,
|
||||||
|
metadata_etag text,
|
||||||
|
metadata_updated_at timestamptz,
|
||||||
|
last_sync_at timestamptz,
|
||||||
|
next_sync_at timestamptz NOT NULL DEFAULT now(),
|
||||||
|
sync_failure_count integer NOT NULL DEFAULT 0 CHECK (sync_failure_count >= 0),
|
||||||
|
last_error_category text,
|
||||||
|
created_at timestamptz NOT NULL DEFAULT now(),
|
||||||
|
updated_at timestamptz NOT NULL DEFAULT now(),
|
||||||
|
UNIQUE(issuer,application_id,external_tenant_id),
|
||||||
|
UNIQUE(id,gateway_tenant_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gateway_oidc_tenant_bindings_sync
|
||||||
|
ON gateway_oidc_tenant_bindings(next_sync_at)
|
||||||
|
WHERE access_status='active';
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS gateway_oidc_user_bindings (
|
||||||
|
id uuid PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
tenant_binding_id uuid NOT NULL REFERENCES gateway_oidc_tenant_bindings(id) ON DELETE CASCADE,
|
||||||
|
subject text NOT NULL,
|
||||||
|
gateway_user_id uuid NOT NULL REFERENCES gateway_users(id) ON DELETE CASCADE,
|
||||||
|
created_at timestamptz NOT NULL DEFAULT now(),
|
||||||
|
updated_at timestamptz NOT NULL DEFAULT now(),
|
||||||
|
UNIQUE(tenant_binding_id,subject),
|
||||||
|
UNIQUE(id,gateway_user_id)
|
||||||
|
);
|
||||||
|
|
||||||
|
ALTER TABLE gateway_oidc_sessions
|
||||||
|
ADD COLUMN IF NOT EXISTS oidc_user_binding_id uuid
|
||||||
|
REFERENCES gateway_oidc_user_bindings(id) ON DELETE CASCADE,
|
||||||
|
ADD COLUMN IF NOT EXISTS oidc_client_id text;
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_gateway_oidc_sessions_user_binding
|
||||||
|
ON gateway_oidc_sessions(oidc_user_binding_id,created_at DESC)
|
||||||
|
WHERE oidc_user_binding_id IS NOT NULL;
|
||||||
|
|
||||||
|
ALTER TABLE gateway_security_event_receipts
|
||||||
|
ADD COLUMN IF NOT EXISTS application_id text,
|
||||||
|
ADD COLUMN IF NOT EXISTS subject_type text
|
||||||
|
CHECK (subject_type IS NULL OR subject_type IN ('principal','tenant'));
|
||||||
|
|
||||||
|
ALTER TABLE gateway_oidc_revocation_watermarks
|
||||||
|
ADD COLUMN IF NOT EXISTS application_id text NOT NULL DEFAULT '',
|
||||||
|
ADD COLUMN IF NOT EXISTS subject_type text NOT NULL DEFAULT 'principal'
|
||||||
|
CHECK (subject_type IN ('principal','tenant'));
|
||||||
|
|
||||||
|
ALTER TABLE gateway_oidc_revocation_watermarks
|
||||||
|
DROP CONSTRAINT IF EXISTS gateway_oidc_revocation_watermarks_pkey;
|
||||||
|
|
||||||
|
ALTER TABLE gateway_oidc_revocation_watermarks
|
||||||
|
ADD PRIMARY KEY (issuer,application_id,tenant_id,subject_type,subject);
|
||||||
|
|
||||||
|
DROP INDEX IF EXISTS idx_gateway_oidc_revocation_watermarks_lookup;
|
||||||
|
CREATE INDEX idx_gateway_oidc_revocation_watermarks_lookup
|
||||||
|
ON gateway_oidc_revocation_watermarks(issuer,application_id,tenant_id,subject_type,subject,revoked_at);
|
||||||
@@ -25,6 +25,8 @@ export function AuthPanel(props: {
|
|||||||
oidcEnabled: boolean;
|
oidcEnabled: boolean;
|
||||||
onOIDCLogin: () => void;
|
onOIDCLogin: () => void;
|
||||||
}) {
|
}) {
|
||||||
|
const visibleTabs = props.oidcEnabled ? tabs.filter((tab) => tab.value !== 'register') : tabs;
|
||||||
|
const visibleMode = props.oidcEnabled && props.authMode === 'register' ? 'login' : props.authMode;
|
||||||
return (
|
return (
|
||||||
<section className="authShell" aria-label="登录">
|
<section className="authShell" aria-label="登录">
|
||||||
<Card className="authCard">
|
<Card className="authCard">
|
||||||
@@ -41,10 +43,10 @@ export function AuthPanel(props: {
|
|||||||
使用统一认证中心登录
|
使用统一认证中心登录
|
||||||
</Button>
|
</Button>
|
||||||
)}
|
)}
|
||||||
<Tabs value={props.authMode} tabs={tabs} onValueChange={props.onAuthModeChange} />
|
<Tabs value={visibleMode} tabs={visibleTabs} onValueChange={props.onAuthModeChange} />
|
||||||
{props.authMode === 'login' && <LoginFormView {...props} />}
|
{visibleMode === 'login' && <LoginFormView {...props} />}
|
||||||
{props.authMode === 'register' && <RegisterFormView {...props} />}
|
{visibleMode === 'register' && <RegisterFormView {...props} />}
|
||||||
{props.authMode === 'external' && <ExternalTokenForm {...props} />}
|
{visibleMode === 'external' && <ExternalTokenForm {...props} />}
|
||||||
</CardContent>
|
</CardContent>
|
||||||
</Card>
|
</Card>
|
||||||
</section>
|
</section>
|
||||||
|
|||||||
Reference in New Issue
Block a user