From 4b72419a964f2719c97a26486ab3d6e19c3cc512 Mon Sep 17 00:00:00 2001 From: ryan Date: Sun, 7 Jun 2026 23:55:19 +0800 Subject: [PATCH] =?UTF-8?q?=E5=8E=BB=E9=99=A4=E9=BB=98=E8=AE=A4OIDC?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 1 + Agents.md | 1 + README.md | 11 +- README_zh.md | 11 +- config.example.yaml | 13 +- docker-compose.yml | 61 +++++++ docs/docs.go | 4 +- docs/swagger.json | 4 +- docs/swagger.yaml | 4 +- frontend/components/auth/login-form.tsx | 45 +++--- .../components/common/settings/security.tsx | 140 +++++++++++++++- internal/apps/oauth/config.go | 73 --------- internal/apps/oauth/logics.go | 152 ------------------ internal/apps/oauth/oauth_test.go | 139 ++++++++++------ internal/apps/oauth/sources.go | 118 +++----------- internal/apps/upload/tasks.go | 2 +- internal/config/model.go | 12 -- internal/model/auth_source.go | 2 +- 18 files changed, 356 insertions(+), 437 deletions(-) create mode 100644 docker-compose.yml delete mode 100644 internal/apps/oauth/config.go delete mode 100644 internal/apps/oauth/logics.go diff --git a/.gitignore b/.gitignore index 244998fd..b324e755 100644 --- a/.gitignore +++ b/.gitignore @@ -44,3 +44,4 @@ uploads/* s3_cache /frontend/.next/ +/data/ diff --git a/Agents.md b/Agents.md index 0837bba0..fc832340 100644 --- a/Agents.md +++ b/Agents.md @@ -46,6 +46,7 @@ Refreshing/ # 项目根目录(模块名: github.com/lin ├── config.example.yaml # 配置模板(需提交) ├── Makefile # 常用命令(swagger/tidy/license) ├── Dockerfile # 后端容器镜像构建 +├── docker-compose.yml # 本地依赖服务(PostgreSQL / Redis / ClickHouse) ├── .editorconfig # 编辑器格式规范 ├── .gitignore ├── docs/ # Swagger 自动生成文档(不要手动编辑) diff --git a/README.md b/README.md index c3064feb..17938ab7 100644 --- a/README.md +++ b/README.md @@ -98,12 +98,18 @@ cd refreshing cp config.example.yaml config.yaml ``` -Edit `config.yaml` to configure your database, Redis, and at least one auth source (OIDC or password-based). +Edit `config.yaml` to configure your database and Redis. OIDC auth sources are configured at runtime in the admin settings page. ### 3. Initialize Database ```bash -# Create the database +# Start local dependencies (PostgreSQL + Redis) +docker compose up -d + +# Optional: also start ClickHouse +docker compose --profile clickhouse up -d + +# If you use an external PostgreSQL instance instead of Docker, create the database manually createdb -h -p 5432 -U postgres refreshing # Database schema is auto-migrated on first startup @@ -159,7 +165,6 @@ Key configuration options (see `config.example.yaml` for the full reference): | `database.database` | Database name | `refreshing` | | `redis.host` | Redis host | `127.0.0.1` | | `storage.endpoint` | S3-compatible endpoint | `s3.amazonaws.com` | -| `oauth2.client_id` | Default OIDC client ID | `your_client_id` | ## 🔧 Development Guide diff --git a/README_zh.md b/README_zh.md index 1335b828..bacf6cd1 100644 --- a/README_zh.md +++ b/README_zh.md @@ -98,12 +98,18 @@ cd refreshing cp config.example.yaml config.yaml ``` -编辑 `config.yaml`,配置数据库、Redis,以及至少一个认证源(OIDC 或密码登录)。 +编辑 `config.yaml`,配置数据库和 Redis。OIDC 认证源统一在管理后台的系统设置页面运行时配置。 ### 3. 初始化数据库 ```bash -# 创建数据库 +# 启动本地依赖服务(PostgreSQL + Redis) +docker compose up -d + +# 可选:同时启动 ClickHouse +docker compose --profile clickhouse up -d + +# 如果使用外部 PostgreSQL,而不是 Docker 内置服务,则手动创建数据库 createdb -h <主机> -p 5432 -U postgres refreshing # 数据库表结构在首次启动时自动迁移,无需手动执行 @@ -159,7 +165,6 @@ pnpm dev | `database.database` | 数据库名称 | `refreshing` | | `redis.host` | Redis 主机 | `127.0.0.1` | | `storage.endpoint` | S3 兼容存储端点 | `s3.amazonaws.com` | -| `oauth2.client_id` | 默认 OIDC 客户端 ID | `your_client_id` | ## 🔧 开发指南 diff --git a/config.example.yaml b/config.example.yaml index baa03d6b..03b70c0f 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -18,17 +18,6 @@ app: api_prefix: "/api" frontend_url: "http://localhost:3000" -# ─── Default OAuth2 / OIDC Provider (optional) ────────────────────────────────── -# You can configure additional OIDC providers at runtime via the admin panel. -oauth2: - client_id: "" - client_secret: "" - redirect_uri: "" - issuer: "" # OIDC Issuer URL for auto-discovery (recommended) - authorization_endpoint: "" # Leave empty if issuer is set - token_endpoint: "" - user_endpoint: "" - # ─── PostgreSQL ───────────────────────────────────────────────────────────────── # Supports Standalone and Primary-Replica (read/write split) modes. database: @@ -36,7 +25,7 @@ database: host: "127.0.0.1" port: 5432 username: "postgres" - password: "" + password: "postgres" database: "refreshing" max_idle_conn: 16 max_open_conn: 128 diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 00000000..ef3ad414 --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,61 @@ +services: + postgres: + image: postgres:17-alpine + container_name: refreshing-postgres + restart: unless-stopped + environment: + POSTGRES_DB: ${POSTGRES_DB:-refreshing} + POSTGRES_USER: ${POSTGRES_USER:-postgres} + POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:-postgres} + TZ: ${TZ:-Asia/Shanghai} + ports: + - "${POSTGRES_PORT:-5432}:5432" + volumes: + - ./data/postgres_data:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U ${POSTGRES_USER:-postgres} -d ${POSTGRES_DB:-refreshing}"] + interval: 10s + timeout: 5s + retries: 5 + start_period: 10s +# +# redis: +# image: redis:7-alpine +# container_name: refreshing-redis +# restart: unless-stopped +# command: ["redis-server", "--appendonly", "yes"] +# ports: +# - "${REDIS_PORT:-6379}:6379" +# volumes: +# - ./data/redis_data:/data +# healthcheck: +# test: ["CMD", "redis-cli", "ping"] +# interval: 10s +# timeout: 5s +# retries: 5 +# start_period: 5s +# +# clickhouse: +# image: clickhouse/clickhouse-server:25.3-alpine +# container_name: refreshing-clickhouse +# restart: unless-stopped +# profiles: +# - clickhouse +# environment: +# CLICKHOUSE_DB: ${CLICKHOUSE_DB:-refreshing} +# CLICKHOUSE_USER: ${CLICKHOUSE_USER:-default} +# CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD:-} +# CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: 1 +# TZ: ${TZ:-Asia/Shanghai} +# ports: +# - "${CLICKHOUSE_HTTP_PORT:-8123}:8123" +# - "${CLICKHOUSE_NATIVE_PORT:-9000}:9000" +# volumes: +# - clickhouse_data:/var/lib/clickhouse +# - clickhouse_logs:/var/log/clickhouse-server +# healthcheck: +# test: ["CMD", "clickhouse-client", "--query", "SELECT 1"] +# interval: 10s +# timeout: 5s +# retries: 5 +# start_period: 15s diff --git a/docs/docs.go b/docs/docs.go index 9d4ee190..a256fba5 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -1485,7 +1485,7 @@ const docTemplate = `{ }, "/api/v1/oauth/login": { "get": { - "description": "根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用默认认证源。", + "description": "根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。", "produces": [ "application/json" ], @@ -1496,7 +1496,7 @@ const docTemplate = `{ "parameters": [ { "type": "string", - "description": "认证源名称,为空使用默认源", + "description": "认证源名称,为空使用第一个启用的认证源", "name": "source", "in": "query" } diff --git a/docs/swagger.json b/docs/swagger.json index 44561cfd..4ccaea83 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -1478,7 +1478,7 @@ }, "/api/v1/oauth/login": { "get": { - "description": "根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用默认认证源。", + "description": "根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。", "produces": [ "application/json" ], @@ -1489,7 +1489,7 @@ "parameters": [ { "type": "string", - "description": "认证源名称,为空使用默认源", + "description": "认证源名称,为空使用第一个启用的认证源", "name": "source", "in": "query" } diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 8a3ad1df..5a001adf 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -1436,9 +1436,9 @@ paths: - oauth /api/v1/oauth/login: get: - description: 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用默认认证源。 + description: 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。 parameters: - - description: 认证源名称,为空使用默认源 + - description: 认证源名称,为空使用第一个启用的认证源 in: query name: source type: string diff --git a/frontend/components/auth/login-form.tsx b/frontend/components/auth/login-form.tsx index 66a67281..2264d35f 100644 --- a/frontend/components/auth/login-form.tsx +++ b/frontend/components/auth/login-form.tsx @@ -237,31 +237,28 @@ export function LoginForm() { -
-
- - 第三方认证源 -
-
- {authSources.length > 0 ? ( - authSources.map((source) => ( - - )) - ) : ( -
- 暂无可用认证源 + {authSources.length > 0 ? ( + authSources.map((source) => ( +
+
+ + 第三方认证源
- )} -
-
+
+ +
+
+ )) + ) : null} + ) diff --git a/frontend/components/common/settings/security.tsx b/frontend/components/common/settings/security.tsx index d54b4dd9..13009ad9 100644 --- a/frontend/components/common/settings/security.tsx +++ b/frontend/components/common/settings/security.tsx @@ -2,9 +2,24 @@ import {useEffect, useMemo, useState} from "react" import {useMutation, useQuery, useQueryClient} from "@tanstack/react-query" -import {Fingerprint, Globe, Loader2, Lock, Pencil, Plus, Settings, Trash2, UserPlus} from "lucide-react" +import { + CalendarClock, + Fingerprint, + Globe, + Info, + Loader2, + Lock, + Monitor, + Pencil, + Plus, + Server, + Settings, + Trash2, + UserPlus +} from "lucide-react" import {useRouter} from "next/navigation" import {motion} from "motion/react" +import packageJson from "../../../package.json" import {Button} from "@/components/ui/button" import {Card, CardContent, CardDescription, CardHeader, CardTitle} from "@/components/ui/card" @@ -12,7 +27,7 @@ import {Switch} from "@/components/ui/switch" import {Tabs, TabsContent, TabsList, TabsTrigger} from "@/components/ui/tabs" import {useAuth} from "@/components/providers/auth-provider" import {AuthSourceModal} from "@/components/common/settings/auth-source-modal" -import {AdminService} from "@/lib/services" +import {AdminService, apiConfig} from "@/lib/services" import type {AuthSource, SystemConfig} from "@/lib/services/admin" import {toast} from "sonner" @@ -52,12 +67,33 @@ function systemConfigMap(configs: SystemConfig[]) { }, {}) } +function InfoRow({ label, value }: { label: string; value: React.ReactNode }) { + return ( +
+ {label} + {value || "-"} +
+ ) +} + +function formatBooleanConfig(config?: SystemConfig) { + if (!config) return "未配置" + return config.value === "true" ? "启用" : "禁用" +} + export function SecurityMain() { const queryClient = useQueryClient() const { user, loading } = useAuth() const router = useRouter() const [authSourceModalOpen, setAuthSourceModalOpen] = useState(false) const [selectedSource, setSelectedSource] = useState(null) + const [runtimeInfo, setRuntimeInfo] = useState({ + language: "-", + platform: "-", + timezone: "-", + viewport: "-", + userAgent: "-", + }) const systemConfigsQuery = useQuery({ queryKey: ["admin", "system-configs"], @@ -83,10 +119,14 @@ export function SecurityMain() { }, [user, loading, router]) useEffect(() => { - if (user?.is_admin) { - void systemConfigsQuery.refetch() - } - }, [systemConfigsQuery, user]) + setRuntimeInfo({ + language: navigator.language || "-", + platform: navigator.platform || "-", + timezone: Intl.DateTimeFormat().resolvedOptions().timeZone || "-", + viewport: `${window.innerWidth} x ${window.innerHeight}`, + userAgent: navigator.userAgent || "-", + }) + }, []) const updateConfigMutation = useMutation({ mutationFn: async ({ key, value }: { key: SecurityKey; value: boolean }) => { @@ -341,7 +381,93 @@ export function SecurityMain() { - + +
+ + +
+
+ +
+
+ 应用信息 + 当前前端应用的版本与构建信息 +
+
+
+ + + + + + + +
+ + + +
+
+ +
+
+ 服务连接 + 前端 API 客户端的基础连接参数 +
+
+
+ + + + + + + +
+ + + +
+
+ +
+
+ 运行环境 + 当前浏览器会话的本地运行信息 +
+
+
+ + + + + + + +
+ + + +
+
+ +
+
+ 安全配置概览 + 当前系统登录与注册开关状态 +
+
+
+ + + + + + + +
+
+
将占用者改名并注销 - newParams := map[string]interface{}{ - "username": fmt.Sprintf("%s已注销: %s", holder.Username, uuid.NewString()), - "is_active": false, - } - if updateErr := tx.Model(&holder).Updates(newParams).Error; updateErr != nil { - return updateErr - } - } - - // 根据 ID 处理当前用户的 更新 或 创建 - if queryErr := tx.Where("id = ?", userInfo.GetID()).First(&user).Error; queryErr == nil { - // 用户已存在 -> 更新信息 - if activeErr := user.CheckActive(); activeErr != nil { - return activeErr - } - user.UpdateFromOAuthInfo(&userInfo) - if saveErr := tx.Save(&user).Error; saveErr != nil { - return saveErr - } - } else if errors.Is(queryErr, gorm.ErrRecordNotFound) { - // 用户不存在 -> 创建新用户 - user = model.User{} - if createErr := user.CreateUser(tx, &userInfo); createErr != nil { - return createErr - } - } else { - return queryErr - } - - return nil - }) - if err != nil { - span.SetStatus(codes.Error, err.Error()) - return nil, err - } - - return &user, nil -} diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index 1a402068..df87f603 100644 --- a/internal/apps/oauth/oauth_test.go +++ b/internal/apps/oauth/oauth_test.go @@ -27,6 +27,7 @@ import ( "net/http" "net/http/httptest" "net/url" + "strconv" "strings" "testing" "time" @@ -134,6 +135,17 @@ var ( testJWKS jose.JSONWebKeySet ) +const ( + testIssuerURL = "https://connect.linux.do" + testAuthURL = "https://connect.linux.do/oauth2/authorize" + testTokenURL = "https://connect.linux.do/oauth2/token" + testJWKSURL = "https://connect.linux.do/oauth2/keys" + testClientID = "test_client_id" + testClientSecret = "test_client_secret" + testSourceName = "linuxdo" + testSourceDisplay = "LINUX DO" +) + func init() { var err error testRSAPrivateKey, err = rsa.GenerateKey(rand.Reader, 2048) @@ -151,7 +163,50 @@ func init() { } } +func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) { + t.Helper() + if err := dbConn.Create(&model.AuthSource{ + ID: 100, + Name: testSourceName, + Type: model.AuthSourceTypeOIDC, + DisplayName: testSourceDisplay, + IsActive: true, + ClientID: testClientID, + ClientSecret: testClientSecret, + OpenIDDiscoveryURL: testIssuerURL, + }).Error; err != nil { + t.Fatalf("failed to seed auth source: %v", err) + } +} + +func oidcDiscoveryResponse() *http.Response { + body := fmt.Sprintf(`{ + "issuer": %q, + "authorization_endpoint": %q, + "token_endpoint": %q, + "jwks_uri": %q, + "response_types_supported": ["code"], + "subject_types_supported": ["public"], + "id_token_signing_alg_values_supported": ["RS256"] + }`, testIssuerURL, testAuthURL, testTokenURL, testJWKSURL) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(strings.NewReader(body)), + Header: make(http.Header), + } +} + +func jwksResponse() *http.Response { + jwksJSON, _ := json.Marshal(testJWKS) + return &http.Response{ + StatusCode: http.StatusOK, + Body: io.NopCloser(bytes.NewReader(jwksJSON)), + Header: make(http.Header), + } +} + type mockClaims struct { + ID uint64 `json:"id"` Issuer string `json:"iss"` Subject string `json:"sub"` Audience string `json:"aud"` @@ -169,7 +224,9 @@ func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string) if err != nil { panic(err) } + id, _ := strconv.ParseUint(sub, 10, 64) claims := mockClaims{ + ID: id, Issuer: issuer, Subject: sub, Audience: aud, @@ -236,19 +293,6 @@ func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *ht db.SetDB(dbConn) db.Redis = mockRedis - // Setup oauth vars - oauthConf = &oauth2.Config{ - ClientID: config.Config.OAuth2.ClientID, - ClientSecret: config.Config.OAuth2.ClientSecret, - RedirectURL: config.Config.OAuth2.RedirectURI, - Scopes: []string{"profile", "email"}, - Endpoint: oauth2.Endpoint{ - AuthURL: config.Config.OAuth2.AuthorizationEndpoint, - TokenURL: config.Config.OAuth2.TokenEndpoint, - }, - } - oidcVerifier = nil - api := r.Group("/api/v1") { api.GET("/oauth/sources", GetLoginSources) @@ -288,12 +332,7 @@ func initializeTestConfig() { config.Config.App.SessionCookieName = "test_session_id" config.Config.App.SessionSecret = "test_session_secret" config.Config.App.APIPrefix = "/api" - config.Config.OAuth2.ClientID = "test_client_id" - config.Config.OAuth2.ClientSecret = "test_client_secret" - config.Config.OAuth2.RedirectURI = "http://localhost:3000/callback" - config.Config.OAuth2.AuthorizationEndpoint = "https://connect.linux.do/oauth2/authorize" - config.Config.OAuth2.TokenEndpoint = "https://connect.linux.do/oauth2/token" - config.Config.OAuth2.UserEndpoint = "https://connect.linux.do/api/user" + config.Config.App.FrontendURL = "http://localhost:3000" config.Config.OpenAPIRisk.Enabled = false config.Config.OpenAPIRisk.BaseURL = "" config.Config.OpenAPIRisk.BlockRiskLevels = []string{} @@ -352,13 +391,12 @@ func TestGetLoginSources(t *testing.T) { t.Fatalf("failed to unmarshal response: %v", err) } - if len(resp.Data) != 2 { - t.Fatalf("expected 2 active sources, got %d", len(resp.Data)) + if len(resp.Data) != 1 { + t.Fatalf("expected 1 active source, got %d", len(resp.Data)) } - names := []string{resp.Data[0].Name, resp.Data[1].Name} - if !strings.Contains(strings.Join(names, ","), "default") || !strings.Contains(strings.Join(names, ","), "github") { - t.Errorf("expected default and github sources, got %v", names) + if resp.Data[0].Name != "github" { + t.Errorf("expected github source, got %s", resp.Data[0].Name) } // Test disabling OIDC @@ -379,10 +417,14 @@ func TestGetLoginURL(t *testing.T) { initializeTestConfig() dbConn := setupTestDB(t) mockRedis := newMockRedisClient() + seedTestAuthSource(t, dbConn) httpMock := &http.Client{ Transport: &mockRoundTripper{ roundTripFunc: func(req *http.Request) (*http.Response, error) { + if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") { + return oidcDiscoveryResponse(), nil + } return nil, fmt.Errorf("unexpected request") }, }, @@ -404,7 +446,7 @@ func TestGetLoginURL(t *testing.T) { t.Fatalf("failed to unmarshal response: %v", err) } - if !strings.Contains(resp.Data.AuthorizeURL, config.Config.OAuth2.AuthorizationEndpoint) { + if !strings.Contains(resp.Data.AuthorizeURL, testAuthURL) { t.Errorf("invalid authorize URL: %s", resp.Data.AuthorizeURL) } @@ -424,7 +466,7 @@ func TestGetLoginURL(t *testing.T) { if err != nil { t.Fatalf("failed to decode state payload: %v", err) } - if payload.SourceName != "default" || payload.Purpose != OAuthPurposeLogin { + if payload.SourceName != testSourceName || payload.Purpose != OAuthPurposeLogin { t.Errorf("unexpected payload: %+v", payload) } @@ -525,23 +567,23 @@ func TestCallbackLoginAndUserInfo(t *testing.T) { initializeTestConfig() dbConn := setupTestDB(t) mockRedis := newMockRedisClient() + seedTestAuthSource(t, dbConn) + state := uuid.NewString() // 1. Mock the outgoing HTTP client for token exchange and user info fetching httpMock := &http.Client{ Transport: &mockRoundTripper{ roundTripFunc: func(req *http.Request) (*http.Response, error) { - // Handle Token Exchange - if req.Method == http.MethodPost && req.URL.String() == config.Config.OAuth2.TokenEndpoint { - body := `{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600}` - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(strings.NewReader(body)), - Header: make(http.Header), - }, nil + if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") { + return oidcDiscoveryResponse(), nil } - // Handle User Info fetch - if req.Method == http.MethodGet && req.URL.String() == config.Config.OAuth2.UserEndpoint { - body := `{"id":88888,"sub":"oauth_user_88888","username":"test_oauth_user","email":"oauth@linux.do","name":"Oauth Test User","active":true,"trust_level":3}` + if req.Method == http.MethodGet && req.URL.String() == testJWKSURL { + return jwksResponse(), nil + } + // Handle Token Exchange + if req.Method == http.MethodPost && req.URL.String() == testTokenURL { + idToken := generateMockIDToken(testIssuerURL, "88888", testClientID, state, "test_oauth_user", "oauth@linux.do", "Oauth Test User") + body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken) return &http.Response{ StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(body)), @@ -556,9 +598,8 @@ func TestCallbackLoginAndUserInfo(t *testing.T) { router := setupTestRouter(dbConn, mockRedis, httpMock) // 2. Setup state in Redis - state := uuid.NewString() payloadValue, _ := encodeOAuthStatePayload(oauthStatePayload{ - SourceName: "default", + SourceName: testSourceName, Purpose: OAuthPurposeLogin, }) mockRedis.Set(context.Background(), db.PrefixedKey(fmt.Sprintf(OAuthStateCacheKeyFormat, state)), payloadValue, OAuthStateCacheKeyExpiration) @@ -582,13 +623,13 @@ func TestCallbackLoginAndUserInfo(t *testing.T) { t.Errorf("expected logged_in status, got %s", callbackResp.Data.Status) } - if callbackResp.Data.User.Username != "test_oauth_user" || callbackResp.Data.User.ID != 1 { + if callbackResp.Data.User.Username != "test_oauth_user" || callbackResp.Data.User.ID != 88888 { t.Errorf("unexpected user returned: %+v", callbackResp.Data.User) } // Verify user is created in database var user model.User - if err := dbConn.First(&user, "id = ?", 1).Error; err != nil { + if err := dbConn.First(&user, "id = ?", 88888).Error; err != nil { t.Fatalf("user was not created in DB: %v", err) } @@ -619,15 +660,15 @@ func TestCallbackLoginAndUserInfo(t *testing.T) { httpMock2 := &http.Client{ Transport: &mockRoundTripper{ roundTripFunc: func(req *http.Request) (*http.Response, error) { - if req.Method == http.MethodPost && req.URL.String() == config.Config.OAuth2.TokenEndpoint { - body := `{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600}` - return &http.Response{ - StatusCode: http.StatusOK, - Body: io.NopCloser(strings.NewReader(body)), - }, nil + if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/.well-known/openid-configuration") { + return oidcDiscoveryResponse(), nil } - if req.Method == http.MethodGet && req.URL.String() == config.Config.OAuth2.UserEndpoint { - body := `{"id":99999,"sub":"oauth_user_99999","username":"test_oauth_user","email":"another@linux.do","name":"Another User","active":true,"trust_level":1}` + if req.Method == http.MethodGet && req.URL.String() == testJWKSURL { + return jwksResponse(), nil + } + if req.Method == http.MethodPost && req.URL.String() == testTokenURL { + idToken := generateMockIDToken(testIssuerURL, "99999", testClientID, state2, "test_oauth_user", "another@linux.do", "Another User") + body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken) return &http.Response{ StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(body)), diff --git a/internal/apps/oauth/sources.go b/internal/apps/oauth/sources.go index ce2838f9..8c50577a 100644 --- a/internal/apps/oauth/sources.go +++ b/internal/apps/oauth/sources.go @@ -2,10 +2,8 @@ package oauth import ( "context" - "encoding/json" "errors" "fmt" - "io" "net/http" "strings" "time" @@ -53,33 +51,30 @@ type CallbackRequest struct { Code string `json:"code" binding:"required"` } -func defaultAuthSource() *model.AuthSource { - if config.Config.OAuth2.ClientID == "" || config.Config.OAuth2.RedirectURI == "" { - return nil +func GetUserIDFromSession(s sessions.Session) uint64 { + userID, ok := s.Get(UserIDKey).(uint64) + if !ok { + return 0 } - source := &model.AuthSource{ - Name: "default", - Type: model.AuthSourceTypeOIDC, - DisplayName: "默认认证源", - IsActive: true, - ClientID: config.Config.OAuth2.ClientID, - ClientSecret: config.Config.OAuth2.ClientSecret, - OpenIDDiscoveryURL: config.Config.OAuth2.Issuer, - } - if source.DisplayName == "" { - source.DisplayName = "默认认证源" - } - return source + return userID +} + +func GetUserIDFromContext(c *gin.Context) uint64 { + session := sessions.Default(c) + return GetUserIDFromSession(session) } func resolveAuthSource(sourceName string) (*model.AuthSource, error) { name := strings.TrimSpace(strings.ToLower(sourceName)) - if name == "" || name == "default" { - source := defaultAuthSource() - if source == nil { - return nil, errors.New("默认认证源未配置") + if name == "" { + sources, err := model.GetActiveAuthSources() + if err != nil { + return nil, err } - return source, nil + if len(sources) == 0 { + return nil, errors.New("未配置可用认证源") + } + return &sources[0], nil } return model.GetAuthSourceByName(name) } @@ -90,24 +85,11 @@ func activeLoginSources() []AuthSourceView { return nil } - sources := make([]AuthSourceView, 0, 4) - if source := defaultAuthSource(); source != nil { - source.Sanitize() - sources = append(sources, AuthSourceView{ - ID: source.ID, - Name: source.Name, - Type: source.Type, - DisplayName: source.DisplayName, - IsActive: source.IsActive, - IconURL: source.IconURL, - ClientSecretConfigured: source.ClientSecretConfigured, - }) - } - dbSources, err := model.GetActiveAuthSources() if err != nil { - return sources + return nil } + sources := make([]AuthSourceView, 0, len(dbSources)) for _, source := range dbSources { sources = append(sources, AuthSourceView{ ID: source.ID, @@ -126,9 +108,6 @@ func frontendLoginRedirectURL() string { if config.Config.App.FrontendURL != "" { return strings.TrimRight(config.Config.App.FrontendURL, "/") + "/login" } - if config.Config.OAuth2.RedirectURI != "" { - return config.Config.OAuth2.RedirectURI - } return "/login" } @@ -137,24 +116,6 @@ func buildOAuthConfig(ctx context.Context, source *model.AuthSource, redirectURL return nil, nil, errors.New("认证源不能为空") } - if source.Name == "default" { - scopes := []string{"profile", "email"} - if oidcVerifier != nil { - scopes = append([]string{oidc.ScopeOpenID}, scopes...) - } - return &oauth2.Config{ - ClientID: config.Config.OAuth2.ClientID, - ClientSecret: config.Config.OAuth2.ClientSecret, - RedirectURL: redirectURL, - Scopes: scopes, - Endpoint: oauth2.Endpoint{ - AuthURL: config.Config.OAuth2.AuthorizationEndpoint, - TokenURL: config.Config.OAuth2.TokenEndpoint, - AuthStyle: oauth2.AuthStyleAutoDetect, - }, - }, oidcVerifier, nil - } - if source.OpenIDDiscoveryURL == "" { return nil, nil, errors.New("OIDC 认证源必须配置 Discovery URL") } @@ -260,29 +221,6 @@ func buildOAuthUserInfo(ctx context.Context, source *model.AuthSource, code stri userInfo.Name = userInfo.Username } - if userInfo.Username == "" { - client := authConfig.Client(ctx, token) - userEndpoint := config.Config.OAuth2.UserEndpoint - if source.Name != "default" { - userEndpoint = "" - } - if userEndpoint != "" { - resp, httpErr := client.Get(userEndpoint) - if httpErr != nil { - return nil, httpErr - } - defer resp.Body.Close() - - responseData, readErr := io.ReadAll(resp.Body) - if readErr != nil { - return nil, readErr - } - if unmarshalErr := json.Unmarshal(responseData, userInfo); unmarshalErr != nil { - return nil, unmarshalErr - } - } - } - return userInfo, nil } @@ -336,10 +274,10 @@ func GetLoginSources(c *gin.Context) { // GetLoginURL 获取登录授权地址 // @Summary 获取登录授权地址 -// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用默认认证源。 +// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。 // @Tags oauth // @Produce json -// @Param source query string false "认证源名称,为空使用默认源" +// @Param source query string false "认证源名称,为空使用第一个启用的认证源" // @Success 200 {object} util.ResponseAny{data=oauth.OAuthAuthorizeResponse} "授权 URL" // @Failure 400 {object} util.ResponseAny "认证源不存在或未配置" // @Failure 500 {object} util.ResponseAny "Redis 异常或构造 URL 失败" @@ -522,16 +460,8 @@ func Callback(c *gin.Context) { c.JSON(http.StatusInternalServerError, util.Err(uniqueErr.Error())) return } - user = model.User{ - Username: username, - Nickname: userInfo.Name, - AvatarUrl: userInfo.AvatarUrl, - TrustLevel: userInfo.TrustLevel, - SignKey: util.GenerateUniqueIDSimple(), - IsActive: true, - LastLoginAt: time.Now(), - } - if err := db.DB(ctx).Create(&user).Error; err != nil { + userInfo.Username = username + if err := user.CreateUser(db.DB(ctx), userInfo); err != nil { c.JSON(http.StatusInternalServerError, util.Err(err.Error())) return } diff --git a/internal/apps/upload/tasks.go b/internal/apps/upload/tasks.go index 0bc567d5..5e448b4d 100644 --- a/internal/apps/upload/tasks.go +++ b/internal/apps/upload/tasks.go @@ -92,6 +92,6 @@ func (h *CleanupUnusedUploadsHandler) Execute(ctx context.Context, payload []byt } msg := fmt.Sprintf("共处理 %d 个文件,成功删除 %d 个", totalProcessed, totalDeleted) - task.AppendLog(ctx, msg) + task.AppendLog(ctx, "%s", msg) return &task.TaskResult{Message: msg}, nil } diff --git a/internal/config/model.go b/internal/config/model.go index 86d7a856..6b124396 100644 --- a/internal/config/model.go +++ b/internal/config/model.go @@ -20,7 +20,6 @@ import "time" type configModel struct { App appConfig `mapstructure:"app"` - OAuth2 OAuth2Config `mapstructure:"oauth2"` Database databaseConfig `mapstructure:"database"` Redis redisConfig `mapstructure:"redis"` Log logConfig `mapstructure:"log"` @@ -54,17 +53,6 @@ func (a *appConfig) IsProduction() bool { return a.Env == "production" } -// OAuth2Config OAuth2/OIDC认证配置 -type OAuth2Config struct { - ClientID string `mapstructure:"client_id"` - ClientSecret string `mapstructure:"client_secret"` - RedirectURI string `mapstructure:"redirect_uri"` - Issuer string `mapstructure:"issuer"` - AuthorizationEndpoint string `mapstructure:"authorization_endpoint"` - TokenEndpoint string `mapstructure:"token_endpoint"` - UserEndpoint string `mapstructure:"user_endpoint"` -} - // databaseConfig 数据库配置 type databaseConfig struct { Enabled bool `mapstructure:"enabled"` diff --git a/internal/model/auth_source.go b/internal/model/auth_source.go index a706bd2e..0a22c2a4 100644 --- a/internal/model/auth_source.go +++ b/internal/model/auth_source.go @@ -251,7 +251,7 @@ func ListExternalAccountsByUserID(userID uint64) ([]ExternalAccountView, error) if account.AuthSourceID == 0 { name = "default" sourceType = "oidc" - label = "默认认证源" + label = "历史认证源" } else { source, err := GetAuthSourceByID(account.AuthSourceID) if err != nil {