去除默认OIDC

This commit is contained in:
ryan
2026-06-07 23:55:19 +08:00
parent 53a4c92176
commit 4b72419a96
18 changed files with 356 additions and 437 deletions
+1
View File
@@ -44,3 +44,4 @@ uploads/*
s3_cache
/frontend/.next/
/data/
+1
View File
@@ -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 自动生成文档(不要手动编辑)
+8 -3
View File
@@ -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 <host> -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
+8 -3
View File
@@ -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` |
## 🔧 开发指南
+1 -12
View File
@@ -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: "<OAUTH2_CLIENT_ID>"
client_secret: "<OAUTH2_CLIENT_SECRET>"
redirect_uri: "<OAUTH2_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
+61
View File
@@ -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
+2 -2
View File
@@ -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"
}
+2 -2
View File
@@ -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"
}
+2 -2
View File
@@ -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
+21 -24
View File
@@ -237,31 +237,28 @@ export function LoginForm() {
<Separator />
<div className="space-y-3">
<div className="flex items-center gap-2 text-sm font-medium text-foreground">
<ShieldCheck className="size-4" />
第三方认证源
</div>
<div className="grid gap-2">
{authSources.length > 0 ? (
authSources.map((source) => (
<Button
key={source.id}
type="button"
variant="outline"
className="justify-start"
onClick={() => void handleOAuthLogin(source.name)}
>
{source.display_name || source.name} 登录
</Button>
))
) : (
<div className="rounded-lg border border-dashed border-border/60 px-3 py-4 text-sm text-muted-foreground">
暂无可用认证源
{authSources.length > 0 ? (
authSources.map((source) => (
<div className="space-y-3">
<div className="flex items-center gap-2 text-sm font-medium text-foreground">
<ShieldCheck className="size-4" />
第三方认证源
</div>
)}
</div>
</div>
<div className="grid gap-2">
<Button
key={source.id}
type="button"
variant="outline"
className="justify-start"
onClick={() => void handleOAuthLogin(source.name)}
>
{source.display_name || source.name} 登录
</Button>
</div>
</div>
))
) : null}
</CardContent>
</Card>
)
@@ -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 (
<div className="flex items-center justify-between gap-4 border-b border-dashed py-2 last:border-b-0">
<span className="text-xs text-muted-foreground">{label}</span>
<span className="text-right text-xs font-medium text-foreground break-all">{value || "-"}</span>
</div>
)
}
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<AuthSource | null>(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() {
<TabsContent value="operation" />
<TabsContent value="system" />
<TabsContent value="other" />
<TabsContent value="info" />
<TabsContent value="info" className="pt-4">
<div className="grid grid-cols-1 gap-4 lg:grid-cols-2">
<Card className="border border-dashed shadow-sm">
<CardHeader className="border-b border-dashed pb-4">
<div className="flex items-center gap-2">
<div className="p-1.5 rounded-lg bg-muted text-muted-foreground">
<Info className="size-4" />
</div>
<div>
<CardTitle className="text-base font-semibold">应用信息</CardTitle>
<CardDescription className="text-xs">当前前端应用的版本与构建信息</CardDescription>
</div>
</div>
</CardHeader>
<CardContent className="pt-4">
<InfoRow label="应用名称" value={packageJson.name} />
<InfoRow label="版本号" value={packageJson.version} />
<InfoRow label="构建时间" value={packageJson.buildDate} />
<InfoRow label="Next.js" value={packageJson.dependencies.next} />
<InfoRow label="React" value={packageJson.dependencies.react} />
</CardContent>
</Card>
<Card className="border border-dashed shadow-sm">
<CardHeader className="border-b border-dashed pb-4">
<div className="flex items-center gap-2">
<div className="p-1.5 rounded-lg bg-muted text-muted-foreground">
<Server className="size-4" />
</div>
<div>
<CardTitle className="text-base font-semibold">服务连接</CardTitle>
<CardDescription className="text-xs">前端 API 客户端的基础连接参数</CardDescription>
</div>
</div>
</CardHeader>
<CardContent className="pt-4">
<InfoRow label="API Base URL" value={apiConfig.baseURL || "同源"} />
<InfoRow label="请求超时" value={`${apiConfig.timeout}ms`} />
<InfoRow label="携带凭证" value={apiConfig.withCredentials ? "是" : "否"} />
<InfoRow label="系统配置项" value={`${systemConfigsQuery.data?.length ?? 0} 项`} />
<InfoRow label="认证源数量" value={`${authSourcesQuery.data?.length ?? 0} 个`} />
</CardContent>
</Card>
<Card className="border border-dashed shadow-sm">
<CardHeader className="border-b border-dashed pb-4">
<div className="flex items-center gap-2">
<div className="p-1.5 rounded-lg bg-muted text-muted-foreground">
<Monitor className="size-4" />
</div>
<div>
<CardTitle className="text-base font-semibold">运行环境</CardTitle>
<CardDescription className="text-xs">当前浏览器会话的本地运行信息</CardDescription>
</div>
</div>
</CardHeader>
<CardContent className="pt-4">
<InfoRow label="语言" value={runtimeInfo.language} />
<InfoRow label="平台" value={runtimeInfo.platform} />
<InfoRow label="时区" value={runtimeInfo.timezone} />
<InfoRow label="视口" value={runtimeInfo.viewport} />
<InfoRow label="User Agent" value={runtimeInfo.userAgent} />
</CardContent>
</Card>
<Card className="border border-dashed shadow-sm">
<CardHeader className="border-b border-dashed pb-4">
<div className="flex items-center gap-2">
<div className="p-1.5 rounded-lg bg-muted text-muted-foreground">
<CalendarClock className="size-4" />
</div>
<div>
<CardTitle className="text-base font-semibold">安全配置概览</CardTitle>
<CardDescription className="text-xs">当前系统登录与注册开关状态</CardDescription>
</div>
</div>
</CardHeader>
<CardContent className="pt-4">
<InfoRow label="密码登录" value={formatBooleanConfig(configs.password_login_enabled)} />
<InfoRow label="开放注册" value={formatBooleanConfig(configs.registration_enabled)} />
<InfoRow label="密码注册" value={formatBooleanConfig(configs.password_register_enabled)} />
<InfoRow label="OIDC 登录" value={formatBooleanConfig(configs.oidc_login_enabled)} />
<InfoRow label="配置加载状态" value={systemConfigsQuery.isFetching ? "刷新中" : "已加载"} />
</CardContent>
</Card>
</div>
</TabsContent>
</Tabs>
<AuthSourceModal
-73
View File
@@ -1,73 +0,0 @@
/*
Copyright 2025 linux.do
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package oauth
import (
"context"
"log"
"strings"
"github.com/coreos/go-oidc/v3/oidc"
"github.com/linux-do/credit/internal/config"
"golang.org/x/oauth2"
)
var (
oauthConf *oauth2.Config
oidcVerifier *oidc.IDTokenVerifier
)
func init() {
cfg := config.Config.OAuth2
if cfg.Issuer != "" {
ctx := context.Background()
// Clean the issuer URL (trim /.well-known/openid-configuration if configured by mistake)
issuer := strings.TrimSuffix(strings.TrimSpace(cfg.Issuer), "/")
issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration")
issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server")
provider, err := oidc.NewProvider(ctx, issuer)
if err != nil {
log.Printf("[OAuth] 初始化 OIDC Provider 失败: %v,将仅使用 OAuth2", err)
} else {
oidcVerifier = provider.Verifier(&oidc.Config{
ClientID: cfg.ClientID,
})
log.Printf("[OAuth] OIDC Provider 初始化成功: %s", issuer)
}
}
// 初始化 OAuth2 配置
scopes := []string{"profile", "email"}
if oidcVerifier != nil {
// 启用 OIDC 时添加 openid scope
scopes = append([]string{oidc.ScopeOpenID}, scopes...)
}
oauthConf = &oauth2.Config{
ClientID: cfg.ClientID,
ClientSecret: cfg.ClientSecret,
RedirectURL: cfg.RedirectURI,
Scopes: scopes,
Endpoint: oauth2.Endpoint{
AuthURL: cfg.AuthorizationEndpoint,
TokenURL: cfg.TokenEndpoint,
AuthStyle: oauth2.AuthStyleAutoDetect,
},
}
}
-152
View File
@@ -1,152 +0,0 @@
/*
Copyright 2025 linux.do
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
*/
package oauth
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"github.com/linux-do/credit/internal/common"
"github.com/linux-do/credit/internal/config"
"github.com/linux-do/credit/internal/db"
"github.com/linux-do/credit/internal/model"
"github.com/linux-do/credit/internal/otel_trace"
"go.opentelemetry.io/otel/codes"
"gorm.io/gorm"
)
func GetUserIDFromSession(s sessions.Session) uint64 {
userID, ok := s.Get(UserIDKey).(uint64)
if !ok {
return 0
}
return userID
}
func GetUserIDFromContext(c *gin.Context) uint64 {
session := sessions.Default(c)
return GetUserIDFromSession(session)
}
// doOAuth 执行 OAuth2/OIDC 流程
func doOAuth(ctx context.Context, code string, nonce string) (*model.User, error) {
ctx, span := otel_trace.Start(ctx, "OAuth")
defer span.End()
var userInfo model.OAuthUserInfo
// 使用授权码换取 Token
token, err := oauthConf.Exchange(ctx, code)
if err != nil {
span.SetStatus(codes.Error, err.Error())
return nil, err
}
if oidcVerifier != nil {
if rawIDToken, ok := token.Extra("id_token").(string); ok {
idToken, verifyErr := oidcVerifier.Verify(ctx, rawIDToken)
if verifyErr != nil {
err := fmt.Errorf("%s: %w", IDTokenVerifyFailed, verifyErr)
span.SetStatus(codes.Error, err.Error())
return nil, err
}
if nonce != "" && idToken.Nonce != nonce {
span.SetStatus(codes.Error, NonceMismatch)
return nil, errors.New(NonceMismatch)
}
if claimsErr := idToken.Claims(&userInfo); claimsErr != nil {
span.SetStatus(codes.Error, claimsErr.Error())
return nil, claimsErr
}
}
}
if userInfo.GetID() == 0 {
client := oauthConf.Client(ctx, token)
resp, httpErr := client.Get(config.Config.OAuth2.UserEndpoint)
if httpErr != nil {
span.SetStatus(codes.Error, httpErr.Error())
return nil, httpErr
}
defer resp.Body.Close()
responseData, readErr := io.ReadAll(resp.Body)
if readErr != nil {
span.SetStatus(codes.Error, readErr.Error())
return nil, readErr
}
if unmarshalErr := json.Unmarshal(responseData, &userInfo); unmarshalErr != nil {
span.SetStatus(codes.Error, unmarshalErr.Error())
return nil, unmarshalErr
}
}
if !userInfo.Active {
err := errors.New(common.BannedAccount)
span.SetStatus(codes.Error, err.Error())
return nil, err
}
var user model.User
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var holder model.User
if conflictErr := tx.Where("username = ? AND id != ?", userInfo.Username, userInfo.GetID()).First(&holder).Error; conflictErr == nil {
// 存在冲突 -> 将占用者改名并注销
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
}
+90 -49
View File
@@ -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)),
+24 -94
View File
@@ -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
}
+1 -1
View File
@@ -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
}
-12
View File
@@ -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"`
+1 -1
View File
@@ -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 {