移除 risk

This commit is contained in:
ryan
2026-06-08 13:55:00 +08:00
parent 47c9bc53fa
commit 62bd5d09d4
11 changed files with 129 additions and 381 deletions
@@ -93,6 +93,7 @@ export function SecurityMain() {
const [capTTL, setCapTTL] = useState("")
const [capTokenTTL, setCapTokenTTL] = useState("")
const [capAutoSolve, setCapAutoSolve] = useState(true)
const [serverAddress, setServerAddress] = useState("")
const systemConfigsQuery = useQuery({
queryKey: ["admin", "system-configs"],
@@ -126,6 +127,7 @@ export function SecurityMain() {
setCapTTL(cfgMap["cap_challenge_ttl_seconds"]?.value || "600")
setCapTokenTTL(cfgMap["cap_token_ttl_seconds"]?.value || "1200")
setCapAutoSolve(cfgMap["cap_auto_solve"]?.value !== "false")
setServerAddress(cfgMap["server_address"]?.value || "")
}
}, [systemConfigsQuery.data])
@@ -214,6 +216,28 @@ export function SecurityMain() {
saveCapMutation.mutate()
}
const saveSystemMutation = useMutation({
mutationFn: async () => {
const currentCfg = configs["server_address"]
await AdminService.updateSystemConfig("server_address", {
value: serverAddress,
description: currentCfg?.description || "服务器地址",
})
},
onSuccess: async () => {
await queryClient.invalidateQueries({ queryKey: ["admin", "system-configs"] })
toast.success("通用配置已成功保存")
},
onError: (error: Error) => {
toast.error(error.message || "保存配置失败")
},
})
const handleSystemSave = (e: React.FormEvent) => {
e.preventDefault()
saveSystemMutation.mutate()
}
if (loading || !user || !user.is_admin) {
return (
<div className="flex items-center justify-center min-h-[400px]">
@@ -541,7 +565,57 @@ export function SecurityMain() {
</div>
</TabsContent>
<TabsContent value="operation" />
<TabsContent value="system" />
<TabsContent value="system" className="pt-4">
<div className="space-y-6">
<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-indigo-500/10 text-indigo-500">
<Server className="size-4" />
</div>
<div>
<CardTitle className="text-base font-semibold">通用设置</CardTitle>
<CardDescription className="text-xs">配置系统的全局通用参数</CardDescription>
</div>
</div>
</CardHeader>
<CardContent className="pt-6">
<form onSubmit={handleSystemSave} className="space-y-6">
<div className="space-y-1.5">
<Label htmlFor="server_address" className="text-xs font-semibold">服务器地址</Label>
<Input
id="server_address"
type="text"
value={serverAddress}
onChange={(e) => setServerAddress(e.target.value)}
placeholder="例如: https://example.com"
className="bg-card border-dashed text-xs"
/>
<p className="text-[10px] text-muted-foreground leading-normal">
这里可以编辑更改服务器地址。默认不设定,允许从任意源(*)访问 API,此时存在跨域安全风险;如果手动设置服务器地址,CORS 允许源将更新为该地址,消除跨域安全隐患。
</p>
</div>
<div className="flex justify-end pt-4 border-t border-dashed">
<Button
type="submit"
size="sm"
disabled={saveSystemMutation.isPending}
>
{saveSystemMutation.isPending ? (
<>
<Loader2 className="mr-1.5 size-3.5 animate-spin" />
保存中...
</>
) : (
"保存配置"
)}
</Button>
</div>
</form>
</CardContent>
</Card>
</div>
</TabsContent>
<TabsContent value="status" className="pt-4">
<SystemStatusManager />
</TabsContent>
@@ -143,9 +143,9 @@ func TestListSystemConfigs(t *testing.T) {
var configs []model.SystemConfig
json.Unmarshal(dataBytes, &configs)
// Defaults seed 7 configurations
if len(configs) != 7 {
t.Errorf("expected 7 default configs, got %d", len(configs))
// Defaults seed 15 configurations
if len(configs) != 15 {
t.Errorf("expected 15 default configs, got %d", len(configs))
}
})
-6
View File
@@ -92,12 +92,6 @@ func LoginRequired() gin.HandlerFunc {
// set user info
util.SetToContext(c, UserObjKey, &user)
if risk, ok := checkOpenAPIUserRisk(ctx, user.ID); ok {
if blocked := applyOpenAPIUserRisk(c, risk); blocked {
return
}
}
// next
c.Next()
}
-99
View File
@@ -333,10 +333,6 @@ func initializeTestConfig() {
config.Config.App.SessionSecret = "test_session_secret"
config.Config.App.APIPrefix = "/api"
config.Config.App.FrontendURL = "http://localhost:3000"
config.Config.OpenAPIRisk.Enabled = false
config.Config.OpenAPIRisk.BaseURL = ""
config.Config.OpenAPIRisk.BlockRiskLevels = []string{}
config.Config.OpenAPIRisk.PromptRiskLevels = []string{}
}
// -----------------------------------------------------------------------------
@@ -986,98 +982,3 @@ func TestExternalAccountsListAndDelete(t *testing.T) {
t.Error("binding record was not deleted from DB")
}
}
func TestLoginRequiredAndRiskChecks(t *testing.T) {
initializeTestConfig()
dbConn := setupTestDB(t)
mockRedis := newMockRedisClient()
// Create user
dbConn.Create(&model.User{
ID: 1122,
Username: "risk_tester",
IsActive: true,
})
// Enable risk checks
config.Config.OpenAPIRisk.Enabled = true
config.Config.OpenAPIRisk.BaseURL = "https://risk.test"
config.Config.OpenAPIRisk.BlockRiskLevels = []string{"BLOCK"}
config.Config.OpenAPIRisk.PromptRiskLevels = []string{"WARN"}
// Mock Risk HTTP Endpoint (Returns BLOCK)
httpMock := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/api/open/v1/risk/users/1122") {
body := `{"risky":true,"risk_level":"BLOCK","risks":[{"label":"IP_ABUSE","value":"high","desc":"Blocked IP address"}]}`
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
return nil, fmt.Errorf("unexpected request: %s", req.URL)
},
},
}
util.SetHTTPClient(httpMock)
router := setupTestRouter(dbConn, mockRedis, httpMock)
router.GET("/test-helper/login-1122", func(c *gin.Context) {
session := sessions.Default(c)
session.Set(UserIDKey, uint64(1122))
session.Save()
c.String(200, "ok")
})
wLogin := performRequest(router, http.MethodGet, "/test-helper/login-1122", nil, nil, nil)
var activeCookie *http.Cookie
for _, cookie := range wLogin.Result().Cookies() {
if cookie.Name == config.Config.App.SessionCookieName {
activeCookie = cookie
break
}
}
// Trigger request - UserInfo requires LoginRequired middleware, which will check risk
wBlock := performRequest(router, http.MethodGet, "/api/v1/oauth/user-info", nil, nil, []*http.Cookie{activeCookie})
if wBlock.Code != http.StatusForbidden {
t.Fatalf("expected 403 Forbidden when user is risk-blocked, got %d, body: %s", wBlock.Code, wBlock.Body.String())
}
if !strings.Contains(wBlock.Body.String(), "RISK_BLOCKED") {
t.Errorf("response body should contain RISK_BLOCKED: %s", wBlock.Body.String())
}
// Clear risk cache in redis
mockRedis.store = make(map[string]string)
// Mock Risk HTTP Endpoint (Returns WARN)
httpMock2 := &http.Client{
Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) {
if req.Method == http.MethodGet && strings.Contains(req.URL.String(), "/api/open/v1/risk/users/1122") {
body := `{"risky":true,"risk_level":"WARN","risks":[{"label":"VPN_DETECTED","value":"yes","desc":"Warning VPN"}]}`
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)),
}, nil
}
return nil, fmt.Errorf("unexpected request")
},
},
}
util.SetHTTPClient(httpMock2)
router2 := setupTestRouter(dbConn, mockRedis, httpMock2)
wWarn := performRequest(router2, http.MethodGet, "/api/v1/oauth/user-info", nil, nil, []*http.Cookie{activeCookie})
if wWarn.Code != http.StatusOK {
t.Fatalf("expected 200 OK when user is only warned, got %d, body: %s", wWarn.Code, wWarn.Body.String())
}
// Check if risk headers are present
hLevel := wWarn.Header().Get("X-Credit-Risk-Level")
if hLevel != "WARN" {
t.Errorf("expected X-Credit-Risk-Level WARN, got %s", hLevel)
}
}
-251
View File
@@ -1,251 +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/base64"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/linux-do/credit/internal/config"
"github.com/linux-do/credit/internal/db"
"github.com/linux-do/credit/internal/logger"
"github.com/linux-do/credit/internal/util"
"github.com/redis/go-redis/v9"
)
const (
openAPIRiskCacheKeyFormat = "openapi_risk:user:%d"
minOpenAPIRiskCacheTTL = time.Hour
riskLevelHeader = "X-Credit-Risk-Level"
riskLabelsHeader = "X-Credit-Risk-Labels"
riskItemsHeader = "X-Credit-Risks"
exposeHeader = "Access-Control-Expose-Headers"
riskBlockedCode = "RISK_BLOCKED"
riskBlockedMsg = "账号存在风险"
)
type openAPIUserRiskItem struct {
Label string `json:"label"`
Value string `json:"value"`
Desc string `json:"desc"`
}
type openAPIUserRiskResponse struct {
Risky bool `json:"risky"`
RiskLevel string `json:"risk_level"`
Risks []openAPIUserRiskItem `json:"risks"`
}
type riskBlockDetails struct {
RiskLevel string `json:"risk_level"`
RiskLabels []string `json:"risk_labels"`
Risks []openAPIUserRiskItem `json:"risks"`
}
func checkOpenAPIUserRisk(ctx context.Context, userID uint64) (*openAPIUserRiskResponse, bool) {
cfg := config.Config.OpenAPIRisk
if !cfg.Enabled || strings.TrimSpace(cfg.BaseURL) == "" {
return nil, false
}
if db.Redis == nil {
logger.ErrorF(ctx, "[OpenAPIRisk] redis is not initialized, skip risk check")
return nil, false
}
cacheKey := fmt.Sprintf(openAPIRiskCacheKeyFormat, userID)
var cached openAPIUserRiskResponse
if err := db.GetJSON(ctx, cacheKey, &cached); err == nil {
return &cached, true
} else if err != nil && !errors.Is(err, redis.Nil) {
logger.ErrorF(ctx, "[OpenAPIRisk] read cache failed, skip risk check: %v", err)
return nil, false
}
risk, err := fetchOpenAPIUserRisk(ctx, userID)
if err != nil {
logger.ErrorF(ctx, "[OpenAPIRisk] fetch user risk failed, skip risk check: %v", err)
return nil, false
}
if err := db.SetJSON(ctx, cacheKey, risk, openAPIRiskCacheTTL()); err != nil {
logger.ErrorF(ctx, "[OpenAPIRisk] write cache failed, skip risk check: %v", err)
return nil, false
}
return risk, true
}
func fetchOpenAPIUserRisk(ctx context.Context, userID uint64) (*openAPIUserRiskResponse, error) {
cfg := config.Config.OpenAPIRisk
endpoint := fmt.Sprintf(
"%s/api/open/v1/risk/users/%d",
strings.TrimRight(cfg.BaseURL, "/"),
userID,
)
headers := map[string]string{
"Accept": "application/json",
}
if cfg.Username != "" || cfg.Password != "" {
token := base64.StdEncoding.EncodeToString([]byte(cfg.Username + ":" + cfg.Password))
headers["Authorization"] = "Basic " + token
}
resp, err := util.Request(ctx, http.MethodGet, endpoint, nil, headers, nil)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("unexpected status code: %d", resp.StatusCode)
}
var risk openAPIUserRiskResponse
if err := json.NewDecoder(resp.Body).Decode(&risk); err != nil {
return nil, fmt.Errorf("decode response failed: %w", err)
}
return &risk, nil
}
func openAPIRiskCacheTTL() time.Duration {
ttl := time.Duration(config.Config.OpenAPIRisk.CacheTTLSeconds) * time.Second
if ttl < minOpenAPIRiskCacheTTL {
return minOpenAPIRiskCacheTTL
}
return ttl
}
func applyOpenAPIUserRisk(c *gin.Context, risk *openAPIUserRiskResponse) bool {
if risk == nil || !risk.Risky {
return false
}
labels := riskLabels(risk)
items := riskItems(risk)
cfg := config.Config.OpenAPIRisk
if containsString(cfg.BlockRiskLevels, risk.RiskLevel) {
setRiskHeaders(c, risk.RiskLevel, labels, items)
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{
"error_code": riskBlockedCode,
"error_msg": riskBlockedMsg,
"details": riskBlockDetails{
RiskLevel: risk.RiskLevel,
RiskLabels: labels,
Risks: items,
},
})
return true
}
if containsString(cfg.PromptRiskLevels, risk.RiskLevel) {
setRiskHeaders(c, risk.RiskLevel, labels, items)
}
return false
}
func setRiskHeaders(c *gin.Context, riskLevel string, labels []string, items []openAPIUserRiskItem) {
labelsJSON, err := json.Marshal(labels)
if err != nil {
logger.ErrorF(c.Request.Context(), "[OpenAPIRisk] marshal risk labels failed: %v", err)
return
}
itemsJSON, err := json.Marshal(items)
if err != nil {
logger.ErrorF(c.Request.Context(), "[OpenAPIRisk] marshal risk items failed: %v", err)
return
}
c.Header(riskLevelHeader, riskLevel)
c.Header(riskLabelsHeader, base64.StdEncoding.EncodeToString(labelsJSON))
c.Header(riskItemsHeader, base64.StdEncoding.EncodeToString(itemsJSON))
appendExposeHeaders(c, riskLevelHeader, riskLabelsHeader, riskItemsHeader)
}
func appendExposeHeaders(c *gin.Context, names ...string) {
existing := c.Writer.Header().Get(exposeHeader)
exposed := make([]string, 0, len(names)+1)
if existing != "" {
exposed = append(exposed, strings.Split(existing, ",")...)
}
exposed = append(exposed, names...)
seen := make(map[string]struct{}, len(exposed))
normalized := make([]string, 0, len(exposed))
for _, header := range exposed {
header = strings.TrimSpace(header)
if header == "" {
continue
}
key := strings.ToLower(header)
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
normalized = append(normalized, header)
}
c.Header(exposeHeader, strings.Join(normalized, ", "))
}
func riskLabels(risk *openAPIUserRiskResponse) []string {
labels := make([]string, 0, len(risk.Risks))
for _, item := range risk.Risks {
label := strings.TrimSpace(item.Label)
if label == "" {
continue
}
labels = append(labels, label)
}
return labels
}
func riskItems(risk *openAPIUserRiskResponse) []openAPIUserRiskItem {
items := make([]openAPIUserRiskItem, 0, len(risk.Risks))
for _, item := range risk.Risks {
item.Label = strings.TrimSpace(item.Label)
item.Value = strings.TrimSpace(item.Value)
item.Desc = strings.TrimSpace(item.Desc)
if item.Label == "" {
continue
}
items = append(items, item)
}
return items
}
func containsString(values []string, target string) bool {
target = strings.TrimSpace(target)
for _, value := range values {
if strings.TrimSpace(value) == target {
return true
}
}
return false
}
+9 -21
View File
@@ -19,16 +19,15 @@ package config
import "time"
type configModel struct {
App appConfig `mapstructure:"app"`
Database databaseConfig `mapstructure:"database"`
Redis redisConfig `mapstructure:"redis"`
Log logConfig `mapstructure:"log"`
Scheduler schedulerConfig `mapstructure:"scheduler"`
Worker workerConfig `mapstructure:"worker"`
ClickHouse clickHouseConfig `mapstructure:"clickhouse"`
OpenAPIRisk openAPIRiskConfig `mapstructure:"openapi_risk"`
Otel otelConfig `mapstructure:"otel"`
S3 s3Config `mapstructure:"s3"`
App appConfig `mapstructure:"app"`
Database databaseConfig `mapstructure:"database"`
Redis redisConfig `mapstructure:"redis"`
Log logConfig `mapstructure:"log"`
Scheduler schedulerConfig `mapstructure:"scheduler"`
Worker workerConfig `mapstructure:"worker"`
ClickHouse clickHouseConfig `mapstructure:"clickhouse"`
Otel otelConfig `mapstructure:"otel"`
S3 s3Config `mapstructure:"s3"`
}
// appConfig 应用基本配置
@@ -149,17 +148,6 @@ type QueueConfig struct {
Priority int `mapstructure:"priority"`
}
// openAPIRiskConfig OpenAPI 用户风险配置
type openAPIRiskConfig struct {
Enabled bool `mapstructure:"enabled"`
BaseURL string `mapstructure:"base_url"`
Username string `mapstructure:"username"`
Password string `mapstructure:"password" json:"-"`
CacheTTLSeconds int `mapstructure:"cache_ttl_seconds"`
PromptRiskLevels []string `mapstructure:"prompt_risk_levels"`
BlockRiskLevels []string `mapstructure:"block_risk_levels"`
}
// otelConfig OpenTelemetry 配置
type otelConfig struct {
SamplingRate float64 `mapstructure:"sampling_rate"`
+7
View File
@@ -89,6 +89,7 @@ func initSystemConfigs() {
ensureConfigKeyExists(model.ConfigKeyCapChallengeDifficulty, "4", "system", "人机验证 PoW 难度(目标前缀长度)")
ensureConfigKeyExists(model.ConfigKeyCapChallengeTTL, "600", "system", "人机验证难题有效时间(秒)")
ensureConfigKeyExists(model.ConfigKeyCapTokenTTL, "1200", "system", "人机验证兑换凭证有效时间(秒)")
ensureConfigKeyExists(model.ConfigKeyServerAddress, "", "system", "服务器地址(用于跨域源控制,不设定则允许任意源)")
return
}
@@ -135,6 +136,12 @@ func initSystemConfigs() {
Type: "system",
Description: "人机验证兑换凭证有效时间(秒)",
},
{
Key: model.ConfigKeyServerAddress,
Value: "",
Type: "system",
Description: "服务器地址(用于跨域源控制,不设定则允许任意源)",
},
{
Key: model.ConfigKeyUploadAllowedExtensions,
Value: "jpg,png,webp",
+1
View File
@@ -45,6 +45,7 @@ const (
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" // 人机验证兑换凭证有效时间(秒)
ConfigKeyServerAddress = "server_address" // 服务器地址
)
const (
+27
View File
@@ -17,12 +17,14 @@ limitations under the License.
package router
import (
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/linux-do/credit/internal/config"
"github.com/linux-do/credit/internal/logger"
"github.com/linux-do/credit/internal/model"
"github.com/linux-do/credit/internal/otel_trace"
"go.opentelemetry.io/otel/codes"
"go.opentelemetry.io/otel/trace"
@@ -76,3 +78,28 @@ func loggerMiddleware() gin.HandlerFunc {
}
}
}
func corsMiddleware() gin.HandlerFunc {
return func(c *gin.Context) {
origin := c.Request.Header.Get("Origin")
if origin != "" {
var sc model.SystemConfig
// Fetch from system config. We use request context which supports trace
if err := sc.GetByKey(c.Request.Context(), model.ConfigKeyServerAddress); err == nil && sc.Value != "" {
c.Writer.Header().Set("Access-Control-Allow-Origin", sc.Value)
} else {
c.Writer.Header().Set("Access-Control-Allow-Origin", origin)
}
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With, X-Access-Token, X-Cap-Token")
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE, PATCH")
}
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(http.StatusNoContent)
return
}
c.Next()
}
}
+1
View File
@@ -63,6 +63,7 @@ func Serve() {
// 初始化路由
r := gin.New()
r.Use(gin.Recovery())
r.Use(corsMiddleware())
cfg := config.Config.Redis
addrs := cfg.Addrs
+6
View File
@@ -176,6 +176,12 @@ func seedDefaultConfigs(t *testing.T, tx *gorm.DB) {
Type: "system",
Description: "人机验证兑换凭证有效时间(秒)",
},
{
Key: model.ConfigKeyServerAddress,
Value: "",
Type: "system",
Description: "服务器地址(用于跨域源控制,不设定则允许任意源)",
},
}
if err := tx.Create(&defaultConfigs).Error; err != nil {