fix: 修复 OIDC 用户预配与跨标签页登录态

增加受控 JIT 本地用户投影,并使用 HttpOnly Cookie 在新标签页恢复认证中心登录态。补充错误语义、安全校验、自动化测试与回滚配置。
This commit is contained in:
2026-07-13 17:07:52 +08:00
parent 17b1f77e1d
commit a81a7b5200
37 changed files with 2694 additions and 179 deletions
+4
View File
@@ -30,6 +30,10 @@ func main() {
logger := slog.New(slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{
Level: cfg.LogLevel,
}))
if err := cfg.Validate(); err != nil {
logger.Error("invalid gateway configuration", "error", err)
os.Exit(1)
}
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer stop()
+175 -12
View File
@@ -3777,23 +3777,29 @@
"$ref": "#/definitions/httpapi.PlayableAPIKeyListResponse"
}
},
"400": {
"description": "Bad Request",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"401": {
"description": "Unauthorized",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -3826,11 +3832,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
},
@@ -3881,11 +3899,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -3912,23 +3942,29 @@
"$ref": "#/definitions/httpapi.AccessRuleListResponse"
}
},
"400": {
"description": "Bad Request",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"401": {
"description": "Unauthorized",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -4243,6 +4279,49 @@
}
}
},
"/api/v1/auth/oidc/session": {
"post": {
"description": "验证 Auth Center Access Token 后写入 HttpOnly 会话 Cookie;不会签发 Gateway JWT。",
"tags": [
"auth"
],
"summary": "建立 OIDC 浏览器会话",
"responses": {
"204": {
"description": "No Content"
},
"400": {
"description": "Bad Request",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"401": {
"description": "Unauthorized",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"404": {
"description": "Not Found",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
},
"delete": {
"tags": [
"auth"
],
"summary": "注销 OIDC 浏览器会话",
"responses": {
"204": {
"description": "No Content"
}
}
}
},
"/api/v1/auth/register": {
"post": {
"description": "在 standalone 或 hybrid 身份模式下创建本地用户,并返回 24 小时 JWT。",
@@ -4763,6 +4842,18 @@
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -5607,11 +5698,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -6224,11 +6327,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -6475,11 +6590,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -6521,11 +6648,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -6610,11 +6749,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
@@ -7666,11 +7817,23 @@
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"403": {
"description": "Forbidden",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"500": {
"description": "Internal Server Error",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
},
"503": {
"description": "Service Unavailable",
"schema": {
"$ref": "#/definitions/httpapi.ErrorEnvelope"
}
}
}
}
+117 -8
View File
@@ -4951,18 +4951,22 @@ paths:
description: OK
schema:
$ref: '#/definitions/httpapi.PlayableAPIKeyListResponse'
"400":
description: Bad Request
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"401":
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 列出 Playground API Key
@@ -4982,10 +4986,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 列出 API Key
@@ -5017,10 +5029,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 创建 API Key
@@ -5153,18 +5173,22 @@ paths:
description: OK
schema:
$ref: '#/definitions/httpapi.AccessRuleListResponse'
"400":
description: Bad Request
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"401":
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 列出 API Key 访问规则
@@ -5252,6 +5276,35 @@ paths:
summary: 本地登录
tags:
- auth
/api/v1/auth/oidc/session:
delete:
responses:
"204":
description: No Content
summary: 注销 OIDC 浏览器会话
tags:
- auth
post:
description: 验证 Auth Center Access Token 后写入 HttpOnly 会话 Cookie;不会签发 Gateway
JWT。
responses:
"204":
description: No Content
"400":
description: Bad Request
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"401":
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"404":
description: Not Found
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
summary: 建立 OIDC 浏览器会话
tags:
- auth
/api/v1/auth/register:
post:
consumes:
@@ -5589,6 +5642,14 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 获取当前用户
@@ -6135,10 +6196,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 列出任务
@@ -6532,10 +6601,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 列出任务
@@ -6691,10 +6768,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 获取当前用户组策略
@@ -6720,10 +6805,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 获取钱包摘要
@@ -6778,10 +6871,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 列出钱包交易
@@ -7473,10 +7574,18 @@ paths:
description: Unauthorized
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"403":
description: Forbidden
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"500":
description: Internal Server Error
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
"503":
description: Service Unavailable
schema:
$ref: '#/definitions/httpapi.ErrorEnvelope'
security:
- BearerAuth: []
summary: 列出任务
+44 -22
View File
@@ -22,6 +22,8 @@ import (
type Permission string
const (
OIDCSessionCookieName = "easyai_gateway_oidc_session"
PermissionPublic Permission = "public"
PermissionBasic Permission = "basic"
PermissionCreat Permission = "creat"
@@ -30,23 +32,24 @@ const (
)
type User struct {
ID string `json:"sub"`
Username string `json:"username"`
Roles []string `json:"role,omitempty"`
TenantID string `json:"tenantId,omitempty"`
GatewayTenantID string `json:"gatewayTenantId,omitempty"`
TenantKey string `json:"tenantKey,omitempty"`
SSOID string `json:"sso_id,omitempty"`
Source string `json:"source,omitempty"`
GatewayUserID string `json:"gatewayUserId,omitempty"`
UserGroupID string `json:"userGroupId,omitempty"`
UserGroupKey string `json:"userGroupKey,omitempty"`
UserGroupKeys []string `json:"userGroupKeys,omitempty"`
APIKeyID string `json:"apiKeyId,omitempty"`
APIKeySecret string `json:"apiKeySecret,omitempty"`
APIKeyName string `json:"apiKeyName,omitempty"`
APIKeyPrefix string `json:"apiKeyPrefix,omitempty"`
APIKeyScopes []string `json:"apiKeyScopes,omitempty"`
ID string `json:"sub"`
Username string `json:"username"`
Roles []string `json:"role,omitempty"`
TenantID string `json:"tenantId,omitempty"`
GatewayTenantID string `json:"gatewayTenantId,omitempty"`
TenantKey string `json:"tenantKey,omitempty"`
SSOID string `json:"sso_id,omitempty"`
Source string `json:"source,omitempty"`
GatewayUserID string `json:"gatewayUserId,omitempty"`
UserGroupID string `json:"userGroupId,omitempty"`
UserGroupKey string `json:"userGroupKey,omitempty"`
UserGroupKeys []string `json:"userGroupKeys,omitempty"`
APIKeyID string `json:"apiKeyId,omitempty"`
APIKeySecret string `json:"apiKeySecret,omitempty"`
APIKeyName string `json:"apiKeyName,omitempty"`
APIKeyPrefix string `json:"apiKeyPrefix,omitempty"`
APIKeyScopes []string `json:"apiKeyScopes,omitempty"`
TokenExpiresAt time.Time `json:"-"`
}
type contextKey string
@@ -84,6 +87,10 @@ func UserFromContext(ctx context.Context) (*User, bool) {
return user, ok
}
func WithUser(ctx context.Context, user *User) context.Context {
return context.WithValue(ctx, userContextKey, user)
}
func (a *Authenticator) Require(permission Permission, next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, err := a.Authenticate(r)
@@ -102,12 +109,13 @@ func (a *Authenticator) Require(permission Permission, next http.Handler) http.H
http.Error(w, "forbidden", http.StatusForbidden)
return
}
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), userContextKey, user)))
next.ServeHTTP(w, r.WithContext(WithUser(r.Context(), user)))
})
}
func (a *Authenticator) Authenticate(r *http.Request) (*User, error) {
token := extractBearer(r.Header.Get("Authorization"))
fromOIDCSessionCookie := false
if token == "" {
token = strings.TrimSpace(r.Header.Get("x-comfy-api-key"))
}
@@ -118,6 +126,15 @@ func (a *Authenticator) Authenticate(r *http.Request) (*User, error) {
token = strings.TrimSpace(r.URL.Query().Get("key"))
}
if token == "" {
if cookie, err := r.Cookie(OIDCSessionCookieName); err == nil {
token = strings.TrimSpace(cookie.Value)
fromOIDCSessionCookie = token != ""
}
}
if token == "" {
return nil, ErrUnauthorized
}
if fromOIDCSessionCookie && jwtAlgorithm(token) != "RS256" && jwtAlgorithm(token) != "ES256" {
return nil, ErrUnauthorized
}
if strings.HasPrefix(token, "sk-") {
@@ -125,10 +142,7 @@ func (a *Authenticator) Authenticate(r *http.Request) (*User, error) {
}
algorithm := jwtAlgorithm(token)
if algorithm == "RS256" || algorithm == "ES256" {
if a.OIDCVerifier == nil {
return nil, ErrUnauthorized
}
return a.OIDCVerifier.Verify(r.Context(), token)
return a.AuthenticateOIDCAccessToken(r.Context(), token)
}
if !a.LegacyJWTEnabled {
return nil, ErrUnauthorized
@@ -136,6 +150,14 @@ func (a *Authenticator) Authenticate(r *http.Request) (*User, error) {
return a.verifyJWT(token)
}
func (a *Authenticator) AuthenticateOIDCAccessToken(ctx context.Context, token string) (*User, error) {
algorithm := jwtAlgorithm(token)
if a.OIDCVerifier == nil || algorithm != "RS256" && algorithm != "ES256" {
return nil, ErrUnauthorized
}
return a.OIDCVerifier.Verify(ctx, token)
}
func (a *Authenticator) verifyJWT(tokenString string) (*User, error) {
token, err := jwt.ParseWithClaims(tokenString, jwt.MapClaims{}, func(token *jwt.Token) (any, error) {
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
+5 -1
View File
@@ -117,6 +117,10 @@ func (v *OIDCVerifier) Verify(ctx context.Context, raw string) (*User, error) {
if !ok || stringClaim(claims, "sub") == "" || stringClaim(claims, "tid") != v.config.TenantID {
return nil, oidcUnauthorized("stable identity claims are invalid", nil)
}
expiresAt, ok := numericDateClaim(claims["exp"])
if !ok {
return nil, oidcUnauthorized("exp is invalid", nil)
}
if _, ok := numericDateClaim(claims["nbf"]); !ok {
return nil, oidcUnauthorized("nbf is missing", nil)
}
@@ -143,7 +147,7 @@ func (v *OIDCVerifier) Verify(ctx context.Context, raw string) (*User, error) {
}
return &User{
ID: stringClaim(claims, "sub"), Username: username, Roles: roles,
TenantID: v.config.TenantID, Source: "oidc",
TenantID: v.config.TenantID, Source: "oidc", TokenExpiresAt: expiresAt,
}, nil
}
@@ -0,0 +1,86 @@
package auth
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
func TestAuthenticateAcceptsValidatedOIDCSessionCookie(t *testing.T) {
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
var issuer string
issuerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/.well-known/openid-configuration":
_ = json.NewEncoder(w).Encode(map[string]any{"issuer": issuer, "jwks_uri": issuer + "/jwks"})
case "/jwks":
_ = json.NewEncoder(w).Encode(map[string]any{"keys": []any{ecJWK("ec-key", &key.PublicKey)}})
default:
http.NotFound(w, r)
}
}))
defer issuerServer.Close()
issuer = issuerServer.URL
verifier, err := NewOIDCVerifier(OIDCConfig{
Issuer: issuer, Audience: "gateway-api", TenantID: "tenant-1", RolePrefix: "gateway.",
RequiredScopes: []string{"gateway.access"}, HTTPClient: issuerServer.Client(),
})
if err != nil {
t.Fatal(err)
}
authenticator := New("local-jwt-secret", "", "")
authenticator.OIDCVerifier = verifier
raw := signedOIDCToken(t, issuer, "ec-key", jwt.SigningMethodES256, key, nil)
request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil)
request.AddCookie(&http.Cookie{Name: OIDCSessionCookieName, Value: raw})
user, err := authenticator.Authenticate(request)
if err != nil {
t.Fatalf("authenticate OIDC session cookie: %v", err)
}
if user.ID != "platform-subject" || user.Source != "oidc" {
t.Fatalf("unexpected session user: %#v", user)
}
if user.TokenExpiresAt.Before(time.Now().Add(50 * time.Minute)) {
t.Fatalf("token expiry was not retained: %v", user.TokenExpiresAt)
}
}
func TestAuthenticateBearerTakesPrecedenceOverOIDCSessionCookie(t *testing.T) {
authenticator := New("local-jwt-secret", "", "")
localToken, err := authenticator.SignJWT(&User{ID: "local-user", Source: "gateway", Roles: []string{"user"}}, time.Hour)
if err != nil {
t.Fatal(err)
}
request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil)
request.Header.Set("Authorization", "Bearer "+localToken)
request.AddCookie(&http.Cookie{Name: OIDCSessionCookieName, Value: "not-a-token"})
user, err := authenticator.Authenticate(request)
if err != nil {
t.Fatalf("authenticate bearer token: %v", err)
}
if user.ID != "local-user" || user.Source != "gateway" {
t.Fatalf("cookie overrode explicit bearer credentials: %#v", user)
}
}
func TestAuthenticateRejectsInvalidOIDCSessionCookie(t *testing.T) {
authenticator := New("local-jwt-secret", "", "")
request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil)
request.AddCookie(&http.Cookie{Name: OIDCSessionCookieName, Value: "not-a-token"})
if _, err := authenticator.Authenticate(request); err == nil {
t.Fatal("invalid OIDC session cookie was accepted")
}
}
+45 -7
View File
@@ -1,6 +1,7 @@
package config
import (
"errors"
"log/slog"
"net/url"
"os"
@@ -34,6 +35,10 @@ type Config struct {
OIDCIntrospectionEnabled bool
OIDCIntrospectionClientID string
OIDCIntrospectionClientSecret string
OIDCJITProvisioningEnabled bool
OIDCGatewayTenantKey string
OIDCBrowserSessionEnabled bool
OIDCSessionCookieSecure bool
PublicBaseURL string
WebBaseURL string
LocalGeneratedStorageDir string
@@ -51,8 +56,9 @@ type Config struct {
func Load() Config {
globalProxy := LoadGlobalHTTPProxyStatus()
appEnv := env("APP_ENV", "development")
return Config{
AppEnv: env("APP_ENV", "development"),
AppEnv: appEnv,
HTTPAddr: env("HTTP_ADDR", ":8088"),
DatabaseURL: gatewayDatabaseURL(),
IdentityMode: env("IDENTITY_MODE", "hybrid"),
@@ -75,12 +81,18 @@ func Load() Config {
OIDCIntrospectionEnabled: env("OIDC_INTROSPECTION_ENABLED", "false") == "true",
OIDCIntrospectionClientID: env("OIDC_INTROSPECTION_CLIENT_ID", ""),
OIDCIntrospectionClientSecret: env("OIDC_INTROSPECTION_CLIENT_SECRET", ""),
PublicBaseURL: strings.TrimRight(env("AI_GATEWAY_PUBLIC_BASE_URL", env("PUBLIC_BASE_URL", "")), "/"),
WebBaseURL: strings.TrimRight(env("AI_GATEWAY_WEB_BASE_URL", env("GATEWAY_WEB_BASE_URL", env("PUBLIC_WEB_BASE_URL", ""))), "/"),
LocalGeneratedStorageDir: env("AI_GATEWAY_GENERATED_STORAGE_DIR", env("LOCAL_GENERATED_STORAGE_DIR", env("AI_GATEWAY_STATIC_STORAGE_DIR", DefaultLocalGeneratedStorageDir))),
LocalUploadedStorageDir: env("AI_GATEWAY_UPLOADED_STORAGE_DIR", env("LOCAL_UPLOADED_STORAGE_DIR", DefaultLocalUploadedStorageDir)),
LocalTempAssetTTLHours: envInt("AI_GATEWAY_LOCAL_TEMP_ASSET_TTL_HOURS", 24),
TaskProgressCallbackEnabled: env("TASK_PROGRESS_CALLBACK_ENABLED", "true") == "true",
OIDCJITProvisioningEnabled: env("OIDC_JIT_PROVISIONING_ENABLED", "false") == "true",
OIDCGatewayTenantKey: env("OIDC_GATEWAY_TENANT_KEY", ""),
OIDCBrowserSessionEnabled: env("OIDC_BROWSER_SESSION_ENABLED", "true") == "true",
OIDCSessionCookieSecure: env("OIDC_SESSION_COOKIE_SECURE",
strconv.FormatBool(!isLocalEnvironment(appEnv)),
) == "true",
PublicBaseURL: strings.TrimRight(env("AI_GATEWAY_PUBLIC_BASE_URL", env("PUBLIC_BASE_URL", "")), "/"),
WebBaseURL: strings.TrimRight(env("AI_GATEWAY_WEB_BASE_URL", env("GATEWAY_WEB_BASE_URL", env("PUBLIC_WEB_BASE_URL", ""))), "/"),
LocalGeneratedStorageDir: env("AI_GATEWAY_GENERATED_STORAGE_DIR", env("LOCAL_GENERATED_STORAGE_DIR", env("AI_GATEWAY_STATIC_STORAGE_DIR", DefaultLocalGeneratedStorageDir))),
LocalUploadedStorageDir: env("AI_GATEWAY_UPLOADED_STORAGE_DIR", env("LOCAL_UPLOADED_STORAGE_DIR", DefaultLocalUploadedStorageDir)),
LocalTempAssetTTLHours: envInt("AI_GATEWAY_LOCAL_TEMP_ASSET_TTL_HOURS", 24),
TaskProgressCallbackEnabled: env("TASK_PROGRESS_CALLBACK_ENABLED", "true") == "true",
TaskProgressCallbackURL: env("TASK_PROGRESS_CALLBACK_URL",
strings.TrimRight(env("SERVER_MAIN_BASE_URL", "http://localhost:3000"), "/")+"/internal/platform/task-progress-callbacks",
),
@@ -93,6 +105,32 @@ func Load() Config {
}
}
func (c Config) Validate() error {
if c.OIDCJITProvisioningEnabled && strings.TrimSpace(c.OIDCGatewayTenantKey) == "" {
return errors.New("OIDC_GATEWAY_TENANT_KEY is required when OIDC_JIT_PROVISIONING_ENABLED=true")
}
if c.OIDCEnabled && c.OIDCBrowserSessionEnabled {
if !isLocalEnvironment(c.AppEnv) && !c.OIDCSessionCookieSecure {
return errors.New("OIDC_SESSION_COOKIE_SECURE must be true outside local development and tests")
}
for _, origin := range strings.Split(c.CORSAllowedOrigin, ",") {
if strings.TrimSpace(origin) == "*" {
return errors.New("CORS_ALLOWED_ORIGIN cannot contain * when OIDC browser sessions are enabled")
}
}
}
return nil
}
func isLocalEnvironment(value string) bool {
switch strings.ToLower(strings.TrimSpace(value)) {
case "development", "dev", "local", "test":
return true
default:
return false
}
}
type GlobalHTTPProxyStatus struct {
HTTPProxy string
Source string
+81
View File
@@ -0,0 +1,81 @@
package config
import (
"strings"
"testing"
)
func TestLoadOIDCJITProvisioningDefaultsToDisabled(t *testing.T) {
t.Setenv("OIDC_JIT_PROVISIONING_ENABLED", "")
t.Setenv("OIDC_GATEWAY_TENANT_KEY", "")
cfg := Load()
if cfg.OIDCJITProvisioningEnabled {
t.Fatal("OIDC JIT provisioning must be disabled by default")
}
if cfg.OIDCGatewayTenantKey != "" {
t.Fatalf("unexpected gateway tenant key: %q", cfg.OIDCGatewayTenantKey)
}
}
func TestValidateRequiresGatewayTenantKeyWhenOIDCJITIsEnabled(t *testing.T) {
cfg := Config{
OIDCEnabled: true,
OIDCJITProvisioningEnabled: true,
}
err := cfg.Validate()
if err == nil || !strings.Contains(err.Error(), "OIDC_GATEWAY_TENANT_KEY") {
t.Fatalf("Validate() error = %v, want missing OIDC_GATEWAY_TENANT_KEY", err)
}
cfg.OIDCGatewayTenantKey = "default"
if err := cfg.Validate(); err != nil {
t.Fatalf("Validate() with tenant key: %v", err)
}
}
func TestLoadOIDCBrowserSessionUsesSafeEnvironmentDefaults(t *testing.T) {
t.Setenv("APP_ENV", "development")
t.Setenv("OIDC_BROWSER_SESSION_ENABLED", "")
t.Setenv("OIDC_SESSION_COOKIE_SECURE", "")
cfg := Load()
if !cfg.OIDCBrowserSessionEnabled {
t.Fatal("OIDC browser session should be enabled by default")
}
if cfg.OIDCSessionCookieSecure {
t.Fatal("development cookie should allow localhost HTTP by default")
}
t.Setenv("APP_ENV", "production")
cfg = Load()
if !cfg.OIDCSessionCookieSecure {
t.Fatal("production OIDC session cookie must default to Secure")
}
t.Setenv("APP_ENV", "staging")
cfg = Load()
if !cfg.OIDCSessionCookieSecure {
t.Fatal("staging OIDC session cookie must default to Secure")
}
}
func TestValidateRejectsInsecureNonLocalOIDCBrowserSession(t *testing.T) {
cfg := Config{
AppEnv: "staging",
OIDCEnabled: true,
OIDCBrowserSessionEnabled: true,
OIDCSessionCookieSecure: false,
CORSAllowedOrigin: "https://gateway.example.com",
}
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "OIDC_SESSION_COOKIE_SECURE") {
t.Fatalf("Validate() error = %v, want insecure non-local cookie rejection", err)
}
cfg.OIDCSessionCookieSecure = true
cfg.CORSAllowedOrigin = "*"
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "CORS_ALLOWED_ORIGIN") {
t.Fatalf("Validate() error = %v, want wildcard credentialed CORS rejection", err)
}
}
@@ -38,8 +38,9 @@ func (s *Server) listAccessRules(w http.ResponseWriter, r *http.Request) {
// @Produce json
// @Security BearerAuth
// @Success 200 {object} AccessRuleListResponse
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/v1/api-keys/access-rules [get]
func (s *Server) listAPIKeyAccessRules(w http.ResponseWriter, r *http.Request) {
@@ -47,7 +48,7 @@ func (s *Server) listAPIKeyAccessRules(w http.ResponseWriter, r *http.Request) {
items, err := s.store.ListAPIKeyAccessRules(r.Context(), user)
if err != nil {
if errors.Is(err, store.ErrLocalUserRequired) {
writeError(w, http.StatusBadRequest, err.Error())
writeLocalUserRequired(w)
return
}
s.logger.Error("list api key access rules failed", "error", err)
@@ -157,7 +158,7 @@ func (s *Server) batchAPIKeyAccessRules(w http.ResponseWriter, r *http.Request)
items, err := s.store.BatchAPIKeyAccessRules(r.Context(), input, user)
if err != nil {
if errors.Is(err, store.ErrLocalUserRequired) {
writeError(w, http.StatusBadRequest, err.Error())
writeLocalUserRequired(w)
return
}
if store.IsNotFound(err) {
+1 -1
View File
@@ -45,7 +45,7 @@ var geminiGenerateContentRoutePrefixes = []string{
}
func (s *Server) registerGeminiGenerateContentRoutes(mux *http.ServeMux) {
handler := s.auth.Require(auth.PermissionBasic, http.HandlerFunc(s.geminiGenerateContent))
handler := s.requireUser(auth.PermissionBasic, http.HandlerFunc(s.geminiGenerateContent))
for _, prefix := range geminiGenerateContentRoutePrefixes {
mux.Handle("POST "+prefix, geminiGenerateContentRouteHandler(prefix, handler))
}
+21 -6
View File
@@ -57,6 +57,8 @@ func (s *Server) ready(w http.ResponseWriter, r *http.Request) {
// @Security BearerAuth
// @Success 200 {object} auth.User
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Router /api/v1/me [get]
func (s *Server) me(w http.ResponseWriter, r *http.Request) {
user, _ := auth.UserFromContext(r.Context())
@@ -630,6 +632,8 @@ func (s *Server) listUserGroups(w http.ResponseWriter, r *http.Request) {
// @Security BearerAuth
// @Success 200 {object} UserGroupListResponse
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/workspace/user-groups [get]
func (s *Server) listCurrentUserGroups(w http.ResponseWriter, r *http.Request) {
@@ -669,6 +673,8 @@ func compactAuthStrings(values ...string) []string {
// @Security BearerAuth
// @Success 200 {object} APIKeyListResponse
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/v1/api-keys [get]
func (s *Server) listAPIKeys(w http.ResponseWriter, r *http.Request) {
@@ -689,8 +695,9 @@ func (s *Server) listAPIKeys(w http.ResponseWriter, r *http.Request) {
// @Produce json
// @Security BearerAuth
// @Success 200 {object} PlayableAPIKeyListResponse
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/playground/api-keys [get]
func (s *Server) listPlayableAPIKeys(w http.ResponseWriter, r *http.Request) {
@@ -698,7 +705,7 @@ func (s *Server) listPlayableAPIKeys(w http.ResponseWriter, r *http.Request) {
items, err := s.store.ListPlayableAPIKeys(r.Context(), user)
if err != nil {
if errors.Is(err, store.ErrLocalUserRequired) {
writeError(w, http.StatusBadRequest, err.Error())
writeLocalUserRequired(w)
return
}
s.logger.Error("list playable api keys failed", "error", err)
@@ -719,6 +726,8 @@ func (s *Server) listPlayableAPIKeys(w http.ResponseWriter, r *http.Request) {
// @Success 201 {object} store.CreatedAPIKey
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/v1/api-keys [post]
func (s *Server) createAPIKey(w http.ResponseWriter, r *http.Request) {
@@ -731,7 +740,7 @@ func (s *Server) createAPIKey(w http.ResponseWriter, r *http.Request) {
created, err := s.store.CreateAPIKey(r.Context(), input, user)
if err != nil {
if errors.Is(err, store.ErrLocalUserRequired) {
writeError(w, http.StatusBadRequest, err.Error())
writeLocalUserRequired(w)
return
}
s.logger.Error("create api key failed", "error", err)
@@ -768,7 +777,11 @@ func (s *Server) updateAPIKeyScopes(w http.ResponseWriter, r *http.Request) {
writeJSON(w, http.StatusOK, item)
return
}
if errors.Is(err, store.ErrLocalUserRequired) || errors.Is(err, store.ErrInvalidAPIKeyScopes) {
if errors.Is(err, store.ErrLocalUserRequired) {
writeLocalUserRequired(w)
return
}
if errors.Is(err, store.ErrInvalidAPIKeyScopes) {
writeError(w, http.StatusBadRequest, err.Error())
return
}
@@ -801,7 +814,7 @@ func (s *Server) disableAPIKey(w http.ResponseWriter, r *http.Request) {
return
}
if errors.Is(err, store.ErrLocalUserRequired) {
writeError(w, http.StatusBadRequest, err.Error())
writeLocalUserRequired(w)
return
}
if store.IsNotFound(err) {
@@ -833,7 +846,7 @@ func (s *Server) deleteAPIKey(w http.ResponseWriter, r *http.Request) {
return
}
if errors.Is(err, store.ErrLocalUserRequired) {
writeError(w, http.StatusBadRequest, err.Error())
writeLocalUserRequired(w)
return
}
if store.IsNotFound(err) {
@@ -1505,6 +1518,8 @@ func matchedRateLimitRule(policy map[string]any, metric string) map[string]any {
// @Success 200 {object} TaskListResponse
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/workspace/tasks [get]
// @Router /api/v1/tasks [get]
@@ -0,0 +1,328 @@
package httpapi
import (
"bytes"
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"encoding/base64"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"os"
"strings"
"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/store"
"github.com/golang-jwt/jwt/v5"
)
func TestOIDCJITFullGatewayUserFlowAndSecurityBoundary(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run OIDC JIT HTTP integration tests")
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
applyMigration(t, ctx, databaseURL)
db, err := store.Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect store: %v", err)
}
defer db.Close()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("generate test signing key: %v", err)
}
var issuer string
issuerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/.well-known/openid-configuration":
_ = json.NewEncoder(w).Encode(map[string]any{"issuer": issuer, "jwks_uri": issuer + "/jwks"})
case "/jwks":
_ = json.NewEncoder(w).Encode(map[string]any{"keys": []any{oidcJITECJWK("jit-key", &key.PublicKey)}})
default:
http.NotFound(w, r)
}
}))
defer issuerServer.Close()
issuer = issuerServer.URL
suffix := time.Now().UTC().Format("20060102150405.000000000")
validSubject := "platform-http-jit-" + suffix
rejectedSubjects := []string{
"platform-http-scope-" + suffix,
"platform-http-role-" + suffix,
"platform-http-tenant-" + suffix,
"platform-http-disabled-jit-" + suffix,
"platform-http-missing-tenant-" + suffix,
}
allSubjects := append([]string{validSubject}, rejectedSubjects...)
t.Cleanup(func() {
_, _ = db.Pool().Exec(context.Background(), `
DELETE FROM gateway_audit_logs
WHERE target_id IN (SELECT id::text FROM gateway_users WHERE source = 'oidc' AND external_user_id = ANY($1::text[]));
DELETE FROM gateway_users WHERE source = 'oidc' AND external_user_id = ANY($1::text[]);`, allSubjects)
})
baseConfig := config.Config{
AppEnv: "test",
HTTPAddr: ":0",
DatabaseURL: databaseURL,
IdentityMode: "hybrid",
JWTSecret: "test-only-jwt-secret",
OIDCEnabled: true,
OIDCIssuer: issuer,
OIDCAudience: "gateway-api",
OIDCTenantID: "auth-center-test-tenant",
OIDCRolePrefix: "gateway.",
OIDCRequiredScopes: []string{"gateway.access"},
OIDCJWKSCacheTTLSeconds: 60,
OIDCAcceptLegacyHS256: true,
OIDCJITProvisioningEnabled: true,
OIDCGatewayTenantKey: "default",
OIDCBrowserSessionEnabled: true,
OIDCSessionCookieSecure: false,
LocalGeneratedStorageDir: t.TempDir(),
LocalUploadedStorageDir: t.TempDir(),
LocalTempAssetTTLHours: 1,
CORSAllowedOrigin: "http://localhost:5178",
TaskProgressCallbackEnabled: false,
}
server := httptest.NewServer(NewServerWithContext(ctx, baseConfig, db, slog.New(slog.NewTextHandler(io.Discard, nil))))
defer server.Close()
validToken := signedOIDCJITToken(t, key, issuer, validSubject, nil)
var me auth.User
doOIDCJITJSON(t, server.URL, http.MethodGet, "/api/v1/me", validToken, nil, http.StatusOK, &me)
if me.ID != validSubject || me.Source != "oidc" || me.GatewayUserID == "" || me.GatewayTenantID == "" || me.TenantKey != "default" || me.UserGroupID == "" {
t.Fatalf("OIDC /me did not include the local Gateway projection")
}
sessionCookie := createOIDCSessionCookie(t, server.URL, validToken)
request, err := http.NewRequest(http.MethodGet, server.URL+"/api/v1/me", nil)
if err != nil {
t.Fatal(err)
}
request.AddCookie(sessionCookie)
response, err := http.DefaultClient.Do(request)
if err != nil {
t.Fatalf("execute cookie-authenticated /me: %v", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
t.Fatalf("cookie-authenticated /me status = %d, want 200", response.StatusCode)
}
var cookieMe auth.User
if err := json.NewDecoder(response.Body).Decode(&cookieMe); err != nil {
t.Fatalf("decode cookie-authenticated /me: %v", err)
}
if cookieMe.GatewayUserID != me.GatewayUserID || cookieMe.ID != me.ID {
t.Fatalf("new-tab cookie resolved a different Gateway user: %#v", cookieMe)
}
var automaticallyCreatedAPIKeys int
if err := db.Pool().QueryRow(ctx, `SELECT count(*) FROM gateway_api_keys WHERE gateway_user_id = $1::uuid`, me.GatewayUserID).Scan(&automaticallyCreatedAPIKeys); err != nil {
t.Fatalf("count pre-created API keys: %v", err)
}
if automaticallyCreatedAPIKeys != 0 {
t.Fatalf("OIDC JIT created %d API keys before explicit user action", automaticallyCreatedAPIKeys)
}
for _, path := range []string{
"/api/workspace/user-groups",
"/api/workspace/wallet",
"/api/workspace/tasks",
"/api/v1/api-keys",
} {
doOIDCJITJSON(t, server.URL, http.MethodGet, path, validToken, nil, http.StatusOK, nil)
}
var createdKey struct {
Secret string `json:"secret"`
APIKey struct {
ID string `json:"id"`
} `json:"apiKey"`
}
doOIDCJITJSON(t, server.URL, http.MethodPost, "/api/v1/api-keys", validToken, map[string]any{"name": "OIDC JIT integration key"}, http.StatusCreated, &createdKey)
if createdKey.Secret == "" || createdKey.APIKey.ID == "" {
t.Fatal("OIDC user API Key creation returned incomplete data")
}
doOIDCJITJSON(t, server.URL, http.MethodGet, "/api/v1/api-keys", validToken, nil, http.StatusOK, nil)
doOIDCJITJSON(t, server.URL, http.MethodDelete, "/api/v1/api-keys/"+createdKey.APIKey.ID, validToken, nil, http.StatusNoContent, nil)
var users struct {
Items []store.GatewayUser `json:"items"`
}
doOIDCJITJSON(t, server.URL, http.MethodGet, "/api/admin/users", validToken, nil, http.StatusOK, &users)
foundOIDCUser := false
for _, user := range users.Items {
if user.ID == me.GatewayUserID {
foundOIDCUser = user.Source == "oidc" && user.ExternalUserID == validSubject
break
}
}
if !foundOIDCUser {
t.Fatal("admin user list did not expose the OIDC Gateway projection")
}
negativeTokens := []struct {
subject string
mutate func(jwt.MapClaims)
}{
{rejectedSubjects[0], func(claims jwt.MapClaims) { claims["scope"] = "openid" }},
{rejectedSubjects[1], func(claims jwt.MapClaims) { claims["roles"] = []string{"other.admin"} }},
{rejectedSubjects[2], func(claims jwt.MapClaims) { claims["tid"] = "wrong-tenant" }},
}
for _, negative := range negativeTokens {
token := signedOIDCJITToken(t, key, issuer, negative.subject, negative.mutate)
doOIDCJITJSON(t, server.URL, http.MethodGet, "/api/v1/me", token, nil, http.StatusUnauthorized, nil)
}
var rejectedWrites int
if err := db.Pool().QueryRow(ctx, `SELECT count(*) FROM gateway_users WHERE source = 'oidc' AND external_user_id = ANY($1::text[])`, rejectedSubjects[:3]).Scan(&rejectedWrites); err != nil {
t.Fatalf("count rejected OIDC writes: %v", err)
}
if rejectedWrites != 0 {
t.Fatalf("rejected OIDC tokens created %d local users", rejectedWrites)
}
disabledJITConfig := baseConfig
disabledJITConfig.OIDCJITProvisioningEnabled = false
disabledJITServer := httptest.NewServer(NewServerWithContext(ctx, disabledJITConfig, db, slog.New(slog.NewTextHandler(io.Discard, nil))))
defer disabledJITServer.Close()
assertOIDCJITError(t, disabledJITServer.URL, signedOIDCJITToken(t, key, issuer, rejectedSubjects[3], nil), http.StatusForbidden, errorCodeGatewayUserNotProvisioned)
missingTenantConfig := baseConfig
missingTenantConfig.OIDCGatewayTenantKey = "missing-tenant-" + suffix
missingTenantServer := httptest.NewServer(NewServerWithContext(ctx, missingTenantConfig, db, slog.New(slog.NewTextHandler(io.Discard, nil))))
defer missingTenantServer.Close()
assertOIDCJITError(t, missingTenantServer.URL, signedOIDCJITToken(t, key, issuer, rejectedSubjects[4], nil), http.StatusServiceUnavailable, errorCodeGatewayTenantUnavailable)
if _, err := db.Pool().Exec(ctx, `UPDATE gateway_users SET status = 'disabled' WHERE id = $1::uuid`, me.GatewayUserID); err != nil {
t.Fatalf("disable projected user: %v", err)
}
assertOIDCJITError(t, server.URL, validToken, http.StatusForbidden, errorCodeGatewayUserDisabled)
if _, err := db.Pool().Exec(ctx, `UPDATE gateway_users SET status = 'active' WHERE id = $1::uuid`, me.GatewayUserID); err != nil {
t.Fatalf("restore projected user for delete test: %v", err)
}
if err := db.DeleteGatewayUser(ctx, me.GatewayUserID); err != nil {
t.Fatalf("delete projected user: %v", err)
}
assertOIDCJITError(t, server.URL, validToken, http.StatusForbidden, errorCodeGatewayUserDisabled)
}
func createOIDCSessionCookie(t *testing.T, baseURL string, token string) *http.Cookie {
t.Helper()
request, err := http.NewRequest(http.MethodPost, baseURL+"/api/v1/auth/oidc/session", nil)
if err != nil {
t.Fatal(err)
}
request.Header.Set("Authorization", "Bearer "+token)
request.Header.Set("Origin", "http://localhost:5178")
response, err := http.DefaultClient.Do(request)
if err != nil {
t.Fatalf("create OIDC browser session: %v", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusNoContent {
t.Fatalf("create OIDC browser session status = %d, want 204", response.StatusCode)
}
for _, cookie := range response.Cookies() {
if cookie.Name == auth.OIDCSessionCookieName {
if !cookie.HttpOnly || cookie.SameSite != http.SameSiteStrictMode {
t.Fatalf("unsafe OIDC browser session cookie: %#v", cookie)
}
return cookie
}
}
t.Fatal("OIDC browser session cookie was not returned")
return nil
}
func assertOIDCJITError(t *testing.T, baseURL string, token string, expectedStatus int, expectedCode string) {
t.Helper()
var envelope struct {
Error struct {
Code string `json:"code"`
} `json:"error"`
}
doOIDCJITJSON(t, baseURL, http.MethodGet, "/api/v1/me", token, nil, expectedStatus, &envelope)
if envelope.Error.Code != expectedCode {
t.Fatalf("error code = %q, want %q", envelope.Error.Code, expectedCode)
}
}
func doOIDCJITJSON(t *testing.T, baseURL string, method string, path string, token string, payload any, expectedStatus int, output any) {
t.Helper()
var body io.Reader
if payload != nil {
raw, err := json.Marshal(payload)
if err != nil {
t.Fatalf("marshal OIDC JIT request: %v", err)
}
body = bytes.NewReader(raw)
}
request, err := http.NewRequest(method, baseURL+path, body)
if err != nil {
t.Fatalf("build %s %s request: %v", method, path, err)
}
request.Header.Set("Authorization", "Bearer "+token)
if payload != nil {
request.Header.Set("Content-Type", "application/json")
}
response, err := http.DefaultClient.Do(request)
if err != nil {
t.Fatalf("execute %s %s: %v", method, path, err)
}
defer response.Body.Close()
raw, err := io.ReadAll(io.LimitReader(response.Body, 1<<20))
if err != nil {
t.Fatalf("read %s %s response: %v", method, path, err)
}
if response.StatusCode != expectedStatus {
t.Fatalf("%s %s status=%d, want=%d", method, path, response.StatusCode, expectedStatus)
}
if output != nil && len(raw) > 0 {
if err := json.Unmarshal(raw, output); err != nil {
t.Fatalf("decode %s %s response: %v", method, path, err)
}
}
}
func signedOIDCJITToken(t *testing.T, key *ecdsa.PrivateKey, issuer string, subject string, mutate func(jwt.MapClaims)) string {
t.Helper()
now := time.Now()
claims := jwt.MapClaims{
"iss": issuer, "aud": "gateway-api", "sub": subject, "tid": "auth-center-test-tenant",
"preferred_username": "oidc-jit-acceptance", "roles": []string{"gateway.admin"},
"scope": "openid gateway.access", "iat": now.Unix(), "nbf": now.Add(-time.Second).Unix(),
"exp": now.Add(time.Hour).Unix(),
}
if mutate != nil {
mutate(claims)
}
token := jwt.NewWithClaims(jwt.SigningMethodES256, claims)
token.Header["kid"] = "jit-key"
raw, err := token.SignedString(key)
if err != nil {
t.Fatalf("sign OIDC JIT test token: %v", err)
}
return raw
}
func oidcJITECJWK(kid string, key *ecdsa.PublicKey) map[string]any {
return map[string]any{
"kid": kid,
"kty": "EC",
"use": "sig",
"alg": "ES256",
"crv": "P-256",
"x": base64.RawURLEncoding.EncodeToString(key.X.FillBytes(make([]byte, 32))),
"y": base64.RawURLEncoding.EncodeToString(key.Y.FillBytes(make([]byte, 32))),
}
}
+131
View File
@@ -0,0 +1,131 @@
package httpapi
import (
"net/http"
"strings"
"time"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
)
const (
errorCodeOIDCBrowserSessionDisabled = "OIDC_BROWSER_SESSION_DISABLED"
errorCodeOIDCSessionInvalid = "OIDC_SESSION_INVALID"
errorCodeOIDCSessionTooLarge = "OIDC_SESSION_TOKEN_TOO_LARGE"
errorCodeOIDCSessionCSRF = "OIDC_SESSION_CSRF_REJECTED"
maxOIDCSessionCookieTokenBytes = 3800
)
// createOIDCBrowserSession godoc
// @Summary 建立 OIDC 浏览器会话
// @Description 验证 Auth Center Access Token 后写入 HttpOnly 会话 Cookie;不会签发 Gateway JWT。
// @Tags auth
// @Success 204
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 404 {object} ErrorEnvelope
// @Router /api/v1/auth/oidc/session [post]
func (s *Server) createOIDCBrowserSession(w http.ResponseWriter, r *http.Request) {
if !s.cfg.OIDCEnabled || !s.cfg.OIDCBrowserSessionEnabled || s.auth == nil || s.auth.OIDCVerifier == nil {
writeError(w, http.StatusNotFound, "OIDC browser session is disabled", errorCodeOIDCBrowserSessionDisabled)
return
}
raw := bearerToken(r.Header.Get("Authorization"))
if raw == "" {
writeError(w, http.StatusUnauthorized, "valid OIDC access token is required", errorCodeOIDCSessionInvalid)
return
}
if len(raw) > maxOIDCSessionCookieTokenBytes {
writeError(w, http.StatusBadRequest, "OIDC access token is too large for browser session", errorCodeOIDCSessionTooLarge)
return
}
user, err := s.auth.AuthenticateOIDCAccessToken(r.Context(), raw)
if err != nil || user == nil || !strings.EqualFold(strings.TrimSpace(user.Source), "oidc") {
writeError(w, http.StatusUnauthorized, "valid OIDC access token is required", errorCodeOIDCSessionInvalid)
return
}
now := time.Now()
if user.TokenExpiresAt.IsZero() || !user.TokenExpiresAt.After(now) {
writeError(w, http.StatusUnauthorized, "OIDC access token has expired", errorCodeOIDCSessionInvalid)
return
}
maxAge := int(time.Until(user.TokenExpiresAt).Seconds())
if maxAge < 1 {
maxAge = 1
}
http.SetCookie(w, &http.Cookie{
Name: auth.OIDCSessionCookieName,
Value: raw,
Path: "/",
Expires: user.TokenExpiresAt,
MaxAge: maxAge,
HttpOnly: true,
Secure: s.cfg.OIDCSessionCookieSecure,
SameSite: http.SameSiteStrictMode,
})
w.Header().Set("Cache-Control", "no-store")
w.WriteHeader(http.StatusNoContent)
}
// deleteOIDCBrowserSession godoc
// @Summary 注销 OIDC 浏览器会话
// @Tags auth
// @Success 204
// @Router /api/v1/auth/oidc/session [delete]
func (s *Server) deleteOIDCBrowserSession(w http.ResponseWriter, _ *http.Request) {
http.SetCookie(w, &http.Cookie{
Name: auth.OIDCSessionCookieName,
Value: "",
Path: "/",
Expires: time.Unix(1, 0),
MaxAge: -1,
HttpOnly: true,
Secure: s.cfg.OIDCSessionCookieSecure,
SameSite: http.SameSiteStrictMode,
})
w.Header().Set("Cache-Control", "no-store")
w.WriteHeader(http.StatusNoContent)
}
func (s *Server) protectOIDCSessionCookie(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !s.cfg.OIDCEnabled || !s.cfg.OIDCBrowserSessionEnabled || isSafeHTTPMethod(r.Method) || hasExplicitCredential(r) {
next.ServeHTTP(w, r)
return
}
if _, err := r.Cookie(auth.OIDCSessionCookieName); err != nil {
next.ServeHTTP(w, r)
return
}
origin := strings.TrimSpace(r.Header.Get("Origin"))
if origin == "" || !originAllowed(origin, s.cfg.CORSAllowedOrigin) {
writeError(w, http.StatusForbidden, "browser session request origin was rejected", errorCodeOIDCSessionCSRF)
return
}
next.ServeHTTP(w, r)
})
}
func bearerToken(value string) string {
fields := strings.Fields(value)
if len(fields) == 2 && strings.EqualFold(fields[0], "bearer") {
return fields[1]
}
return ""
}
func isSafeHTTPMethod(method string) bool {
switch method {
case http.MethodGet, http.MethodHead, http.MethodOptions:
return true
default:
return false
}
}
func hasExplicitCredential(r *http.Request) bool {
return strings.TrimSpace(r.Header.Get("Authorization")) != "" ||
strings.TrimSpace(r.Header.Get("x-comfy-api-key")) != "" ||
strings.TrimSpace(r.Header.Get("x-goog-api-key")) != "" ||
strings.TrimSpace(r.URL.Query().Get("key")) != ""
}
@@ -0,0 +1,220 @@
package httpapi
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
)
func TestCreateOIDCBrowserSessionSetsProtectedSharedCookie(t *testing.T) {
server, token, closeIssuer := newOIDCSessionTestServer(t)
defer closeIssuer()
request := httptest.NewRequest(http.MethodPost, "/api/v1/auth/oidc/session", nil)
request.Header.Set("Authorization", "Bearer "+token)
recorder := httptest.NewRecorder()
server.createOIDCBrowserSession(recorder, request)
response := recorder.Result()
defer response.Body.Close()
if response.StatusCode != http.StatusNoContent {
t.Fatalf("session creation status = %d, want 204", response.StatusCode)
}
var sessionCookie *http.Cookie
for _, cookie := range response.Cookies() {
if cookie.Name == auth.OIDCSessionCookieName {
sessionCookie = cookie
break
}
}
if sessionCookie == nil {
t.Fatal("OIDC session cookie was not set")
}
if !sessionCookie.HttpOnly || sessionCookie.SameSite != http.SameSiteStrictMode || sessionCookie.Path != "/" {
t.Fatalf("unsafe OIDC session cookie attributes: %#v", sessionCookie)
}
if sessionCookie.MaxAge <= 0 || sessionCookie.Expires.IsZero() {
t.Fatalf("OIDC session cookie did not inherit token expiration: %#v", sessionCookie)
}
body, err := io.ReadAll(response.Body)
if err != nil {
t.Fatal(err)
}
if strings.Contains(string(body), token) {
t.Fatal("OIDC access token leaked into session response body")
}
}
func TestCreateOIDCBrowserSessionRejectsNonOIDCCredential(t *testing.T) {
server, _, closeIssuer := newOIDCSessionTestServer(t)
defer closeIssuer()
localToken, err := server.auth.SignJWT(&auth.User{ID: "local-user", Source: "gateway", Roles: []string{"user"}}, 0)
if err != nil {
t.Fatal(err)
}
request := httptest.NewRequest(http.MethodPost, "/api/v1/auth/oidc/session", nil)
request.Header.Set("Authorization", "Bearer "+localToken)
recorder := httptest.NewRecorder()
server.createOIDCBrowserSession(recorder, request)
if recorder.Code != http.StatusUnauthorized {
t.Fatalf("local credential session creation status = %d, want 401", recorder.Code)
}
}
func TestCreateOIDCBrowserSessionRejectsOversizedTokenBeforeCookieWrite(t *testing.T) {
server, _, closeIssuer := newOIDCSessionTestServer(t)
defer closeIssuer()
request := httptest.NewRequest(http.MethodPost, "/api/v1/auth/oidc/session", nil)
request.Header.Set("Authorization", "Bearer "+strings.Repeat("a", maxOIDCSessionCookieTokenBytes+1))
recorder := httptest.NewRecorder()
server.createOIDCBrowserSession(recorder, request)
if recorder.Code != http.StatusBadRequest || recorder.Header().Get("Set-Cookie") != "" {
t.Fatalf("oversized token response status=%d cookie=%q", recorder.Code, recorder.Header().Get("Set-Cookie"))
}
}
func TestCreateOIDCBrowserSessionHonorsDisabledFlag(t *testing.T) {
server, token, closeIssuer := newOIDCSessionTestServer(t)
defer closeIssuer()
server.cfg.OIDCBrowserSessionEnabled = false
request := httptest.NewRequest(http.MethodPost, "/api/v1/auth/oidc/session", nil)
request.Header.Set("Authorization", "Bearer "+token)
recorder := httptest.NewRecorder()
server.createOIDCBrowserSession(recorder, request)
if recorder.Code != http.StatusNotFound || recorder.Header().Get("Set-Cookie") != "" {
t.Fatalf("disabled session response status=%d cookie=%q", recorder.Code, recorder.Header().Get("Set-Cookie"))
}
}
func TestDeleteOIDCBrowserSessionExpiresCookie(t *testing.T) {
server, _, closeIssuer := newOIDCSessionTestServer(t)
defer closeIssuer()
request := httptest.NewRequest(http.MethodDelete, "/api/v1/auth/oidc/session", nil)
request.AddCookie(&http.Cookie{Name: auth.OIDCSessionCookieName, Value: "session-token"})
recorder := httptest.NewRecorder()
server.deleteOIDCBrowserSession(recorder, request)
response := recorder.Result()
defer response.Body.Close()
if response.StatusCode != http.StatusNoContent {
t.Fatalf("session deletion status = %d, want 204", response.StatusCode)
}
cookies := response.Cookies()
if len(cookies) != 1 || cookies[0].Name != auth.OIDCSessionCookieName || cookies[0].MaxAge >= 0 {
t.Fatalf("OIDC session cookie was not expired: %#v", cookies)
}
}
func TestOIDCSessionCSRFAcceptsOnlyAllowedOriginForCookieWrites(t *testing.T) {
server := &Server{cfg: config.Config{
OIDCEnabled: true,
OIDCBrowserSessionEnabled: true,
CORSAllowedOrigin: "https://gateway.example.com",
}}
next := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusNoContent) })
handler := server.protectOIDCSessionCookie(next)
for _, test := range []struct {
name string
method string
origin string
bearer bool
wantStatus int
}{
{name: "missing origin", method: http.MethodPost, wantStatus: http.StatusForbidden},
{name: "foreign origin", method: http.MethodDelete, origin: "https://evil.example", wantStatus: http.StatusForbidden},
{name: "allowed origin", method: http.MethodPatch, origin: "https://gateway.example.com", wantStatus: http.StatusNoContent},
{name: "safe request", method: http.MethodGet, wantStatus: http.StatusNoContent},
{name: "explicit bearer bypasses cookie csrf", method: http.MethodPost, bearer: true, wantStatus: http.StatusNoContent},
} {
t.Run(test.name, func(t *testing.T) {
request := httptest.NewRequest(test.method, "/api/workspace/tasks", nil)
request.AddCookie(&http.Cookie{Name: auth.OIDCSessionCookieName, Value: "session-token"})
if test.origin != "" {
request.Header.Set("Origin", test.origin)
}
if test.bearer {
request.Header.Set("Authorization", "Bearer explicit-token")
}
recorder := httptest.NewRecorder()
handler.ServeHTTP(recorder, request)
if recorder.Code != test.wantStatus {
t.Fatalf("status = %d, want %d", recorder.Code, test.wantStatus)
}
})
}
}
func TestOIDCSessionCSRFIsInactiveWhenOIDCIsDisabled(t *testing.T) {
server := &Server{cfg: config.Config{
OIDCEnabled: false,
OIDCBrowserSessionEnabled: true,
}}
request := httptest.NewRequest(http.MethodPost, "/api/v1/auth/login", nil)
request.AddCookie(&http.Cookie{Name: auth.OIDCSessionCookieName, Value: "irrelevant-cookie"})
recorder := httptest.NewRecorder()
server.protectOIDCSessionCookie(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusNoContent)
})).ServeHTTP(recorder, request)
if recorder.Code != http.StatusNoContent {
t.Fatalf("OIDC-disabled request status = %d, want 204", recorder.Code)
}
}
func newOIDCSessionTestServer(t *testing.T) (*Server, string, func()) {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatal(err)
}
var issuer string
issuerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/.well-known/openid-configuration":
_ = json.NewEncoder(w).Encode(map[string]any{"issuer": issuer, "jwks_uri": issuer + "/jwks"})
case "/jwks":
_ = json.NewEncoder(w).Encode(map[string]any{"keys": []any{oidcJITECJWK("jit-key", &key.PublicKey)}})
default:
http.NotFound(w, r)
}
}))
issuer = issuerServer.URL
verifier, err := auth.NewOIDCVerifier(auth.OIDCConfig{
Issuer: issuer, Audience: "gateway-api", TenantID: "auth-center-test-tenant",
RolePrefix: "gateway.", RequiredScopes: []string{"gateway.access"}, HTTPClient: issuerServer.Client(),
})
if err != nil {
issuerServer.Close()
t.Fatal(err)
}
authenticator := auth.New("test-local-jwt-secret", "", "")
authenticator.OIDCVerifier = verifier
server := &Server{
cfg: config.Config{
OIDCEnabled: true,
OIDCBrowserSessionEnabled: true,
OIDCSessionCookieSecure: false,
CORSAllowedOrigin: "http://localhost:5178",
},
auth: authenticator,
logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
}
return server, signedOIDCJITToken(t, key, issuer, "session-user", nil), issuerServer.Close
}
@@ -0,0 +1,90 @@
package httpapi
import (
"context"
"errors"
"net/http"
"strings"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
)
const (
errorCodeGatewayUserNotProvisioned = "GATEWAY_USER_NOT_PROVISIONED"
errorCodeGatewayUserDisabled = "GATEWAY_USER_DISABLED"
errorCodeGatewayTenantUnavailable = "GATEWAY_TENANT_UNAVAILABLE"
errorCodeGatewayProvisioningFailed = "GATEWAY_USER_PROVISIONING_FAILED"
)
type oidcUserResolver interface {
ResolveOrProvisionOIDCUser(context.Context, store.ResolveOrProvisionOIDCUserInput) (store.ResolveOrProvisionOIDCUserResult, error)
}
func (s *Server) requireUser(permission auth.Permission, next http.Handler) http.Handler {
return s.auth.Require(permission, s.resolveGatewayUser(next))
}
func (s *Server) resolveGatewayUser(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, ok := auth.UserFromContext(r.Context())
if !ok || user == nil || !strings.EqualFold(strings.TrimSpace(user.Source), "oidc") {
next.ServeHTTP(w, r)
return
}
if s.oidcUserResolver == nil {
s.writeOIDCUserResolutionError(w, r, errors.New("OIDC user resolver is unavailable"))
return
}
result, err := s.oidcUserResolver.ResolveOrProvisionOIDCUser(r.Context(), store.ResolveOrProvisionOIDCUserInput{
Issuer: s.cfg.OIDCIssuer,
Subject: user.ID,
Username: user.Username,
Roles: user.Roles,
TenantID: user.TenantID,
GatewayTenantKey: s.cfg.OIDCGatewayTenantKey,
ProvisioningEnabled: s.cfg.OIDCJITProvisioningEnabled,
RequestIP: limitAuditText(requestIP(r), 128),
UserAgent: limitAuditText(r.UserAgent(), 512),
})
if err != nil {
s.writeOIDCUserResolutionError(w, r, err)
return
}
if result.User == nil || strings.TrimSpace(result.User.GatewayUserID) == "" {
s.writeOIDCUserResolutionError(w, r, errors.New("OIDC user resolver returned no local user"))
return
}
if result.Created {
s.logger.InfoContext(r.Context(), "OIDC gateway user provisioned",
"gatewayUserId", result.User.GatewayUserID,
"auditId", result.AuditID,
)
}
next.ServeHTTP(w, r.WithContext(auth.WithUser(r.Context(), result.User)))
})
}
func (s *Server) writeOIDCUserResolutionError(w http.ResponseWriter, r *http.Request, err error) {
switch {
case errors.Is(err, store.ErrOIDCUserNotProvisioned):
writeError(w, http.StatusForbidden, "该账号尚未开通 EasyAI Gateway", errorCodeGatewayUserNotProvisioned)
case errors.Is(err, store.ErrOIDCUserDisabled):
writeError(w, http.StatusForbidden, "该 Gateway 账号已停用,请联系管理员", errorCodeGatewayUserDisabled)
case errors.Is(err, store.ErrOIDCTenantUnavailable):
writeError(w, http.StatusServiceUnavailable, "Gateway 租户尚未就绪,请联系管理员", errorCodeGatewayTenantUnavailable)
default:
s.logger.ErrorContext(r.Context(), "resolve OIDC gateway user failed", "error", err, "path", r.URL.Path)
writeError(w, http.StatusServiceUnavailable, "Gateway 账号初始化失败,请稍后重试", errorCodeGatewayProvisioningFailed)
}
}
func limitAuditText(value string, limit int) string {
value = strings.TrimSpace(value)
runes := []rune(value)
if limit > 0 && len(runes) > limit {
return string(runes[:limit])
}
return value
}
@@ -0,0 +1,153 @@
package httpapi
import (
"context"
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"testing"
"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/store"
)
type fakeOIDCUserResolver struct {
result store.ResolveOrProvisionOIDCUserResult
err error
calls int
input store.ResolveOrProvisionOIDCUserInput
}
func (f *fakeOIDCUserResolver) ResolveOrProvisionOIDCUser(_ context.Context, input store.ResolveOrProvisionOIDCUserInput) (store.ResolveOrProvisionOIDCUserResult, error) {
f.calls++
f.input = input
return f.result, f.err
}
func TestResolveGatewayUserAddsLocalOIDCContext(t *testing.T) {
resolver := &fakeOIDCUserResolver{result: store.ResolveOrProvisionOIDCUserResult{User: &auth.User{
ID: "platform-user",
Username: "alice",
Roles: []string{"basic"},
TenantID: "external-tenant",
Source: "oidc",
GatewayUserID: "21dd9ccb-3793-4023-ab31-4d04982ca4d3",
GatewayTenantID: "8f17f3ac-136e-4d0f-b097-655e2a6240a3",
TenantKey: "default",
UserGroupID: "6dcf86f2-8eaf-4b43-8e69-181315db24f0",
}}}
server := &Server{
cfg: config.Config{
OIDCIssuer: "https://auth.test.example/realms/easyai",
OIDCGatewayTenantKey: "default",
OIDCJITProvisioningEnabled: true,
},
oidcUserResolver: resolver,
logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
}
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, ok := auth.UserFromContext(r.Context())
if !ok || user.GatewayUserID == "" || user.GatewayTenantID == "" || user.UserGroupID == "" {
t.Fatalf("resolved Gateway context missing: %+v", user)
}
writeJSON(w, http.StatusOK, user)
})
request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil)
request = request.WithContext(auth.WithUser(request.Context(), &auth.User{
ID: "platform-user",
Username: "alice",
Roles: []string{"basic"},
TenantID: "external-tenant",
Source: "oidc",
}))
recorder := httptest.NewRecorder()
server.resolveGatewayUser(next).ServeHTTP(recorder, request)
if recorder.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", recorder.Code)
}
if resolver.calls != 1 || resolver.input.Subject != "platform-user" || resolver.input.GatewayTenantKey != "default" || !resolver.input.ProvisioningEnabled {
t.Fatalf("unexpected resolver call: calls=%d input=%+v", resolver.calls, resolver.input)
}
}
func TestResolveGatewayUserLeavesNonOIDCIdentityChainsUnchanged(t *testing.T) {
for _, source := range []string{"gateway", "api_key", "server-main"} {
t.Run(source, func(t *testing.T) {
resolver := &fakeOIDCUserResolver{}
server := &Server{oidcUserResolver: resolver, logger: slog.New(slog.NewTextHandler(io.Discard, nil))}
request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil)
original := &auth.User{ID: "local-user", Source: source, GatewayUserID: "local-user"}
request = request.WithContext(auth.WithUser(request.Context(), original))
recorder := httptest.NewRecorder()
server.resolveGatewayUser(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resolved, _ := auth.UserFromContext(r.Context())
if resolved != original {
t.Fatalf("non-OIDC identity context was replaced: %+v", resolved)
}
w.WriteHeader(http.StatusNoContent)
})).ServeHTTP(recorder, request)
if recorder.Code != http.StatusNoContent || resolver.calls != 0 {
t.Fatalf("status=%d resolver calls=%d", recorder.Code, resolver.calls)
}
})
}
}
func TestResolveGatewayUserReturnsStructuredErrors(t *testing.T) {
tests := []struct {
name string
err error
status int
code string
}{
{name: "not provisioned", err: store.ErrOIDCUserNotProvisioned, status: http.StatusForbidden, code: "GATEWAY_USER_NOT_PROVISIONED"},
{name: "disabled", err: store.ErrOIDCUserDisabled, status: http.StatusForbidden, code: "GATEWAY_USER_DISABLED"},
{name: "tenant unavailable", err: store.ErrOIDCTenantUnavailable, status: http.StatusServiceUnavailable, code: "GATEWAY_TENANT_UNAVAILABLE"},
{name: "storage failure", err: errors.New("database unavailable"), status: http.StatusServiceUnavailable, code: "GATEWAY_USER_PROVISIONING_FAILED"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
server := &Server{
cfg: config.Config{OIDCIssuer: "https://auth.test.example", OIDCGatewayTenantKey: "default"},
oidcUserResolver: &fakeOIDCUserResolver{err: test.err},
logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
}
request := httptest.NewRequest(http.MethodGet, "/api/v1/me", nil)
request = request.WithContext(auth.WithUser(request.Context(), &auth.User{
ID: "platform-user", Source: "oidc", TenantID: "external-tenant",
}))
recorder := httptest.NewRecorder()
server.resolveGatewayUser(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
t.Fatal("next handler must not run")
})).ServeHTTP(recorder, request)
if recorder.Code != test.status {
t.Fatalf("status = %d, want %d", recorder.Code, test.status)
}
var envelope struct {
Error struct {
Code string `json:"code"`
Message string `json:"message"`
Status int `json:"status"`
} `json:"error"`
}
if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil {
t.Fatalf("decode error envelope: %v", err)
}
if envelope.Error.Code != test.code || envelope.Error.Status != test.status || envelope.Error.Message == "" {
t.Fatalf("unexpected error envelope: %+v", envelope)
}
if envelope.Error.Message == test.err.Error() {
t.Fatalf("internal error leaked to response: %q", envelope.Error.Message)
}
})
}
}
+4
View File
@@ -33,6 +33,10 @@ func writeErrorWithDetails(w http.ResponseWriter, status int, message string, de
writeJSON(w, status, map[string]any{"error": errorPayload})
}
func writeLocalUserRequired(w http.ResponseWriter) {
writeError(w, http.StatusForbidden, "该账号尚未开通 EasyAI Gateway", errorCodeGatewayUserNotProvisioned)
}
func sendSSE(w http.ResponseWriter, event string, payload any) {
bytes, _ := json.Marshal(payload)
_, _ = fmt.Fprintf(w, "event: %s\n", event)
+87 -83
View File
@@ -18,6 +18,7 @@ type Server struct {
ctx context.Context
cfg config.Config
store *store.Store
oidcUserResolver oidcUserResolver
auth *auth.Authenticator
runner *runner.Service
logger *slog.Logger
@@ -30,12 +31,13 @@ func NewServer(cfg config.Config, db *store.Store, logger *slog.Logger) http.Han
func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Store, logger *slog.Logger) http.Handler {
server := &Server{
ctx: ctx,
cfg: cfg,
store: db,
auth: auth.New(cfg.JWTSecret, cfg.ServerMainBaseURL, cfg.ServerMainInternalToken),
runner: runner.New(cfg, db, logger),
logger: logger,
ctx: ctx,
cfg: cfg,
store: db,
oidcUserResolver: db,
auth: auth.New(cfg.JWTSecret, cfg.ServerMainBaseURL, cfg.ServerMainInternalToken),
runner: runner.New(cfg, db, logger),
logger: logger,
}
server.auth.LegacyJWTEnabled = !cfg.OIDCEnabled || cfg.OIDCAcceptLegacyHS256
server.auth.ServerMainInternalKey = cfg.ServerMainInternalKey
@@ -67,7 +69,9 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
mux.Handle("POST /api/v1/auth/register", server.auth.Require(auth.PermissionPublic, http.HandlerFunc(server.register)))
mux.Handle("POST /api/v1/auth/login", server.auth.Require(auth.PermissionPublic, http.HandlerFunc(server.login)))
mux.Handle("GET /api/v1/me", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.me)))
mux.HandleFunc("POST /api/v1/auth/oidc/session", server.createOIDCBrowserSession)
mux.HandleFunc("DELETE /api/v1/auth/oidc/session", server.deleteOIDCBrowserSession)
mux.Handle("GET /api/v1/me", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.me)))
mux.Handle("GET /api/v1/public/catalog/providers", server.auth.Require(auth.PermissionPublic, http.HandlerFunc(server.listCatalogProviders)))
mux.Handle("GET /api/v1/public/catalog/base-models", server.auth.Require(auth.PermissionPublic, http.HandlerFunc(server.listBaseModels)))
mux.Handle("GET /api/v1/public/client-customization", server.auth.Require(auth.PermissionPublic, http.HandlerFunc(server.getPublicClientCustomizationSettings)))
@@ -101,29 +105,29 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
mux.Handle("POST /api/admin/access-rules/batch", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.batchAccessRules)))
mux.Handle("PATCH /api/admin/access-rules/{ruleID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.updateAccessRule)))
mux.Handle("DELETE /api/admin/access-rules/{ruleID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.deleteAccessRule)))
mux.Handle("GET /api/v1/api-keys", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listAPIKeys)))
mux.Handle("POST /api/v1/api-keys", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.createAPIKey)))
mux.Handle("GET /api/v1/api-keys/access-rules", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listAPIKeyAccessRules)))
mux.Handle("POST /api/v1/api-keys/access-rules/batch", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.batchAPIKeyAccessRules)))
mux.Handle("PATCH /api/v1/api-keys/{apiKeyID}/scopes", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.updateAPIKeyScopes)))
mux.Handle("PATCH /api/v1/api-keys/{apiKeyID}/disable", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.disableAPIKey)))
mux.Handle("DELETE /api/v1/api-keys/{apiKeyID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.deleteAPIKey)))
mux.Handle("GET /api/playground/api-keys", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listPlayableAPIKeys)))
mux.Handle("GET /api/workspace/desktop-config", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.getDesktopConfig)))
mux.Handle("GET /api/workspace/user-groups", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listCurrentUserGroups)))
mux.Handle("GET /api/workspace/wallet", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.getWallet)))
mux.Handle("GET /api/workspace/wallet/transactions", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listWalletTransactions)))
mux.Handle("GET /api/workspace/tasks", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listTasks)))
mux.Handle("GET /api/workspace/tasks/{taskID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.getTask)))
mux.Handle("POST /api/workspace/tasks/{taskID}/cancel", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
mux.Handle("GET /api/workspace/tasks/{taskID}/param-preprocessing", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.taskParamPreprocessing)))
mux.Handle("GET /api/workspace/tasks/{taskID}/events", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.taskEvents)))
mux.Handle("GET /api/v1/api-keys", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listAPIKeys)))
mux.Handle("POST /api/v1/api-keys", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.createAPIKey)))
mux.Handle("GET /api/v1/api-keys/access-rules", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listAPIKeyAccessRules)))
mux.Handle("POST /api/v1/api-keys/access-rules/batch", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.batchAPIKeyAccessRules)))
mux.Handle("PATCH /api/v1/api-keys/{apiKeyID}/scopes", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.updateAPIKeyScopes)))
mux.Handle("PATCH /api/v1/api-keys/{apiKeyID}/disable", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.disableAPIKey)))
mux.Handle("DELETE /api/v1/api-keys/{apiKeyID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.deleteAPIKey)))
mux.Handle("GET /api/playground/api-keys", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listPlayableAPIKeys)))
mux.Handle("GET /api/workspace/desktop-config", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getDesktopConfig)))
mux.Handle("GET /api/workspace/user-groups", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listCurrentUserGroups)))
mux.Handle("GET /api/workspace/wallet", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getWallet)))
mux.Handle("GET /api/workspace/wallet/transactions", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listWalletTransactions)))
mux.Handle("GET /api/workspace/tasks", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listTasks)))
mux.Handle("GET /api/workspace/tasks/{taskID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getTask)))
mux.Handle("POST /api/workspace/tasks/{taskID}/cancel", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
mux.Handle("GET /api/workspace/tasks/{taskID}/param-preprocessing", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.taskParamPreprocessing)))
mux.Handle("GET /api/workspace/tasks/{taskID}/events", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.taskEvents)))
mux.Handle("GET /api/admin/pricing/rules", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listPricingRules)))
mux.Handle("GET /api/admin/pricing/rule-sets", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listPricingRuleSets)))
mux.Handle("POST /api/admin/pricing/rule-sets", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.createPricingRuleSet)))
mux.Handle("PATCH /api/admin/pricing/rule-sets/{ruleSetID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.updatePricingRuleSet)))
mux.Handle("DELETE /api/admin/pricing/rule-sets/{ruleSetID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.deletePricingRuleSet)))
mux.Handle("POST /api/v1/pricing/estimate", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.estimatePricing)))
mux.Handle("POST /api/v1/pricing/estimate", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.estimatePricing)))
mux.Handle("GET /api/admin/runtime/policy-sets", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listRuntimePolicySets)))
mux.Handle("POST /api/admin/runtime/policy-sets", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.createRuntimePolicySet)))
mux.Handle("PATCH /api/admin/runtime/policy-sets/{policySetID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.updateRuntimePolicySet)))
@@ -150,71 +154,71 @@ func NewServerWithContext(ctx context.Context, cfg config.Config, db *store.Stor
mux.Handle("POST /api/admin/platform-models", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.createPlatformModel)))
mux.Handle("DELETE /api/admin/platform-models/{modelID}", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.deletePlatformModel)))
mux.Handle("GET /api/admin/models", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listModels)))
mux.Handle("GET /api/v1/model-catalog", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listModelCatalog)))
mux.Handle("GET /api/v1/platforms", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listPlayablePlatforms)))
mux.Handle("GET /api/v1/models", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listPlayableModels)))
mux.Handle("GET /api/v1/playground/models", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listPlayableModels)))
mux.Handle("GET /api/v1/model-catalog", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listModelCatalog)))
mux.Handle("GET /api/v1/platforms", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listPlayablePlatforms)))
mux.Handle("GET /api/v1/models", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listPlayableModels)))
mux.Handle("GET /api/v1/playground/models", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listPlayableModels)))
mux.Handle("GET /api/admin/runtime/rate-limit-windows", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listRateLimitWindows)))
mux.Handle("GET /api/admin/runtime/model-rate-limits", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listModelRateLimitStatuses)))
mux.Handle("POST /api/v1/chat/completions", server.auth.Require(auth.PermissionBasic, server.createAPIV1ChatCompletions()))
mux.Handle("POST /api/v1/responses", server.auth.Require(auth.PermissionBasic, server.createTask("responses", false)))
mux.Handle("POST /api/v1/embeddings", server.auth.Require(auth.PermissionBasic, server.createTask("embeddings", false)))
mux.Handle("POST /api/v1/reranks", server.auth.Require(auth.PermissionBasic, server.createTask("reranks", false)))
mux.Handle("POST /api/v1/images/generations", server.auth.Require(auth.PermissionBasic, server.createTask("images.generations", false)))
mux.Handle("POST /api/v1/images/edits", server.auth.Require(auth.PermissionBasic, server.createTask("images.edits", false)))
mux.Handle("POST /api/v1/videos/generations", server.auth.Require(auth.PermissionBasic, server.createTask("videos.generations", false)))
mux.Handle("POST /api/v1/song/generations", server.auth.Require(auth.PermissionBasic, server.createTask("song.generations", true)))
mux.Handle("POST /api/v1/music/generations", server.auth.Require(auth.PermissionBasic, server.createTask("music.generations", true)))
mux.Handle("POST /api/v1/speech/generations", server.auth.Require(auth.PermissionBasic, server.createTask("speech.generations", true)))
mux.Handle("POST /api/v1/voice_clone", server.auth.Require(auth.PermissionBasic, server.createTask("voice.clone", true)))
mux.Handle("GET /api/v1/voice_clone/voices", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listClonedVoices)))
mux.Handle("DELETE /api/v1/voice_clone/voices/{voiceID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.deleteClonedVoice)))
mux.Handle("POST /api/v1/files/upload", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.uploadFile)))
mux.Handle("POST /api/v1/chat/completions", server.requireUser(auth.PermissionBasic, server.createAPIV1ChatCompletions()))
mux.Handle("POST /api/v1/responses", server.requireUser(auth.PermissionBasic, server.createTask("responses", false)))
mux.Handle("POST /api/v1/embeddings", server.requireUser(auth.PermissionBasic, server.createTask("embeddings", false)))
mux.Handle("POST /api/v1/reranks", server.requireUser(auth.PermissionBasic, server.createTask("reranks", false)))
mux.Handle("POST /api/v1/images/generations", server.requireUser(auth.PermissionBasic, server.createTask("images.generations", false)))
mux.Handle("POST /api/v1/images/edits", server.requireUser(auth.PermissionBasic, server.createTask("images.edits", false)))
mux.Handle("POST /api/v1/videos/generations", server.requireUser(auth.PermissionBasic, server.createTask("videos.generations", false)))
mux.Handle("POST /api/v1/song/generations", server.requireUser(auth.PermissionBasic, server.createTask("song.generations", true)))
mux.Handle("POST /api/v1/music/generations", server.requireUser(auth.PermissionBasic, server.createTask("music.generations", true)))
mux.Handle("POST /api/v1/speech/generations", server.requireUser(auth.PermissionBasic, server.createTask("speech.generations", true)))
mux.Handle("POST /api/v1/voice_clone", server.requireUser(auth.PermissionBasic, server.createTask("voice.clone", true)))
mux.Handle("GET /api/v1/voice_clone/voices", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listClonedVoices)))
mux.Handle("DELETE /api/v1/voice_clone/voices/{voiceID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.deleteClonedVoice)))
mux.Handle("POST /api/v1/files/upload", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.uploadFile)))
server.registerGeminiGenerateContentRoutes(mux)
mux.Handle("POST /upload/{version}/files", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.geminiFilesUpload)))
mux.Handle("POST /upload/{version}/files/{uploadID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.geminiFilesUploadFinalize)))
mux.Handle("GET /api/v1/tasks", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listTasks)))
mux.Handle("GET /api/v1/tasks/{taskID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.getTask)))
mux.Handle("POST /api/v1/tasks/{taskID}/cancel", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
mux.Handle("GET /api/v1/tasks/{taskID}/param-preprocessing", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.taskParamPreprocessing)))
mux.Handle("GET /api/v1/tasks/{taskID}/events", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.taskEvents)))
mux.Handle("GET /tasks", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listTasks)))
mux.Handle("GET /tasks/{taskID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.getTask)))
mux.Handle("POST /tasks/{taskID}/cancel", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
mux.Handle("GET /tasks/{taskID}/param-preprocessing", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.taskParamPreprocessing)))
mux.Handle("GET /tasks/{taskID}/events", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.taskEvents)))
mux.Handle("POST /chat/completions", server.auth.Require(auth.PermissionBasic, server.createTask("chat.completions", true)))
mux.Handle("POST /v1/chat/completions", server.auth.Require(auth.PermissionBasic, server.createTask("chat.completions", true)))
mux.Handle("POST /responses", server.auth.Require(auth.PermissionBasic, server.createTask("responses", true)))
mux.Handle("POST /v1/responses", server.auth.Require(auth.PermissionBasic, server.createTask("responses", true)))
mux.Handle("POST /embeddings", server.auth.Require(auth.PermissionBasic, server.createTask("embeddings", true)))
mux.Handle("POST /v1/embeddings", server.auth.Require(auth.PermissionBasic, server.createTask("embeddings", true)))
mux.Handle("POST /reranks", server.auth.Require(auth.PermissionBasic, server.createTask("reranks", true)))
mux.Handle("POST /v1/reranks", server.auth.Require(auth.PermissionBasic, server.createTask("reranks", true)))
mux.Handle("POST /images/generations", server.auth.Require(auth.PermissionBasic, server.createTask("images.generations", true)))
mux.Handle("POST /v1/images/generations", server.auth.Require(auth.PermissionBasic, server.createTask("images.generations", true)))
mux.Handle("POST /images/edits", server.auth.Require(auth.PermissionBasic, server.createTask("images.edits", true)))
mux.Handle("POST /v1/images/edits", server.auth.Require(auth.PermissionBasic, server.createTask("images.edits", true)))
mux.Handle("POST /song/generations", server.auth.Require(auth.PermissionBasic, server.createTask("song.generations", true)))
mux.Handle("POST /v1/song/generations", server.auth.Require(auth.PermissionBasic, server.createTask("song.generations", true)))
mux.Handle("POST /music/generations", server.auth.Require(auth.PermissionBasic, server.createTask("music.generations", true)))
mux.Handle("POST /v1/music/generations", server.auth.Require(auth.PermissionBasic, server.createTask("music.generations", true)))
mux.Handle("POST /speech/generations", server.auth.Require(auth.PermissionBasic, server.createTask("speech.generations", true)))
mux.Handle("POST /v1/speech/generations", server.auth.Require(auth.PermissionBasic, server.createTask("speech.generations", true)))
mux.Handle("POST /voice_clone", server.auth.Require(auth.PermissionBasic, server.createTask("voice.clone", true)))
mux.Handle("POST /v1/voice_clone", server.auth.Require(auth.PermissionBasic, server.createTask("voice.clone", true)))
mux.Handle("GET /voice_clone/voices", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listClonedVoices)))
mux.Handle("GET /v1/voice_clone/voices", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.listClonedVoices)))
mux.Handle("DELETE /voice_clone/voices/{voiceID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.deleteClonedVoice)))
mux.Handle("DELETE /v1/voice_clone/voices/{voiceID}", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.deleteClonedVoice)))
mux.Handle("POST /v1/files/upload", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.uploadFile)))
mux.Handle("POST /v1/tasks/{taskID}/cancel", server.auth.Require(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
mux.Handle("POST /upload/{version}/files", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.geminiFilesUpload)))
mux.Handle("POST /upload/{version}/files/{uploadID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.geminiFilesUploadFinalize)))
mux.Handle("GET /api/v1/tasks", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listTasks)))
mux.Handle("GET /api/v1/tasks/{taskID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getTask)))
mux.Handle("POST /api/v1/tasks/{taskID}/cancel", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
mux.Handle("GET /api/v1/tasks/{taskID}/param-preprocessing", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.taskParamPreprocessing)))
mux.Handle("GET /api/v1/tasks/{taskID}/events", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.taskEvents)))
mux.Handle("GET /tasks", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listTasks)))
mux.Handle("GET /tasks/{taskID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.getTask)))
mux.Handle("POST /tasks/{taskID}/cancel", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
mux.Handle("GET /tasks/{taskID}/param-preprocessing", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.taskParamPreprocessing)))
mux.Handle("GET /tasks/{taskID}/events", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.taskEvents)))
mux.Handle("POST /chat/completions", server.requireUser(auth.PermissionBasic, server.createTask("chat.completions", true)))
mux.Handle("POST /v1/chat/completions", server.requireUser(auth.PermissionBasic, server.createTask("chat.completions", true)))
mux.Handle("POST /responses", server.requireUser(auth.PermissionBasic, server.createTask("responses", true)))
mux.Handle("POST /v1/responses", server.requireUser(auth.PermissionBasic, server.createTask("responses", true)))
mux.Handle("POST /embeddings", server.requireUser(auth.PermissionBasic, server.createTask("embeddings", true)))
mux.Handle("POST /v1/embeddings", server.requireUser(auth.PermissionBasic, server.createTask("embeddings", true)))
mux.Handle("POST /reranks", server.requireUser(auth.PermissionBasic, server.createTask("reranks", true)))
mux.Handle("POST /v1/reranks", server.requireUser(auth.PermissionBasic, server.createTask("reranks", true)))
mux.Handle("POST /images/generations", server.requireUser(auth.PermissionBasic, server.createTask("images.generations", true)))
mux.Handle("POST /v1/images/generations", server.requireUser(auth.PermissionBasic, server.createTask("images.generations", true)))
mux.Handle("POST /images/edits", server.requireUser(auth.PermissionBasic, server.createTask("images.edits", true)))
mux.Handle("POST /v1/images/edits", server.requireUser(auth.PermissionBasic, server.createTask("images.edits", true)))
mux.Handle("POST /song/generations", server.requireUser(auth.PermissionBasic, server.createTask("song.generations", true)))
mux.Handle("POST /v1/song/generations", server.requireUser(auth.PermissionBasic, server.createTask("song.generations", true)))
mux.Handle("POST /music/generations", server.requireUser(auth.PermissionBasic, server.createTask("music.generations", true)))
mux.Handle("POST /v1/music/generations", server.requireUser(auth.PermissionBasic, server.createTask("music.generations", true)))
mux.Handle("POST /speech/generations", server.requireUser(auth.PermissionBasic, server.createTask("speech.generations", true)))
mux.Handle("POST /v1/speech/generations", server.requireUser(auth.PermissionBasic, server.createTask("speech.generations", true)))
mux.Handle("POST /voice_clone", server.requireUser(auth.PermissionBasic, server.createTask("voice.clone", true)))
mux.Handle("POST /v1/voice_clone", server.requireUser(auth.PermissionBasic, server.createTask("voice.clone", true)))
mux.Handle("GET /voice_clone/voices", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listClonedVoices)))
mux.Handle("GET /v1/voice_clone/voices", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.listClonedVoices)))
mux.Handle("DELETE /voice_clone/voices/{voiceID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.deleteClonedVoice)))
mux.Handle("DELETE /v1/voice_clone/voices/{voiceID}", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.deleteClonedVoice)))
mux.Handle("POST /v1/files/upload", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.uploadFile)))
mux.Handle("POST /v1/tasks/{taskID}/cancel", server.requireUser(auth.PermissionBasic, http.HandlerFunc(server.cancelTask)))
return server.recover(server.cors(mux))
return server.recover(server.cors(server.protectOIDCSessionCookie(mux)))
}
func (s *Server) requireAdmin(permission auth.Permission, next http.Handler) http.Handler {
return s.auth.Require(permission, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
return s.requireUser(permission, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, _ := auth.UserFromContext(r.Context())
if user != nil && strings.TrimSpace(user.APIKeyID) != "" {
writeError(w, http.StatusForbidden, "admin api does not accept api key credentials")
@@ -16,6 +16,8 @@ import (
// @Param currency query string false "币种" default(USD)
// @Success 200 {object} store.WalletSummary
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/workspace/wallet [get]
func (s *Server) getWallet(w http.ResponseWriter, r *http.Request) {
@@ -49,6 +51,8 @@ func (s *Server) getWallet(w http.ResponseWriter, r *http.Request) {
// @Success 200 {object} WalletTransactionListResponse
// @Failure 400 {object} ErrorEnvelope
// @Failure 401 {object} ErrorEnvelope
// @Failure 403 {object} ErrorEnvelope
// @Failure 503 {object} ErrorEnvelope
// @Failure 500 {object} ErrorEnvelope
// @Router /api/workspace/wallet/transactions [get]
func (s *Server) listWalletTransactions(w http.ResponseWriter, r *http.Request) {
+1 -1
View File
@@ -197,7 +197,7 @@ UPDATE gateway_users
SET deleted_at = now(),
status = 'deleted',
user_key = user_key || ':deleted:' || left(id::text, 8),
external_user_id = NULL,
external_user_id = CASE WHEN source = 'oidc' THEN external_user_id ELSE NULL END,
email = NULL,
updated_at = now()
WHERE id = $1::uuid AND deleted_at IS NULL`, id)
+347
View File
@@ -0,0 +1,347 @@
package store
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"strings"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/jackc/pgx/v5"
)
var (
ErrOIDCUserNotProvisioned = errors.New("OIDC gateway user is not provisioned")
ErrOIDCUserDisabled = errors.New("OIDC gateway user is disabled")
ErrOIDCTenantUnavailable = errors.New("OIDC gateway tenant is unavailable")
)
type ResolveOrProvisionOIDCUserInput struct {
Issuer string
Subject string
Username string
Roles []string
TenantID string
GatewayTenantKey string
ProvisioningEnabled bool
RequestIP string
UserAgent string
}
type ResolveOrProvisionOIDCUserResult struct {
User *auth.User
Created bool
AuditID string
}
type oidcUserProjection struct {
user GatewayUser
userGroupKey string
tenantStatus string
tenantDeleted bool
groupStatus string
userDeleted bool
}
func (s *Store) ResolveOrProvisionOIDCUser(ctx context.Context, input ResolveOrProvisionOIDCUserInput) (ResolveOrProvisionOIDCUserResult, error) {
input = normalizeOIDCUserInput(input)
if input.Issuer == "" || input.Subject == "" || input.TenantID == "" {
return ResolveOrProvisionOIDCUserResult{}, errors.New("invalid OIDC user projection input")
}
tx, err := s.pool.Begin(ctx)
if err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
defer func() { _ = tx.Rollback(ctx) }()
projection, err := loadOIDCUserProjection(ctx, tx, input.Subject)
if err == nil {
result, resolveErr := s.syncExistingOIDCUser(ctx, tx, projection, input)
if resolveErr != nil {
return ResolveOrProvisionOIDCUserResult{}, resolveErr
}
if err := tx.Commit(ctx); err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
return result, nil
}
if !errors.Is(err, pgx.ErrNoRows) {
return ResolveOrProvisionOIDCUserResult{}, err
}
if !input.ProvisioningEnabled {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCUserNotProvisioned
}
if input.GatewayTenantKey == "" {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCTenantUnavailable
}
tenantID, userGroupID, userGroupKey, err := loadOIDCProvisioningTenant(ctx, tx, input.GatewayTenantKey)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCTenantUnavailable
}
return ResolveOrProvisionOIDCUserResult{}, err
}
rolesJSON, err := json.Marshal(input.Roles)
if err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
userKey := deriveOIDCUserKey(input.Issuer, input.Subject)
username := input.Username
if username == "" {
username = "oidc-" + strings.TrimPrefix(userKey, "oidc:")[:12]
}
metadataJSON := `{"provisioningMode":"oidc-jit"}`
createdUser, 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', $2, $3, $4::uuid, $5, $6, $7::uuid, $8::jsonb, '{}'::jsonb, $9::jsonb,
'active', now(), now(), now())
ON CONFLICT DO NOTHING
RETURNING `+userColumns,
userKey,
input.Subject,
username,
tenantID,
input.TenantID,
input.GatewayTenantKey,
userGroupID,
string(rolesJSON),
metadataJSON,
))
if err != nil && !errors.Is(err, pgx.ErrNoRows) {
return ResolveOrProvisionOIDCUserResult{}, err
}
if errors.Is(err, pgx.ErrNoRows) {
projection, err = loadOIDCUserProjection(ctx, tx, input.Subject)
if err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
result, resolveErr := s.syncExistingOIDCUser(ctx, tx, projection, input)
if resolveErr != nil {
return ResolveOrProvisionOIDCUserResult{}, resolveErr
}
if err := tx.Commit(ctx); err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
return result, nil
}
if _, err := s.ensureWalletAccount(ctx, tx, createdUser.ID, "resource"); err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
subjectHash := sha256.Sum256([]byte(input.Subject))
audit, err := s.RecordAuditLogTx(ctx, tx, AuditLogInput{
Category: "identity",
Action: "identity.oidc_user.provisioned",
ActorGatewayUserID: createdUser.ID,
ActorUsername: createdUser.Username,
ActorSource: "oidc",
ActorRoles: createdUser.Roles,
TargetType: "gateway_user",
TargetID: createdUser.ID,
TargetGatewayUserID: createdUser.ID,
TargetGatewayTenantID: createdUser.GatewayTenantID,
RequestIP: input.RequestIP,
UserAgent: input.UserAgent,
AfterState: map[string]any{
"source": "oidc",
"tenantKey": createdUser.TenantKey,
"userGroupId": createdUser.DefaultUserGroupID,
},
Metadata: map[string]any{
"provisioningMode": "oidc-jit",
"externalSubjectHash": hex.EncodeToString(subjectHash[:])[:16],
},
})
if err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
if err := tx.Commit(ctx); err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
return ResolveOrProvisionOIDCUserResult{
User: authUserFromOIDCProjection(createdUser, userGroupKey),
Created: true,
AuditID: audit.ID,
}, nil
}
func (s *Store) syncExistingOIDCUser(ctx context.Context, tx pgx.Tx, projection oidcUserProjection, input ResolveOrProvisionOIDCUserInput) (ResolveOrProvisionOIDCUserResult, error) {
if projection.userDeleted || projection.user.Status != "active" {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCUserDisabled
}
if projection.tenantDeleted || projection.tenantStatus != "active" || projection.groupStatus != "active" ||
projection.user.GatewayTenantID == "" || projection.user.DefaultUserGroupID == "" || projection.userGroupKey == "" {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCTenantUnavailable
}
if input.GatewayTenantKey != "" && projection.user.TenantKey != input.GatewayTenantKey {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCTenantUnavailable
}
if projection.user.TenantID != "" && projection.user.TenantID != input.TenantID {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCTenantUnavailable
}
rolesJSON, err := json.Marshal(input.Roles)
if err != nil {
return ResolveOrProvisionOIDCUserResult{}, err
}
updated, 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'
AND deleted_at IS NULL
AND status = 'active'
RETURNING `+userColumns,
projection.user.ID,
input.Username,
string(rolesJSON),
))
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return ResolveOrProvisionOIDCUserResult{}, ErrOIDCUserDisabled
}
return ResolveOrProvisionOIDCUserResult{}, err
}
return ResolveOrProvisionOIDCUserResult{User: authUserFromOIDCProjection(updated, projection.userGroupKey)}, nil
}
func loadOIDCUserProjection(ctx context.Context, tx pgx.Tx, subject string) (oidcUserProjection, error) {
var projection oidcUserProjection
var roles []byte
var authProfile []byte
var metadata []byte
err := tx.QueryRow(ctx, `
SELECT
u.id::text, u.user_key, u.source, COALESCE(u.external_user_id, ''), u.username,
COALESCE(u.display_name, ''), COALESCE(u.email, ''), COALESCE(u.phone, ''), COALESCE(u.avatar_url, ''),
COALESCE(u.gateway_tenant_id::text, ''), COALESCE(u.tenant_id, ''), COALESCE(u.tenant_key, ''),
COALESCE(u.default_user_group_id::text, ''), u.roles, u.auth_profile, u.metadata,
u.status, COALESCE(u.last_login_at::text, ''), COALESCE(u.synced_at::text, ''), COALESCE(u.source_updated_at::text, ''),
u.created_at, u.updated_at,
COALESCE(g.group_key, ''), COALESCE(t.status, ''), t.deleted_at IS NOT NULL,
COALESCE(g.status, ''), u.deleted_at IS NOT NULL
FROM gateway_users u
LEFT JOIN gateway_tenants t ON t.id = u.gateway_tenant_id
LEFT JOIN gateway_user_groups g ON g.id = u.default_user_group_id
WHERE u.source = 'oidc' AND u.external_user_id = $1
FOR UPDATE OF u`, subject).Scan(
&projection.user.ID,
&projection.user.UserKey,
&projection.user.Source,
&projection.user.ExternalUserID,
&projection.user.Username,
&projection.user.DisplayName,
&projection.user.Email,
&projection.user.Phone,
&projection.user.AvatarURL,
&projection.user.GatewayTenantID,
&projection.user.TenantID,
&projection.user.TenantKey,
&projection.user.DefaultUserGroupID,
&roles,
&authProfile,
&metadata,
&projection.user.Status,
&projection.user.LastLoginAt,
&projection.user.SyncedAt,
&projection.user.SourceUpdatedAt,
&projection.user.CreatedAt,
&projection.user.UpdatedAt,
&projection.userGroupKey,
&projection.tenantStatus,
&projection.tenantDeleted,
&projection.groupStatus,
&projection.userDeleted,
)
if err != nil {
return oidcUserProjection{}, err
}
projection.user.Roles = decodeStringArray(roles)
projection.user.AuthProfile = decodeObject(authProfile)
projection.user.Metadata = decodeObject(metadata)
return projection, nil
}
func loadOIDCProvisioningTenant(ctx context.Context, tx pgx.Tx, tenantKey string) (string, string, string, error) {
var tenantID string
var groupID string
var groupKey string
err := tx.QueryRow(ctx, `
SELECT t.id::text, t.default_user_group_id::text, g.group_key
FROM gateway_tenants t
JOIN gateway_user_groups g ON g.id = t.default_user_group_id
WHERE t.tenant_key = $1
AND t.status = 'active'
AND t.deleted_at IS NULL
AND g.status = 'active'`, tenantKey).Scan(&tenantID, &groupID, &groupKey)
return tenantID, groupID, groupKey, err
}
func normalizeOIDCUserInput(input ResolveOrProvisionOIDCUserInput) ResolveOrProvisionOIDCUserInput {
input.Issuer = strings.TrimRight(strings.TrimSpace(input.Issuer), "/")
input.Subject = strings.TrimSpace(input.Subject)
input.Username = strings.TrimSpace(input.Username)
input.TenantID = strings.TrimSpace(input.TenantID)
input.GatewayTenantKey = strings.TrimSpace(input.GatewayTenantKey)
input.RequestIP = strings.TrimSpace(input.RequestIP)
input.UserAgent = strings.TrimSpace(input.UserAgent)
input.Roles = normalizeOIDCRoles(input.Roles)
return input
}
func normalizeOIDCRoles(roles []string) []string {
result := make([]string, 0, len(roles))
seen := make(map[string]struct{}, len(roles))
for _, role := range roles {
role = strings.TrimSpace(role)
if role == "" {
continue
}
if _, ok := seen[role]; ok {
continue
}
seen[role] = struct{}{}
result = append(result, role)
}
return result
}
func deriveOIDCUserKey(issuer string, subject string) string {
issuer = strings.TrimRight(strings.TrimSpace(issuer), "/")
sum := sha256.Sum256([]byte(issuer + "\x00" + strings.TrimSpace(subject)))
return fmt.Sprintf("oidc:%x", sum)
}
func authUserFromOIDCProjection(user GatewayUser, userGroupKey string) *auth.User {
groupKeys := []string(nil)
if userGroupKey != "" {
groupKeys = []string{userGroupKey}
}
return &auth.User{
ID: user.ExternalUserID,
Username: user.Username,
Roles: user.Roles,
TenantID: user.TenantID,
GatewayTenantID: user.GatewayTenantID,
TenantKey: user.TenantKey,
Source: "oidc",
GatewayUserID: user.ID,
UserGroupID: user.DefaultUserGroupID,
UserGroupKey: userGroupKey,
UserGroupKeys: groupKeys,
}
}
@@ -0,0 +1,326 @@
package store
import (
"context"
"os"
"path/filepath"
"runtime"
"sort"
"strings"
"sync"
"testing"
"time"
"github.com/jackc/pgx/v5/pgxpool"
)
func TestResolveOrProvisionOIDCUserLifecycleAndConcurrency(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run OIDC JIT PostgreSQL integration tests")
}
ctx := context.Background()
applyOIDCJITTestMigrations(t, ctx, databaseURL)
db, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect store: %v", err)
}
defer db.Close()
suffix := time.Now().UTC().Format("20060102150405.000000000")
subject := "platform-jit-" + suffix
input := ResolveOrProvisionOIDCUserInput{
Issuer: "https://auth.test.example/realms/easyai",
Subject: subject,
Username: "jit-user-" + suffix,
Roles: []string{"basic"},
TenantID: "auth-center-test-tenant",
GatewayTenantKey: "default",
ProvisioningEnabled: true,
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `
DELETE FROM gateway_audit_logs
WHERE target_gateway_user_id IN (
SELECT id FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1
)`, subject)
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1`, subject)
})
const callers = 12
results := make([]ResolveOrProvisionOIDCUserResult, callers)
errs := make([]error, callers)
var wg sync.WaitGroup
for index := 0; index < callers; index++ {
wg.Add(1)
go func(index int) {
defer wg.Done()
results[index], errs[index] = db.ResolveOrProvisionOIDCUser(ctx, input)
}(index)
}
wg.Wait()
firstID := ""
createdCount := 0
auditID := ""
for index, err := range errs {
if err != nil {
t.Fatalf("concurrent resolve %d: %v", index, err)
}
result := results[index]
if result.User == nil || result.User.GatewayUserID == "" {
t.Fatalf("concurrent resolve %d returned no local user: %+v", index, result)
}
if firstID == "" {
firstID = result.User.GatewayUserID
}
if result.User.GatewayUserID != firstID {
t.Fatalf("concurrent resolve returned different users: %q and %q", firstID, result.User.GatewayUserID)
}
if result.Created {
createdCount++
auditID = result.AuditID
}
}
if createdCount != 1 {
t.Fatalf("created count = %d, want 1", createdCount)
}
if auditID == "" {
t.Fatal("first provision must return an audit ID")
}
var users, wallets, audits int
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1`, subject).Scan(&users); err != nil {
t.Fatalf("count OIDC users: %v", err)
}
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_wallet_accounts WHERE gateway_user_id = $1::uuid AND currency = 'resource'`, firstID).Scan(&wallets); err != nil {
t.Fatalf("count wallets: %v", err)
}
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_audit_logs WHERE action = 'identity.oidc_user.provisioned' AND target_gateway_user_id = $1::uuid`, firstID).Scan(&audits); err != nil {
t.Fatalf("count audits: %v", err)
}
if users != 1 || wallets != 1 || audits != 1 {
t.Fatalf("users=%d wallets=%d audits=%d, want one of each", users, wallets, audits)
}
var auditProjection string
if err := db.pool.QueryRow(ctx, `
SELECT COALESCE(actor_user_id, '') || metadata::text || after_state::text
FROM gateway_audit_logs
WHERE id = $1::uuid`, auditID).Scan(&auditProjection); err != nil {
t.Fatalf("read OIDC provisioning audit: %v", err)
}
if strings.Contains(auditProjection, subject) || strings.Contains(auditProjection, input.Issuer) {
t.Fatal("OIDC provisioning audit exposed raw external identity claims")
}
if _, err := db.pool.Exec(ctx, `
UPDATE gateway_users
SET display_name = 'Manual Display Name',
email = 'manual-profile@example.test',
metadata = metadata || '{"manualProfile":true}'::jsonb
WHERE id = $1::uuid`, firstID); err != nil {
t.Fatalf("seed manually managed profile fields: %v", err)
}
input.Username = "jit-user-renamed-" + suffix
input.Roles = []string{"basic", "admin"}
input.ProvisioningEnabled = false
repeated, err := db.ResolveOrProvisionOIDCUser(ctx, input)
if err != nil {
t.Fatalf("repeat resolve: %v", err)
}
if repeated.Created || repeated.AuditID != "" || repeated.User.GatewayUserID != firstID {
t.Fatalf("unexpected repeat result: %+v", repeated)
}
if repeated.User.Username != input.Username || !containsOIDCTestRole(repeated.User.Roles, "admin") {
t.Fatalf("repeat resolve did not sync token projection: %+v", repeated.User)
}
var displayName, email string
var manualProfile bool
if err := db.pool.QueryRow(ctx, `
SELECT COALESCE(display_name, ''), COALESCE(email, ''), COALESCE((metadata->>'manualProfile')::boolean, false)
FROM gateway_users WHERE id = $1::uuid`, firstID).Scan(&displayName, &email, &manualProfile); err != nil {
t.Fatalf("read manually managed profile fields: %v", err)
}
if displayName != "Manual Display Name" || email != "manual-profile@example.test" || !manualProfile {
t.Fatalf("repeat resolve overwrote manually managed profile fields")
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_users SET status = 'disabled' WHERE id = $1::uuid`, firstID); err != nil {
t.Fatalf("disable OIDC user: %v", err)
}
if _, err := db.ResolveOrProvisionOIDCUser(ctx, input); !errorsIs(err, ErrOIDCUserDisabled) {
t.Fatalf("disabled resolve error = %v, want ErrOIDCUserDisabled", err)
}
var status string
if err := db.pool.QueryRow(ctx, `SELECT status FROM gateway_users WHERE id = $1::uuid`, firstID).Scan(&status); err != nil {
t.Fatalf("read disabled status: %v", err)
}
if status != "disabled" {
t.Fatalf("disabled OIDC user was reactivated: %q", status)
}
}
func TestResolveOrProvisionOIDCUserRejectsMissingMappingWithoutWrites(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run OIDC JIT PostgreSQL integration tests")
}
ctx := context.Background()
applyOIDCJITTestMigrations(t, ctx, databaseURL)
db, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect store: %v", err)
}
defer db.Close()
subject := "platform-jit-missing-" + time.Now().UTC().Format("20060102150405.000000000")
input := ResolveOrProvisionOIDCUserInput{
Issuer: "https://auth.test.example/realms/easyai",
Subject: subject,
Username: "missing-user",
Roles: []string{"basic"},
TenantID: "auth-center-test-tenant",
GatewayTenantKey: "missing-tenant-key",
}
if _, err := db.ResolveOrProvisionOIDCUser(ctx, input); !errorsIs(err, ErrOIDCUserNotProvisioned) {
t.Fatalf("disabled JIT error = %v, want ErrOIDCUserNotProvisioned", err)
}
input.ProvisioningEnabled = true
if _, err := db.ResolveOrProvisionOIDCUser(ctx, input); !errorsIs(err, ErrOIDCTenantUnavailable) {
t.Fatalf("missing tenant error = %v, want ErrOIDCTenantUnavailable", err)
}
var users int
if err := db.pool.QueryRow(ctx, `SELECT count(*) FROM gateway_users WHERE source = 'oidc' AND external_user_id = $1`, subject).Scan(&users); err != nil {
t.Fatalf("count rejected users: %v", err)
}
if users != 0 {
t.Fatalf("rejected OIDC request created %d users", users)
}
suffix := time.Now().UTC().Format("20060102150405.000000000")
groupKey := "jit-disabled-group-" + suffix
tenantKey := "jit-disabled-tenant-" + suffix
var groupID string
if err := db.pool.QueryRow(ctx, `
INSERT INTO gateway_user_groups (group_key, name, status)
VALUES ($1, 'OIDC JIT disabled group test', 'active')
RETURNING id::text`, groupKey).Scan(&groupID); err != nil {
t.Fatalf("create disabled-mapping test group: %v", err)
}
var tenantID string
if err := db.pool.QueryRow(ctx, `
INSERT INTO gateway_tenants (tenant_key, name, default_user_group_id, status)
VALUES ($1, 'OIDC JIT disabled tenant test', $2::uuid, 'disabled')
RETURNING id::text`, tenantKey, groupID).Scan(&tenantID); err != nil {
t.Fatalf("create disabled-mapping test tenant: %v", err)
}
t.Cleanup(func() {
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_tenants WHERE id = $1::uuid`, tenantID)
_, _ = db.pool.Exec(context.Background(), `DELETE FROM gateway_user_groups WHERE id = $1::uuid`, groupID)
})
disabledMappingInput := input
disabledMappingInput.Subject += "-disabled-mapping"
disabledMappingInput.GatewayTenantKey = tenantKey
if _, err := db.ResolveOrProvisionOIDCUser(ctx, disabledMappingInput); !errorsIs(err, ErrOIDCTenantUnavailable) {
t.Fatalf("disabled tenant error = %v, want ErrOIDCTenantUnavailable", err)
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_tenants SET status = 'active' WHERE id = $1::uuid`, tenantID); err != nil {
t.Fatalf("enable test tenant: %v", err)
}
if _, err := db.pool.Exec(ctx, `UPDATE gateway_user_groups SET status = 'disabled' WHERE id = $1::uuid`, groupID); err != nil {
t.Fatalf("disable test group: %v", err)
}
if _, err := db.ResolveOrProvisionOIDCUser(ctx, disabledMappingInput); !errorsIs(err, ErrOIDCTenantUnavailable) {
t.Fatalf("disabled user group error = %v, want ErrOIDCTenantUnavailable", err)
}
}
func TestOIDCUserKeyIsStableAndDoesNotExposeClaims(t *testing.T) {
issuer := "https://auth.test.example/realms/easyai"
subject := "platform-sensitive-subject"
first := deriveOIDCUserKey(issuer, subject)
second := deriveOIDCUserKey(issuer+"/", subject)
if first == "" || first != second {
t.Fatalf("OIDC user key is not stable: %q != %q", first, second)
}
if strings.Contains(first, subject) || strings.Contains(first, issuer) {
t.Fatalf("OIDC user key exposes raw claims: %q", first)
}
if first == deriveOIDCUserKey(issuer, subject+"-other") {
t.Fatal("different subjects produced the same OIDC user key")
}
}
func errorsIs(err error, target error) bool {
for err != nil {
if err == target {
return true
}
type unwrapper interface{ Unwrap() error }
wrapped, ok := err.(unwrapper)
if !ok {
return false
}
err = wrapped.Unwrap()
}
return false
}
func containsOIDCTestRole(roles []string, expected string) bool {
for _, role := range roles {
if role == expected {
return true
}
}
return false
}
func applyOIDCJITTestMigrations(t *testing.T, ctx context.Context, databaseURL string) {
t.Helper()
_, filename, _, _ := runtime.Caller(0)
migrationFiles, err := filepath.Glob(filepath.Join(filepath.Dir(filename), "..", "..", "migrations", "*.sql"))
if err != nil {
t.Fatalf("read migration files: %v", err)
}
sort.Strings(migrationFiles)
pool, err := pgxpool.New(ctx, databaseURL)
if err != nil {
t.Fatalf("connect migration db: %v", err)
}
defer pool.Close()
if _, err := pool.Exec(ctx, `CREATE TABLE IF NOT EXISTS schema_migrations (version text PRIMARY KEY, applied_at timestamptz NOT NULL DEFAULT now())`); err != nil {
t.Fatalf("ensure schema migrations: %v", err)
}
for _, migrationPath := range migrationFiles {
version := strings.TrimSuffix(filepath.Base(migrationPath), filepath.Ext(migrationPath))
var exists bool
if err := pool.QueryRow(ctx, `SELECT EXISTS (SELECT 1 FROM schema_migrations WHERE version = $1)`, version).Scan(&exists); err != nil {
t.Fatalf("check migration %s: %v", version, err)
}
if exists {
continue
}
migration, err := os.ReadFile(migrationPath)
if err != nil {
t.Fatalf("read migration %s: %v", version, err)
}
tx, err := pool.Begin(ctx)
if err != nil {
t.Fatalf("begin migration %s: %v", version, err)
}
if _, err := tx.Exec(ctx, string(migration)); err != nil {
_ = tx.Rollback(ctx)
t.Fatalf("apply migration %s: %v", version, err)
}
if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version) VALUES($1)`, version); err != nil {
_ = tx.Rollback(ctx)
t.Fatalf("record migration %s: %v", version, err)
}
if err := tx.Commit(ctx); err != nil {
t.Fatalf("commit migration %s: %v", version, err)
}
}
}
+80 -14
View File
@@ -50,6 +50,7 @@ import {
deleteApiKey,
deleteFileStorageChannel,
deleteGatewayUser,
deleteOIDCBrowserSession,
deletePlatform,
deleteTenant,
deleteUserGroup,
@@ -87,6 +88,7 @@ import {
listUserGroups,
listUsers,
loginLocalAccount,
OIDC_BROWSER_SESSION_CREDENTIAL,
pollTaskUntilSettled,
registerLocalAccount,
rechargeUserWalletBalance,
@@ -110,8 +112,20 @@ import { LoginRequiredPanel } from './components/LoginRequiredPanel';
import { useCatalogOperations } from './hooks/useCatalogOperations';
import { usePricingRuleSetOperations } from './hooks/usePricingRuleSetOperations';
import { useRuntimePolicySetOperations } from './hooks/useRuntimePolicySetOperations';
import { persistAccessToken, readStoredAccessToken } from './lib/auth-storage';
import { completeOIDCLogin, oidcLoginEnabled, startOIDCLogin, startOIDCLogout } from './lib/oidc';
import {
persistAccessToken,
persistLegacyOIDCAccessToken,
readLegacyOIDCAccessToken,
readStoredAccessToken,
} from './lib/auth-storage';
import { activateOIDCBrowserSession, restoreOIDCBrowserSession } from './lib/oidc-browser-session';
import {
completeOIDCLogin,
oidcBrowserSessionEnabled,
oidcLoginEnabled,
startOIDCLogin,
startOIDCLogout,
} from './lib/oidc';
import { runTask } from './lib/run-task';
import { AdminPage } from './pages/AdminPage';
import { ApiDocsPage } from './pages/ApiDocsPage';
@@ -272,17 +286,45 @@ export function App() {
currentTransactionQueryKeyRef.current = transactionListRequestKey;
useEffect(() => {
void completeOIDCLogin()
.then((result) => {
if (!result) return;
persistAccessToken(result.accessToken, 'session');
setToken(result.accessToken);
let cancelled = false;
void (async () => {
const result = await completeOIDCLogin();
if (cancelled) return;
if (result) {
if (!oidcBrowserSessionEnabled()) {
persistLegacyOIDCAccessToken(result.accessToken);
setToken(result.accessToken);
applyRoute(parseAppRoute(result.returnTo));
return;
}
const credential = await activateOIDCBrowserSession(result.accessToken);
if (cancelled) return;
setToken(credential);
applyRoute(parseAppRoute(result.returnTo));
})
.catch((err) => {
setState('error');
setError(err instanceof Error ? err.message : '统一认证登录失败');
});
return;
}
if (!oidcBrowserSessionEnabled()) return;
const legacyOIDCToken = readLegacyOIDCAccessToken();
if (legacyOIDCToken) {
const credential = await activateOIDCBrowserSession(legacyOIDCToken);
if (!cancelled) setToken(credential);
return;
}
if (readStoredAccessToken() || !oidcLoginEnabled()) return;
const restored = await restoreOIDCBrowserSession();
if (!restored || cancelled) return;
setCurrentUser(restored.user);
loadedDataKeysRef.current.add('currentUser');
setToken(restored.credential);
})().catch((err) => {
if (cancelled) return;
setState('error');
setError(err instanceof Error ? err.message : '统一认证登录失败');
});
return () => {
cancelled = true;
};
}, []);
useEffect(() => {
void ensureData(['health']);
@@ -402,6 +444,7 @@ export function App() {
try {
await Promise.all(requestKeys.map((key) => loadDataKey(key, nextToken)));
requestKeys.forEach((key) => loadedDataKeysRef.current.add(key));
setError('');
setState('ready');
} catch (err) {
if (handleAuthExpired(err, nextToken)) return;
@@ -1026,7 +1069,9 @@ export function App() {
const selectedApiKeySecret = selectedPlaygroundApiKeyId ? apiKeySecretsById[selectedPlaygroundApiKeyId] ?? '' : '';
const fallbackApiKeySecret = apiKeys.find((item) => Boolean(apiKeySecretsById[item.id]))?.id;
const credential = selectedApiKeySecret || (fallbackApiKeySecret ? apiKeySecretsById[fallbackApiKeySecret] : '') || apiKeySecret || token;
const credentialLabel = selectedApiKeySecret || fallbackApiKeySecret || apiKeySecret ? '本地 API Key' : '当前 Access Token';
const usingLocalAPIKey = Boolean(selectedApiKeySecret || fallbackApiKeySecret || apiKeySecret);
let credentialLabel = usingLocalAPIKey ? '本地 API Key' : '当前 Access Token';
if (!usingLocalAPIKey && token === OIDC_BROWSER_SESSION_CREDENTIAL) credentialLabel = '当前登录会话';
setCoreState('loading');
setCoreMessage('');
try {
@@ -1104,8 +1149,28 @@ export function App() {
}
async function signOut() {
const wasOIDCSession = token === OIDC_BROWSER_SESSION_CREDENTIAL;
if (wasOIDCSession) {
try {
await deleteOIDCBrowserSession();
} catch (err) {
setState('error');
setError(err instanceof Error ? err.message : '统一认证会话注销失败,请重试');
return;
}
}
const shouldEndOIDCSession = wasOIDCSession || currentUser?.source === 'oidc';
resetAuthenticatedSession();
if (await startOIDCLogout()) return;
if (shouldEndOIDCSession) {
try {
if (await startOIDCLogout()) return;
} catch {
navigatePath('/');
setState('error');
setError('Gateway 会话已注销,但统一认证退出失败');
return;
}
}
navigatePath('/');
}
@@ -1137,6 +1202,7 @@ export function App() {
}
function navigatePath(path: string) {
setError('');
if (`${window.location.pathname}${window.location.search}` !== path) {
window.history.pushState(null, '', path);
}
+79
View File
@@ -0,0 +1,79 @@
import { afterEach, describe, expect, it, vi } from 'vitest';
import {
createOIDCBrowserSession,
deleteOIDCBrowserSession,
GatewayApiError,
gatewayErrorMessage,
getCurrentUser,
OIDC_BROWSER_SESSION_CREDENTIAL,
} from './api';
describe('Gateway provisioning errors', () => {
const cases = [
['GATEWAY_USER_NOT_PROVISIONED', '该账号尚未开通 EasyAI Gateway'],
['GATEWAY_USER_DISABLED', '该 Gateway 账号已停用,请联系管理员'],
['GATEWAY_TENANT_UNAVAILABLE', 'Gateway 租户尚未就绪,请联系管理员'],
['GATEWAY_USER_PROVISIONING_FAILED', 'Gateway 账号初始化失败,请稍后重试'],
['OIDC_BROWSER_SESSION_DISABLED', 'Gateway 浏览器会话尚未启用'],
['OIDC_SESSION_INVALID', '统一认证会话无效,请重新登录'],
['OIDC_SESSION_TOKEN_TOO_LARGE', '统一认证凭证过大,无法建立浏览器会话'],
['OIDC_SESSION_CSRF_REJECTED', '登录会话来源校验失败,请刷新后重试'],
] as const;
for (const [code, expected] of cases) {
it(`maps ${code} to a stable Chinese status`, () => {
expect(gatewayErrorMessage({ code, message: 'internal server message', status: 503 })).toBe(expected);
expect(new GatewayApiError({ code, message: 'internal server message', status: 503 }).message).toContain(expected);
expect(new GatewayApiError({ code, message: 'internal server message', status: 503 }).message).not.toContain('internal server message');
});
}
it('keeps the server message for unrelated errors', () => {
expect(gatewayErrorMessage({ code: 'OTHER_ERROR', message: '模型不可用', status: 503 })).toBe('模型不可用');
});
});
describe('OIDC browser session transport', () => {
afterEach(() => {
vi.unstubAllGlobals();
});
it('exchanges the in-memory access token for a credentialed HttpOnly session', async () => {
const fetchMock = vi.fn().mockResolvedValue(new Response(null, { status: 204 }));
vi.stubGlobal('fetch', fetchMock);
await createOIDCBrowserSession('auth-center-access-token');
expect(fetchMock).toHaveBeenCalledOnce();
const [, init] = fetchMock.mock.calls[0] as [string, RequestInit];
expect(init.credentials).toBe('include');
expect(new Headers(init.headers).get('Authorization')).toBe('Bearer auth-center-access-token');
expect(init.body).toBeUndefined();
});
it('uses the shared cookie without sending a synthetic bearer token', async () => {
const fetchMock = vi.fn().mockResolvedValue(new Response(JSON.stringify({ sub: 'oidc-user' }), {
status: 200,
headers: { 'Content-Type': 'application/json' },
}));
vi.stubGlobal('fetch', fetchMock);
await getCurrentUser(OIDC_BROWSER_SESSION_CREDENTIAL);
const [, init] = fetchMock.mock.calls[0] as [string, RequestInit];
expect(init.credentials).toBe('include');
expect(new Headers(init.headers).has('Authorization')).toBe(false);
});
it('deletes the shared cookie with a credentialed request', async () => {
const fetchMock = vi.fn().mockResolvedValue(new Response(null, { status: 204 }));
vi.stubGlobal('fetch', fetchMock);
await deleteOIDCBrowserSession();
const [, init] = fetchMock.mock.calls[0] as [string, RequestInit];
expect(init.method).toBe('DELETE');
expect(init.credentials).toBe('include');
expect(new Headers(init.headers).has('Authorization')).toBe(false);
});
});
+44 -7
View File
@@ -51,10 +51,12 @@ import type {
WalletSummaryResponse,
} from '@easyai-ai-gateway/contracts';
import type { PlatformCreateInput, PlatformModelBindingInput, WorkspaceTaskQuery } from './types';
import { oidcBrowserSessionEnabled } from './lib/oidc';
const API_BASE = import.meta.env.VITE_GATEWAY_API_BASE_URL ?? 'http://localhost:8088';
export const OIDC_BROWSER_SESSION_CREDENTIAL = '__easyai_gateway_oidc_browser_session__';
interface GatewayErrorDetails {
export interface GatewayErrorDetails {
code?: string;
message: string;
requestId?: string;
@@ -110,6 +112,20 @@ export async function loginLocalAccount(input: { account: string; password: stri
});
}
export async function createOIDCBrowserSession(accessToken: string): Promise<void> {
await request<void>('/api/v1/auth/oidc/session', {
method: 'POST',
token: accessToken,
});
}
export async function deleteOIDCBrowserSession(): Promise<void> {
await request<void>('/api/v1/auth/oidc/session', {
auth: false,
method: 'DELETE',
});
}
export async function getCurrentUser(token: string): Promise<AuthUser> {
return request<AuthUser>('/api/v1/me', { token });
}
@@ -631,9 +647,10 @@ export async function* streamChatCompletionText(
const response = await fetch(`${API_BASE}/v1/chat/completions`, {
body: JSON.stringify({ ...input, stream: true }),
headers: {
Authorization: `Bearer ${token}`,
...authorizationHeader(token),
'Content-Type': 'application/json',
},
credentials: oidcBrowserSessionEnabled() ? 'include' : 'same-origin',
method: 'POST',
signal,
});
@@ -834,9 +851,8 @@ export async function uploadFileToStorage(
const response = await fetch(`${API_BASE}/v1/files/upload`, {
body: form,
headers: {
Authorization: `Bearer ${token}`,
},
credentials: oidcBrowserSessionEnabled() ? 'include' : 'same-origin',
headers: authorizationHeader(token),
method: 'POST',
});
const body = await response.text();
@@ -1030,7 +1046,7 @@ async function request<T>(
options: { token?: string; auth?: boolean; method?: string; body?: unknown; headers?: Record<string, string> } = {},
): Promise<T> {
const headers: Record<string, string> = { ...(options.headers ?? {}) };
if (options.auth !== false && options.token) {
if (options.auth !== false && options.token && options.token !== OIDC_BROWSER_SESSION_CREDENTIAL) {
headers.Authorization = `Bearer ${options.token}`;
}
if (options.body !== undefined) {
@@ -1040,6 +1056,7 @@ async function request<T>(
method: options.method ?? 'GET',
headers,
body: options.body === undefined ? undefined : JSON.stringify(options.body),
credentials: oidcBrowserSessionEnabled() ? 'include' : 'same-origin',
});
if (!response.ok) {
const body = await response.text();
@@ -1051,6 +1068,11 @@ async function request<T>(
return response.json() as Promise<T>;
}
function authorizationHeader(token: string): Record<string, string> {
if (!token || token === OIDC_BROWSER_SESSION_CREDENTIAL) return {};
return { Authorization: `Bearer ${token}` };
}
function delay(ms: number) {
return new Promise((resolve) => window.setTimeout(resolve, ms));
}
@@ -1122,7 +1144,7 @@ function errorDetailsFromParsed(parsed: unknown, status?: number, fallback = '')
}
function formatGatewayErrorDetails(details: GatewayErrorDetails) {
const message = details.message || '请求失败';
const message = gatewayErrorMessage(details);
const meta = [
details.code ? `错误码: ${details.code}` : '',
details.status ? `状态: ${details.status}` : '',
@@ -1132,6 +1154,21 @@ function formatGatewayErrorDetails(details: GatewayErrorDetails) {
return meta.length ? `${message}${meta.join('')}` : message;
}
const gatewayProvisioningErrorMessages: Record<string, string> = {
GATEWAY_USER_NOT_PROVISIONED: '该账号尚未开通 EasyAI Gateway',
GATEWAY_USER_DISABLED: '该 Gateway 账号已停用,请联系管理员',
GATEWAY_TENANT_UNAVAILABLE: 'Gateway 租户尚未就绪,请联系管理员',
GATEWAY_USER_PROVISIONING_FAILED: 'Gateway 账号初始化失败,请稍后重试',
OIDC_BROWSER_SESSION_DISABLED: 'Gateway 浏览器会话尚未启用',
OIDC_SESSION_INVALID: '统一认证会话无效,请重新登录',
OIDC_SESSION_TOKEN_TOO_LARGE: '统一认证凭证过大,无法建立浏览器会话',
OIDC_SESSION_CSRF_REJECTED: '登录会话来源校验失败,请刷新后重试',
};
export function gatewayErrorMessage(details: GatewayErrorDetails) {
return (details.code && gatewayProvisioningErrorMessages[details.code]) || details.message || '请求失败';
}
function recordFromUnknown(value: unknown): Record<string, unknown> | undefined {
if (!value || typeof value !== 'object' || Array.isArray(value)) return undefined;
return value as Record<string, unknown>;
+12 -7
View File
@@ -13,16 +13,21 @@ describe('auth storage', () => {
vi.stubGlobal('window', { localStorage: new MemoryStorage(), sessionStorage: new MemoryStorage() });
});
it('keeps OIDC access tokens in session storage and removes persistent tokens', () => {
persistAccessToken('legacy-token');
persistAccessToken('oidc-token', 'session');
expect(readStoredAccessToken()).toBe('oidc-token');
expect(window.localStorage.getItem('easyai_ai_gateway_access_token')).toBeNull();
expect(window.sessionStorage.getItem('easyai_ai_gateway_oidc_access_token')).toBe('oidc-token');
it('persists only explicit local or externally supplied bearer tokens', () => {
persistAccessToken('local-token');
expect(readStoredAccessToken()).toBe('local-token');
expect(window.localStorage.getItem('easyai_ai_gateway_access_token')).toBe('local-token');
expect(window.sessionStorage.getItem('easyai_ai_gateway_oidc_access_token')).toBeNull();
});
it('reads a legacy OIDC session token only for one-time cookie migration', () => {
window.sessionStorage.setItem('easyai_ai_gateway_oidc_access_token', 'legacy-oidc-token');
expect(readStoredAccessToken()).toBe('legacy-oidc-token');
});
it('clears both token stores on logout', () => {
persistAccessToken('oidc-token', 'session');
window.sessionStorage.setItem('easyai_ai_gateway_oidc_access_token', 'legacy-oidc-token');
persistAccessToken('local-token');
persistAccessToken('');
expect(readStoredAccessToken()).toBe('');
});
+21 -4
View File
@@ -12,15 +12,32 @@ export function readStoredAccessToken() {
}
}
export function persistAccessToken(value: string, storage: 'local' | 'session' = 'local') {
export function readLegacyOIDCAccessToken() {
if (typeof window === 'undefined') return '';
try {
return window.sessionStorage.getItem(OIDC_SESSION_TOKEN_STORAGE_KEY) ?? '';
} catch {
return '';
}
}
export function persistLegacyOIDCAccessToken(value: string) {
if (typeof window === 'undefined') return;
try {
window.localStorage.removeItem(AUTH_TOKEN_STORAGE_KEY);
if (value) window.sessionStorage.setItem(OIDC_SESSION_TOKEN_STORAGE_KEY, value);
else window.sessionStorage.removeItem(OIDC_SESSION_TOKEN_STORAGE_KEY);
} catch {
// Compatibility-only rollback path for deployments with browser sessions disabled.
}
}
export function persistAccessToken(value: string) {
if (typeof window === 'undefined') return;
try {
if (!value) {
window.localStorage.removeItem(AUTH_TOKEN_STORAGE_KEY);
window.sessionStorage.removeItem(OIDC_SESSION_TOKEN_STORAGE_KEY);
} else if (storage === 'session') {
window.localStorage.removeItem(AUTH_TOKEN_STORAGE_KEY);
window.sessionStorage.setItem(OIDC_SESSION_TOKEN_STORAGE_KEY, value);
} else {
window.sessionStorage.removeItem(OIDC_SESSION_TOKEN_STORAGE_KEY);
window.localStorage.setItem(AUTH_TOKEN_STORAGE_KEY, value);
@@ -0,0 +1,47 @@
import { beforeEach, describe, expect, it, vi } from 'vitest';
import { OIDC_BROWSER_SESSION_CREDENTIAL } from '../api';
import { activateOIDCBrowserSession, restoreOIDCBrowserSession } from './oidc-browser-session';
class MemoryStorage {
private values = new Map<string, string>();
getItem(key: string) { return this.values.get(key) ?? null; }
setItem(key: string, value: string) { this.values.set(key, value); }
removeItem(key: string) { this.values.delete(key); }
}
describe('OIDC browser session lifecycle', () => {
beforeEach(() => {
vi.stubGlobal('window', { localStorage: new MemoryStorage(), sessionStorage: new MemoryStorage() });
});
it('moves an OIDC access token into the HttpOnly session and clears browser token storage', async () => {
window.localStorage.setItem('easyai_ai_gateway_access_token', 'old-local-token');
window.sessionStorage.setItem('easyai_ai_gateway_oidc_access_token', 'old-oidc-token');
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(new Response(null, { status: 204 })));
const credential = await activateOIDCBrowserSession('new-oidc-token');
expect(credential).toBe(OIDC_BROWSER_SESSION_CREDENTIAL);
expect(window.localStorage.getItem('easyai_ai_gateway_access_token')).toBeNull();
expect(window.sessionStorage.getItem('easyai_ai_gateway_oidc_access_token')).toBeNull();
});
it('restores a shared cookie session in a fresh tab', async () => {
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(new Response(JSON.stringify({
sub: 'platform-user', source: 'oidc', gatewayUserId: 'gateway-user',
}), { status: 200, headers: { 'Content-Type': 'application/json' } })));
const restored = await restoreOIDCBrowserSession();
expect(restored?.credential).toBe(OIDC_BROWSER_SESSION_CREDENTIAL);
expect(restored?.user.sub).toBe('platform-user');
});
it('treats a missing or expired shared cookie as a signed-out tab', async () => {
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(new Response(JSON.stringify({
error: { message: 'unauthorized', status: 401 },
}), { status: 401, headers: { 'Content-Type': 'application/json' } })));
await expect(restoreOIDCBrowserSession()).resolves.toBeNull();
});
});
+31
View File
@@ -0,0 +1,31 @@
import {
createOIDCBrowserSession,
GatewayApiError,
getCurrentUser,
OIDC_BROWSER_SESSION_CREDENTIAL,
} from '../api';
import { persistAccessToken } from './auth-storage';
type CurrentUser = Awaited<ReturnType<typeof getCurrentUser>>;
export interface RestoredOIDCBrowserSession {
credential: typeof OIDC_BROWSER_SESSION_CREDENTIAL;
user: CurrentUser;
}
export async function activateOIDCBrowserSession(accessToken: string) {
if (!accessToken.trim()) throw new Error('统一认证未返回 Access Token');
await createOIDCBrowserSession(accessToken);
persistAccessToken('');
return OIDC_BROWSER_SESSION_CREDENTIAL;
}
export async function restoreOIDCBrowserSession(): Promise<RestoredOIDCBrowserSession | null> {
try {
const user = await getCurrentUser(OIDC_BROWSER_SESSION_CREDENTIAL);
return { credential: OIDC_BROWSER_SESSION_CREDENTIAL, user };
} catch (error) {
if (error instanceof GatewayApiError && error.details.status === 401) return null;
throw error;
}
}
+7 -3
View File
@@ -1,4 +1,5 @@
const enabled = import.meta.env.VITE_OIDC_ENABLED === 'true';
const browserSessionEnabled = import.meta.env.VITE_OIDC_BROWSER_SESSION_ENABLED !== 'false';
const issuer = (import.meta.env.VITE_OIDC_ISSUER ?? '').replace(/\/$/, '');
const clientId = import.meta.env.VITE_OIDC_CLIENT_ID ?? '';
const configuredRedirect = import.meta.env.VITE_OIDC_REDIRECT_URI ?? '';
@@ -25,6 +26,10 @@ export function oidcLoginEnabled() {
return enabled && Boolean(issuer && clientId);
}
export function oidcBrowserSessionEnabled() {
return browserSessionEnabled;
}
export async function startOIDCLogin() {
assertConfigured();
const discovery = await getDiscovery();
@@ -81,7 +86,6 @@ async function completeOIDCLoginOnce(): Promise<{ accessToken: string; returnTo:
if (!payload.access_token) throw new Error('统一认证未返回 Access Token');
if (!payload.id_token) throw new Error('统一认证未返回 ID Token');
validateIDToken(payload.id_token, transaction.nonce);
window.sessionStorage.setItem(idTokenKey, payload.id_token);
window.history.replaceState({}, '', transaction.returnTo || '/');
return { accessToken: payload.access_token, returnTo: transaction.returnTo || '/' };
}
@@ -91,10 +95,10 @@ export async function startOIDCLogout() {
const idToken = window.sessionStorage.getItem(idTokenKey);
window.sessionStorage.removeItem(idTokenKey);
window.sessionStorage.removeItem(transactionKey);
if (!idToken) return false;
const discovery = await getDiscovery();
if (!discovery.end_session_endpoint) return false;
const params = new URLSearchParams({ id_token_hint: idToken, post_logout_redirect_uri: window.location.origin + '/' });
const params = new URLSearchParams({ client_id: clientId, post_logout_redirect_uri: window.location.origin + '/' });
if (idToken) params.set('id_token_hint', idToken);
window.location.assign(`${discovery.end_session_endpoint}?${params}`);
return true;
}