mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 14:06:36 +08:00
移除 risk
This commit is contained in:
@@ -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))
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -45,6 +45,7 @@ const (
|
||||
ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" // 人机验证 PoW 难度(目标前缀长度)
|
||||
ConfigKeyCapChallengeTTL = "cap_challenge_ttl_seconds" // 人机验证难题有效时间(秒)
|
||||
ConfigKeyCapTokenTTL = "cap_token_ttl_seconds" // 人机验证兑换凭证有效时间(秒)
|
||||
ConfigKeyServerAddress = "server_address" // 服务器地址
|
||||
)
|
||||
|
||||
const (
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -63,6 +63,7 @@ func Serve() {
|
||||
// 初始化路由
|
||||
r := gin.New()
|
||||
r.Use(gin.Recovery())
|
||||
r.Use(corsMiddleware())
|
||||
|
||||
cfg := config.Config.Redis
|
||||
addrs := cfg.Addrs
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user