diff --git a/.env.example b/.env.example index 156951a..f124b14 100644 --- a/.env.example +++ b/.env.example @@ -35,8 +35,16 @@ OIDC_ACCEPT_LEGACY_HS256=true OIDC_INTROSPECTION_ENABLED=false OIDC_INTROSPECTION_CLIENT_ID= OIDC_INTROSPECTION_CLIENT_SECRET= +# Controlled JIT is opt-in. When enabled, bind the validated Auth Center tid to +# one existing active Gateway tenant; tokens never create Gateway tenants. +OIDC_JIT_PROVISIONING_ENABLED=false +OIDC_GATEWAY_TENANT_KEY= +OIDC_BROWSER_SESSION_ENABLED=true +# Staging/production must use true. Local HTTP development may use false. +OIDC_SESSION_COOKIE_SECURE=false OIDC_CLIENT_ID= OIDC_REDIRECT_URI=http://localhost:5178/auth/callback +VITE_OIDC_BROWSER_SESSION_ENABLED=true AI_GATEWAY_WEB_BASE_PATH=/ AI_GATEWAY_GO_BUILD_IMAGE=golang:1.26.3-alpine AI_GATEWAY_API_RUNTIME_IMAGE=alpine:3.22 diff --git a/Dockerfile b/Dockerfile index 1b59cc5..303e9fd 100644 --- a/Dockerfile +++ b/Dockerfile @@ -75,12 +75,14 @@ ARG VITE_OIDC_ENABLED=false ARG VITE_OIDC_ISSUER= ARG VITE_OIDC_CLIENT_ID= ARG VITE_OIDC_REDIRECT_URI= +ARG VITE_OIDC_BROWSER_SESSION_ENABLED=true ARG VITE_BASE_PATH=/ ENV VITE_GATEWAY_API_BASE_URL=$VITE_GATEWAY_API_BASE_URL ENV VITE_OIDC_ENABLED=$VITE_OIDC_ENABLED \ VITE_OIDC_ISSUER=$VITE_OIDC_ISSUER \ VITE_OIDC_CLIENT_ID=$VITE_OIDC_CLIENT_ID \ VITE_OIDC_REDIRECT_URI=$VITE_OIDC_REDIRECT_URI \ + VITE_OIDC_BROWSER_SESSION_ENABLED=$VITE_OIDC_BROWSER_SESSION_ENABLED \ VITE_BASE_PATH=$VITE_BASE_PATH RUN pnpm --filter @easyai-ai-gateway/web build diff --git a/README.md b/README.md index 4533ffb..bd402a3 100644 --- a/README.md +++ b/README.md @@ -36,6 +36,22 @@ pnpm dev - PostgreSQL: 目标版本 18,默认使用宿主机 `localhost:5432` 上的 `easyai-pgvector` 实例,并使用独立库 `easyai_ai_gateway` - 身份模式: 默认 `IDENTITY_MODE=hybrid`,可同时测试 Gateway 本地账号注册登录、可选邀请码和 `server-main` JWT / API Key 对接。 +### Auth Center OIDC 受控 JIT + +OIDC 用户通过签名、Issuer、Audience、`tid`、Scope 和应用角色校验后,Gateway 可在首个受保护请求前创建本地业务投影。该投影只用于 Gateway 的用户组、钱包、API Key、任务归属和审计,不修改 Auth Center Claims,也不迁移或合并历史账号。 + +```dotenv +OIDC_ENABLED=true +OIDC_JIT_PROVISIONING_ENABLED=true +OIDC_GATEWAY_TENANT_KEY=default +OIDC_BROWSER_SESSION_ENABLED=true +OIDC_SESSION_COOKIE_SECURE=false # 本地 HTTP;生产必须为 true +``` + +`OIDC_GATEWAY_TENANT_KEY` 必须指向已有且启用的 Gateway 租户及其默认用户组;JIT 不会从 Token 自动创建租户。开关默认关闭,关闭后不再创建新用户,但仍可解析此前创建的 `source=oidc` 映射。 + +Web Console 默认把已验证的 Auth Center Access Token 转入 `HttpOnly + SameSite=Strict` Cookie,随后清除浏览器 `localStorage/sessionStorage` 中的 OIDC Access Token。Cookie 在同域标签页之间共享且有效期不超过原 Access Token;Gateway 不会二次签发 JWT。Staging、生产及其他非本地环境必须启用 Secure Cookie,并配置明确的 `CORS_ALLOWED_ORIGIN`,禁止 `*`。完整行为、错误码和回滚方式见 [OIDC JIT 接入说明](docs/oidc-jit-provisioning.md)。 + `pnpm dev` 会先创建数据库并执行 migration,然后并行启动: - `api:dev`:通过 `scripts/go-watch.mjs` 运行 Go API,监听 `.go`、`go.mod`、`go.sum` 变化并自动重启后端进程;watcher 会按进程组终止旧的 `go run` 和其子进程,避免热更新时残留进程占用 API 端口。 diff --git a/apps/api/cmd/gateway/main.go b/apps/api/cmd/gateway/main.go index ca2c6b2..73644f3 100644 --- a/apps/api/cmd/gateway/main.go +++ b/apps/api/cmd/gateway/main.go @@ -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() diff --git a/apps/api/docs/swagger.json b/apps/api/docs/swagger.json index 2b89a60..d1770b1 100644 --- a/apps/api/docs/swagger.json +++ b/apps/api/docs/swagger.json @@ -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" + } } } } diff --git a/apps/api/docs/swagger.yaml b/apps/api/docs/swagger.yaml index cdf65cb..dbedfcd 100644 --- a/apps/api/docs/swagger.yaml +++ b/apps/api/docs/swagger.yaml @@ -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: 列出任务 diff --git a/apps/api/internal/auth/auth.go b/apps/api/internal/auth/auth.go index c41e5f2..a97c6f3 100644 --- a/apps/api/internal/auth/auth.go +++ b/apps/api/internal/auth/auth.go @@ -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 { diff --git a/apps/api/internal/auth/oidc.go b/apps/api/internal/auth/oidc.go index 9a1f1e0..1708d52 100644 --- a/apps/api/internal/auth/oidc.go +++ b/apps/api/internal/auth/oidc.go @@ -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 } diff --git a/apps/api/internal/auth/oidc_session_cookie_test.go b/apps/api/internal/auth/oidc_session_cookie_test.go new file mode 100644 index 0000000..e7bfcb1 --- /dev/null +++ b/apps/api/internal/auth/oidc_session_cookie_test.go @@ -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") + } +} diff --git a/apps/api/internal/config/config.go b/apps/api/internal/config/config.go index c8429a5..058e81c 100644 --- a/apps/api/internal/config/config.go +++ b/apps/api/internal/config/config.go @@ -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 diff --git a/apps/api/internal/config/config_test.go b/apps/api/internal/config/config_test.go new file mode 100644 index 0000000..523824c --- /dev/null +++ b/apps/api/internal/config/config_test.go @@ -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) + } +} diff --git a/apps/api/internal/httpapi/access_rule_handlers.go b/apps/api/internal/httpapi/access_rule_handlers.go index 2062818..9601608 100644 --- a/apps/api/internal/httpapi/access_rule_handlers.go +++ b/apps/api/internal/httpapi/access_rule_handlers.go @@ -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) { diff --git a/apps/api/internal/httpapi/gemini_compat.go b/apps/api/internal/httpapi/gemini_compat.go index aac8b4b..72834f3 100644 --- a/apps/api/internal/httpapi/gemini_compat.go +++ b/apps/api/internal/httpapi/gemini_compat.go @@ -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)) } diff --git a/apps/api/internal/httpapi/handlers.go b/apps/api/internal/httpapi/handlers.go index 7d2053c..7c0ee71 100644 --- a/apps/api/internal/httpapi/handlers.go +++ b/apps/api/internal/httpapi/handlers.go @@ -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] diff --git a/apps/api/internal/httpapi/oidc_jit_integration_test.go b/apps/api/internal/httpapi/oidc_jit_integration_test.go new file mode 100644 index 0000000..8969f57 --- /dev/null +++ b/apps/api/internal/httpapi/oidc_jit_integration_test.go @@ -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))), + } +} diff --git a/apps/api/internal/httpapi/oidc_session.go b/apps/api/internal/httpapi/oidc_session.go new file mode 100644 index 0000000..6216db5 --- /dev/null +++ b/apps/api/internal/httpapi/oidc_session.go @@ -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")) != "" +} diff --git a/apps/api/internal/httpapi/oidc_session_test.go b/apps/api/internal/httpapi/oidc_session_test.go new file mode 100644 index 0000000..d4d4caa --- /dev/null +++ b/apps/api/internal/httpapi/oidc_session_test.go @@ -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 +} diff --git a/apps/api/internal/httpapi/oidc_user_middleware.go b/apps/api/internal/httpapi/oidc_user_middleware.go new file mode 100644 index 0000000..035ea0d --- /dev/null +++ b/apps/api/internal/httpapi/oidc_user_middleware.go @@ -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 +} diff --git a/apps/api/internal/httpapi/oidc_user_middleware_test.go b/apps/api/internal/httpapi/oidc_user_middleware_test.go new file mode 100644 index 0000000..3cf80b8 --- /dev/null +++ b/apps/api/internal/httpapi/oidc_user_middleware_test.go @@ -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) + } + }) + } +} diff --git a/apps/api/internal/httpapi/response.go b/apps/api/internal/httpapi/response.go index 5f3a004..27a144c 100644 --- a/apps/api/internal/httpapi/response.go +++ b/apps/api/internal/httpapi/response.go @@ -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) diff --git a/apps/api/internal/httpapi/server.go b/apps/api/internal/httpapi/server.go index 21764e8..08d07ab 100644 --- a/apps/api/internal/httpapi/server.go +++ b/apps/api/internal/httpapi/server.go @@ -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") diff --git a/apps/api/internal/httpapi/wallet_handlers.go b/apps/api/internal/httpapi/wallet_handlers.go index f02f1e9..b9ce78b 100644 --- a/apps/api/internal/httpapi/wallet_handlers.go +++ b/apps/api/internal/httpapi/wallet_handlers.go @@ -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) { diff --git a/apps/api/internal/store/identity_admin.go b/apps/api/internal/store/identity_admin.go index 235f59f..c087004 100644 --- a/apps/api/internal/store/identity_admin.go +++ b/apps/api/internal/store/identity_admin.go @@ -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) diff --git a/apps/api/internal/store/oidc_users.go b/apps/api/internal/store/oidc_users.go new file mode 100644 index 0000000..1a4ad3a --- /dev/null +++ b/apps/api/internal/store/oidc_users.go @@ -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, + } +} diff --git a/apps/api/internal/store/oidc_users_integration_test.go b/apps/api/internal/store/oidc_users_integration_test.go new file mode 100644 index 0000000..5a41b77 --- /dev/null +++ b/apps/api/internal/store/oidc_users_integration_test.go @@ -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) + } + } +} diff --git a/apps/web/src/App.tsx b/apps/web/src/App.tsx index 7459da1..e0364fa 100644 --- a/apps/web/src/App.tsx +++ b/apps/web/src/App.tsx @@ -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); } diff --git a/apps/web/src/api.test.ts b/apps/web/src/api.test.ts new file mode 100644 index 0000000..742b138 --- /dev/null +++ b/apps/web/src/api.test.ts @@ -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); + }); +}); diff --git a/apps/web/src/api.ts b/apps/web/src/api.ts index 3a59d69..b062eca 100644 --- a/apps/web/src/api.ts +++ b/apps/web/src/api.ts @@ -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 { + await request('/api/v1/auth/oidc/session', { + method: 'POST', + token: accessToken, + }); +} + +export async function deleteOIDCBrowserSession(): Promise { + await request('/api/v1/auth/oidc/session', { + auth: false, + method: 'DELETE', + }); +} + export async function getCurrentUser(token: string): Promise { return request('/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( options: { token?: string; auth?: boolean; method?: string; body?: unknown; headers?: Record } = {}, ): Promise { const headers: Record = { ...(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( 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( return response.json() as Promise; } +function authorizationHeader(token: string): Record { + 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 = { + 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 | undefined { if (!value || typeof value !== 'object' || Array.isArray(value)) return undefined; return value as Record; diff --git a/apps/web/src/lib/auth-storage.test.ts b/apps/web/src/lib/auth-storage.test.ts index bee73dc..40284d3 100644 --- a/apps/web/src/lib/auth-storage.test.ts +++ b/apps/web/src/lib/auth-storage.test.ts @@ -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(''); }); diff --git a/apps/web/src/lib/auth-storage.ts b/apps/web/src/lib/auth-storage.ts index f084496..2b4da06 100644 --- a/apps/web/src/lib/auth-storage.ts +++ b/apps/web/src/lib/auth-storage.ts @@ -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); diff --git a/apps/web/src/lib/oidc-browser-session.test.ts b/apps/web/src/lib/oidc-browser-session.test.ts new file mode 100644 index 0000000..7c53b97 --- /dev/null +++ b/apps/web/src/lib/oidc-browser-session.test.ts @@ -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(); + 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(); + }); +}); diff --git a/apps/web/src/lib/oidc-browser-session.ts b/apps/web/src/lib/oidc-browser-session.ts new file mode 100644 index 0000000..baec9bc --- /dev/null +++ b/apps/web/src/lib/oidc-browser-session.ts @@ -0,0 +1,31 @@ +import { + createOIDCBrowserSession, + GatewayApiError, + getCurrentUser, + OIDC_BROWSER_SESSION_CREDENTIAL, +} from '../api'; +import { persistAccessToken } from './auth-storage'; + +type CurrentUser = Awaited>; + +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 { + 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; + } +} diff --git a/apps/web/src/lib/oidc.ts b/apps/web/src/lib/oidc.ts index 7a02183..c1c2263 100644 --- a/apps/web/src/lib/oidc.ts +++ b/apps/web/src/lib/oidc.ts @@ -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; } diff --git a/docker-compose.yml b/docker-compose.yml index aac8ce2..39ae541 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -14,6 +14,10 @@ x-api-environment: &api-environment OIDC_REQUIRED_SCOPES: ${OIDC_REQUIRED_SCOPES:-gateway.access} OIDC_JWKS_CACHE_TTL_SECONDS: ${OIDC_JWKS_CACHE_TTL_SECONDS:-300} OIDC_ACCEPT_LEGACY_HS256: ${OIDC_ACCEPT_LEGACY_HS256:-true} + OIDC_JIT_PROVISIONING_ENABLED: ${OIDC_JIT_PROVISIONING_ENABLED:-false} + OIDC_GATEWAY_TENANT_KEY: ${OIDC_GATEWAY_TENANT_KEY:-} + OIDC_BROWSER_SESSION_ENABLED: ${OIDC_BROWSER_SESSION_ENABLED:-true} + OIDC_SESSION_COOKIE_SECURE: ${OIDC_SESSION_COOKIE_SECURE:-} SERVER_MAIN_BASE_URL: ${AI_GATEWAY_COMPOSE_SERVER_MAIN_BASE_URL:-http://host.docker.internal:3000} SERVER_MAIN_INTERNAL_TOKEN: ${SERVER_MAIN_INTERNAL_TOKEN:-change-me} SERVER_MAIN_INTERNAL_KEY: ${SERVER_MAIN_INTERNAL_KEY:-gateway} @@ -111,6 +115,7 @@ services: VITE_OIDC_ISSUER: ${OIDC_ISSUER:-} VITE_OIDC_CLIENT_ID: ${OIDC_CLIENT_ID:-} VITE_OIDC_REDIRECT_URI: ${OIDC_REDIRECT_URI:-} + VITE_OIDC_BROWSER_SESSION_ENABLED: ${OIDC_BROWSER_SESSION_ENABLED:-true} VITE_BASE_PATH: ${AI_GATEWAY_WEB_BASE_PATH:-/} ports: - "${AI_GATEWAY_WEB_PORT:-5178}:80" diff --git a/docs/design.md b/docs/design.md index c29504f..40ccf90 100644 --- a/docs/design.md +++ b/docs/design.md @@ -397,6 +397,10 @@ Gateway 需要支持三种身份运行模式,默认配置为 `IDENTITY_MODE=hy 4. 根据用户组优先级和策略合并规则得到 effective policy。 5. 创建任务时把 `gateway_user_id`、`user_source`、`user_group_id`、`user_group_key`、`user_group_policy_snapshot` 写入 `gateway_tasks`,后续重试和结算不受同步变更影响。 +Auth Center OIDC 用户采用受控 JIT 本地投影:Token 完成全部安全校验和角色授权后,Gateway 才按 `source=oidc + external_user_id=sub` 解析或幂等创建 `gateway_users` 记录。首次创建、默认用户组、`resource` 钱包和脱敏审计在同一事务完成;不按邮箱或昵称关联历史用户,不自动创建租户,也不向认证中心 Claims 写入 Gateway 本地 ID。详细配置与错误语义见 `docs/oidc-jit-provisioning.md`。 + +Gateway Web Console 默认使用 HttpOnly OIDC Cookie 会话解决标签页隔离问题:Access Token 仍由 Auth Center 签发,Gateway 不二次签发 JWT;前端建立 Cookie 后不再持久化 OIDC Access Token。Cookie 共享到同域标签页,有效期不超过原 Token,并通过 SameSite、Secure、精确 CORS 和写请求 Origin 校验控制 CSRF 风险。 + ### 7.0.1 多租户模型 多租户支持不能只停留在 claim 的 `tenantId` 字符串,Gateway 需要有自己的租户表和执行上下文: diff --git a/docs/oidc-jit-provisioning.md b/docs/oidc-jit-provisioning.md new file mode 100644 index 0000000..aaaf9aa --- /dev/null +++ b/docs/oidc-jit-provisioning.md @@ -0,0 +1,63 @@ +# Auth Center OIDC 用户受控 JIT 预配 + +## 适用边界 + +Gateway 只在 OIDC Token 已通过签名、Issuer、Audience、有效期、`tid`、Scope 和应用角色校验后执行 JIT。认证中心的稳定 `sub` 是外部用户标识;Gateway 不读取、不保存或公开 Keycloak 内部 ID,也不向 Token 增加 `gatewayUserId`。 + +本次保持单 Gateway 租户:部署方用 `OIDC_GATEWAY_TENANT_KEY` 把 Auth Center 的已验证 `tid` 显式绑定到一个已存在、启用且配置了启用中默认用户组的 Gateway 租户。Token 不能创建租户。 + +## 配置 + +```dotenv +OIDC_ENABLED=true +OIDC_JIT_PROVISIONING_ENABLED=true +OIDC_GATEWAY_TENANT_KEY=default +OIDC_BROWSER_SESSION_ENABLED=true +OIDC_SESSION_COOKIE_SECURE=false # 本地 HTTP;生产必须为 true +``` + +- `OIDC_JIT_PROVISIONING_ENABLED` 默认 `false`。关闭时停止创建新用户,但已有 `source=oidc + external_user_id=sub` 映射仍会解析和同步登录时间。 +- JIT 开启时 `OIDC_GATEWAY_TENANT_KEY` 必填,缺失会使 Gateway 启动失败。 +- 本地与 Staging 验收环境应在各自 Git 忽略或 Secret 管理的环境文件中显式开启;生产启用需独立变更审批。 + +## Web Console 浏览器会话 + +- OIDC 回调仍使用 Authorization Code + PKCE 从 Auth Center 获取 Access Token,但 Token 只在回调函数内存中短暂存在。 +- Web 随即调用 `POST /api/v1/auth/oidc/session`;Gateway 再次完成全部 OIDC 校验后,将原 Auth Center Access Token 写入 `HttpOnly + SameSite=Strict` Cookie,不签发第二枚 Gateway JWT。 +- Cookie 有效期不超过 Access Token 的 `exp`,所有 Cookie 请求仍经过 JWKS、Issuer、Audience、`tid`、Scope、角色及可选 Introspection 校验。 +- 建立成功后清除旧的 OIDC `sessionStorage` 和可能冲突的 `localStorage` Token;新标签页通过 Cookie 调用 `/api/v1/me` 自动恢复同一登录态。 +- 所有 Web API 请求使用 `credentials: include`。Cookie 鉴权的 POST、PUT、PATCH、DELETE 必须携带 `CORS_ALLOWED_ORIGIN` 白名单中的 Origin,否则返回结构化 403。 +- Staging、生产及其他非本地环境启动时强制 `OIDC_SESSION_COOKIE_SECURE=true`,并拒绝带 `*` 的凭据型 CORS 配置。本地 HTTP 开发和自动化测试可以显式设为 `false`。 +- `DELETE /api/v1/auth/oidc/session` 使 Cookie 立即过期,然后前端进入 Auth Center OIDC 登出。浏览器会话关闭时不会创建或撤销 Gateway API Key。 +- 回滚时同时设置后端 `OIDC_BROWSER_SESSION_ENABLED=false` 和 Web 构建参数 `VITE_OIDC_BROWSER_SESSION_ENABLED=false`,恢复旧的标签页级 `sessionStorage` 行为。 + +## 数据和事务语义 + +- 查找键固定为 `source=oidc + external_user_id=sub`;`user_key` 由规范化 Issuer 和 `sub` 做 SHA-256 派生,不按邮箱、手机号、昵称匹配。 +- 首次登录在一个事务中创建 `gateway_users` 投影、绑定租户默认用户组、初始化 `resource` 钱包并写入 `identity.oidc_user.provisioned` 审计事件。 +- 重复或并发登录依赖现有唯一约束和冲突处理返回同一个本地用户,不重复创建钱包或首次预配审计。 +- 后续登录只同步 Token 用户名、应用角色、`last_login_at`、`synced_at` 和 `source_updated_at`;不覆盖人工资料或用户组,不重新启用已禁用/删除账号。 +- JIT 不生成 API Key 或任何 secret。用户仅在主动创建 API Key 或进入需要 Key 的工作流时触发生命周期操作。 + +## 对外错误 + +所有错误沿用 `ErrorEnvelope`: + +| HTTP | code | 含义 | +| --- | --- | --- | +| 403 | `GATEWAY_USER_NOT_PROVISIONED` | JIT 关闭且没有已有本地映射 | +| 403 | `GATEWAY_USER_DISABLED` | 本地用户已禁用或删除 | +| 503 | `GATEWAY_TENANT_UNAVAILABLE` | 配置租户或用户组不存在/未启用 | +| 503 | `GATEWAY_USER_PROVISIONING_FAILED` | 事务或存储故障;响应不暴露内部错误 | +| 404 | `OIDC_BROWSER_SESSION_DISABLED` | Gateway 浏览器会话功能未启用 | +| 401 | `OIDC_SESSION_INVALID` | 不是有效的 Auth Center OIDC Access Token 或已过期 | +| 400 | `OIDC_SESSION_TOKEN_TOO_LARGE` | Access Token 超过安全 Cookie 大小限制 | +| 403 | `OIDC_SESSION_CSRF_REJECTED` | Cookie 写请求缺少可信 Origin | + +`ErrLocalUserRequired` 仅作为防御性内部错误保留,对外统一转换为结构化 403。 + +## 验收与回滚 + +自动化门禁通过后,依次验证本地真实 OIDC 和 `auth.51easyai.com` Staging 专用测试租户。证据应包含脱敏 Claims、网络截图、本地 Gateway 用户记录、Trace ID、Gateway/Auth Center 审计 ID及 200/401/403/503 负向证据;不得保存 Token、授权码、密码或 Secret。 + +JIT 回滚时关闭 `OIDC_JIT_PROVISIONING_ENABLED` 并回退 Gateway 镜像。浏览器 Cookie 会话可通过后端和 Web 两侧的独立开关关闭。已经创建的 `source=oidc` 测试投影保留为惰性数据,不自动删除、迁移或合并。 diff --git a/scripts/dev.sh b/scripts/dev.sh index 73199ac..cf14c80 100755 --- a/scripts/dev.sh +++ b/scripts/dev.sh @@ -25,6 +25,7 @@ load_local_env() { export VITE_OIDC_ISSUER="${VITE_OIDC_ISSUER:-${OIDC_ISSUER:-}}" export VITE_OIDC_CLIENT_ID="${VITE_OIDC_CLIENT_ID:-${OIDC_CLIENT_ID:-}}" export VITE_OIDC_REDIRECT_URI="${VITE_OIDC_REDIRECT_URI:-${OIDC_REDIRECT_URI:-}}" + export VITE_OIDC_BROWSER_SESSION_ENABLED="${VITE_OIDC_BROWSER_SESSION_ENABLED:-${OIDC_BROWSER_SESSION_ENABLED:-true}}" } load_local_env