diff --git a/frontend/components/common/settings/security.tsx b/frontend/components/common/settings/security.tsx index d06761f6..9d0a0285 100644 --- a/frontend/components/common/settings/security.tsx +++ b/frontend/components/common/settings/security.tsx @@ -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 (
@@ -541,7 +565,57 @@ export function SecurityMain() {
- + +
+ + +
+
+ +
+
+ 通用设置 + 配置系统的全局通用参数 +
+
+
+ +
+
+ + setServerAddress(e.target.value)} + placeholder="例如: https://example.com" + className="bg-card border-dashed text-xs" + /> +

+ 这里可以编辑更改服务器地址。默认不设定,允许从任意源(*)访问 API,此时存在跨域安全风险;如果手动设置服务器地址,CORS 允许源将更新为该地址,消除跨域安全隐患。 +

+
+
+ +
+
+
+
+
+
diff --git a/internal/apps/admin/system_config/routers_test.go b/internal/apps/admin/system_config/routers_test.go index 1cf40912..b247d980 100644 --- a/internal/apps/admin/system_config/routers_test.go +++ b/internal/apps/admin/system_config/routers_test.go @@ -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)) } }) diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index 4eb2e997..17a09351 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -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() } diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index ac2a121f..f3e1b0ca 100644 --- a/internal/apps/oauth/oauth_test.go +++ b/internal/apps/oauth/oauth_test.go @@ -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) - } -} diff --git a/internal/apps/oauth/risk.go b/internal/apps/oauth/risk.go deleted file mode 100644 index 5358a151..00000000 --- a/internal/apps/oauth/risk.go +++ /dev/null @@ -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 -} diff --git a/internal/config/model.go b/internal/config/model.go index 6b124396..17aa164b 100644 --- a/internal/config/model.go +++ b/internal/config/model.go @@ -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"` diff --git a/internal/db/migrator/migrator.go b/internal/db/migrator/migrator.go index 38de74bf..11e0f23c 100644 --- a/internal/db/migrator/migrator.go +++ b/internal/db/migrator/migrator.go @@ -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", diff --git a/internal/model/system_configs.go b/internal/model/system_configs.go index 40803e21..d961e911 100644 --- a/internal/model/system_configs.go +++ b/internal/model/system_configs.go @@ -45,6 +45,7 @@ const ( ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度) ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒) ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" // 人机验证兑换凭证有效时间(秒) + ConfigKeyServerAddress = "server_address" // 服务器地址 ) const ( diff --git a/internal/router/middlewares.go b/internal/router/middlewares.go index 92c46eb8..f0025e04 100644 --- a/internal/router/middlewares.go +++ b/internal/router/middlewares.go @@ -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() + } +} diff --git a/internal/router/router.go b/internal/router/router.go index 8ceb8876..98f3c35e 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -63,6 +63,7 @@ func Serve() { // 初始化路由 r := gin.New() r.Use(gin.Recovery()) + r.Use(corsMiddleware()) cfg := config.Config.Redis addrs := cfg.Addrs diff --git a/internal/testhelper/test_helper.go b/internal/testhelper/test_helper.go index 2ba19026..66a6a6eb 100644 --- a/internal/testhelper/test_helper.go +++ b/internal/testhelper/test_helper.go @@ -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 {