fix: 修复 OIDC 用户预配与跨标签页登录态
增加受控 JIT 本地用户投影,并使用 HttpOnly Cookie 在新标签页恢复认证中心登录态。补充错误语义、安全校验、自动化测试与回滚配置。
This commit is contained in:
@@ -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
@@ -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
@@ -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: 列出任务
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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))),
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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
@@ -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>;
|
||||
|
||||
@@ -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('');
|
||||
});
|
||||
|
||||
@@ -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();
|
||||
});
|
||||
});
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user