diff --git a/openflare_server/controller/option.go b/openflare_server/controller/option.go index 77d8143d..13a6921b 100644 --- a/openflare_server/controller/option.go +++ b/openflare_server/controller/option.go @@ -22,6 +22,10 @@ var ( openRestyDurationTokenPattern = regexp.MustCompile(`^\d+[smhdwSMHDW]$`) ) +type optionBatchPayload struct { + Options []model.Option `json:"options"` +} + func validateRateLimitOption(key string, value string) error { maxDurationSeconds := int(common.RateLimitKeyExpirationDuration.Seconds()) @@ -193,6 +197,69 @@ func validateOpenRestyOption(key string, value string) error { } } +func buildOptionValidationState(options []model.Option) map[string]string { + common.OptionMapRWMutex.RLock() + state := make(map[string]string, len(common.OptionMap)+len(options)) + for key, value := range common.OptionMap { + state[key] = value + } + common.OptionMapRWMutex.RUnlock() + + for _, option := range options { + state[option.Key] = option.Value + } + return state +} + +func validateOptionWithState(option model.Option, state map[string]string) error { + switch option.Key { + case "GitHubOAuthEnabled": + if option.Value == "true" && strings.TrimSpace(state["GitHubClientId"]) == "" { + return fmt.Errorf("鏃犳硶鍚敤 GitHub OAuth锛岃鍏堝~鍏?GitHub Client ID 浠ュ強 GitHub Client Secret锛?") + } + case "WeChatAuthEnabled": + if option.Value == "true" && strings.TrimSpace(state["WeChatServerAddress"]) == "" { + return fmt.Errorf("鏃犳硶鍚敤寰俊鐧诲綍锛岃鍏堝~鍏ュ井淇$櫥褰曠浉鍏抽厤缃俊鎭紒") + } + case "TurnstileCheckEnabled": + if option.Value == "true" && strings.TrimSpace(state["TurnstileSiteKey"]) == "" { + return fmt.Errorf("鏃犳硶鍚敤 Turnstile 鏍¢獙锛岃鍏堝~鍏?Turnstile 鏍¢獙鐩稿叧閰嶇疆淇℃伅锛?") + } + } + + if err := validateRateLimitOption(option.Key, option.Value); err != nil { + return err + } + if err := validateOpenRestyOption(option.Key, option.Value); err != nil { + return err + } + if err := validateGeoIPOption(option.Key, option.Value); err != nil { + return err + } + if err := validateDatabaseCleanupOption(option.Key, option.Value); err != nil { + return err + } + return nil +} + +func updateOptions(options []model.Option) error { + if len(options) == 0 { + return fmt.Errorf("鏃犳晥鐨勫弬鏁?") + } + + state := buildOptionValidationState(options) + for _, option := range options { + if strings.TrimSpace(option.Key) == "" { + return fmt.Errorf("鏃犳晥鐨勫弬鏁?") + } + if err := validateOptionWithState(option, state); err != nil { + return err + } + } + + return model.UpdateOptions(options) +} + // GetOptions godoc // @Summary List editable options // @Tags Options @@ -307,3 +374,36 @@ func UpdateOption(c *gin.Context) { }) return } + +// UpdateOptionsBatch godoc +// @Summary Batch update options +// @Tags Options +// @Accept json +// @Produce json +// @Param payload body optionBatchPayload true "Batch option payload" +// @Success 200 {object} map[string]interface{} +// @Failure 400 {object} map[string]interface{} +// @Router /api/option/update-batch [post] +func UpdateOptionsBatch(c *gin.Context) { + var payload optionBatchPayload + if err := json.NewDecoder(c.Request.Body).Decode(&payload); err != nil || len(payload.Options) == 0 { + c.JSON(http.StatusBadRequest, gin.H{ + "success": false, + "message": "鏃犳晥鐨勫弬鏁?", + }) + return + } + + if err := updateOptions(payload.Options); err != nil { + c.JSON(http.StatusOK, gin.H{ + "success": false, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "success": true, + "message": "", + }) +} diff --git a/openflare_server/model/option.go b/openflare_server/model/option.go index 5306959f..84be27c7 100644 --- a/openflare_server/model/option.go +++ b/openflare_server/model/option.go @@ -6,6 +6,8 @@ import ( "strconv" "strings" "time" + + "gorm.io/gorm" ) type Option struct { @@ -110,19 +112,38 @@ func InitOptionMap() { } func UpdateOption(key string, value string) error { - // Save to database first - option := Option{ - Key: key, + return UpdateOptions([]Option{{ + Key: key, + Value: value, + }}) +} + +func UpdateOptions(options []Option) error { + if len(options) == 0 { + return nil + } + + if err := DB.Transaction(func(tx *gorm.DB) error { + for _, item := range options { + option := Option{ + Key: item.Key, + } + if err := tx.FirstOrCreate(&option, Option{Key: item.Key}).Error; err != nil { + return err + } + option.Value = item.Value + if err := tx.Save(&option).Error; err != nil { + return err + } + } + return nil + }); err != nil { + return err + } + + for _, item := range options { + updateOptionMap(item.Key, item.Value) } - // https://gorm.io/docs/update.html#Save-All-Fields - DB.FirstOrCreate(&option, Option{Key: key}) - option.Value = value - // Save is a combination function. - // If save value does not contain primary key, it will execute Create, - // otherwise it will execute Update (with all fields). - DB.Save(&option) - // Update OptionMap - updateOptionMap(key, value) return nil } diff --git a/openflare_server/router/api-router.go b/openflare_server/router/api-router.go index 9950c18c..1fd08aef 100644 --- a/openflare_server/router/api-router.go +++ b/openflare_server/router/api-router.go @@ -54,6 +54,7 @@ func SetApiRouter(router *gin.Engine) { { optionRoute.GET("/", controller.GetOptions) optionRoute.POST("/update", controller.UpdateOption) + optionRoute.POST("/update-batch", controller.UpdateOptionsBatch) optionRoute.POST("/geoip/lookup", controller.LookupGeoIP) optionRoute.POST("/database/cleanup", controller.CleanupDatabaseObservability) } diff --git a/openflare_server/router/api_phase2_test.go b/openflare_server/router/api_phase2_test.go index ae9f2225..967b5ca8 100644 --- a/openflare_server/router/api_phase2_test.go +++ b/openflare_server/router/api_phase2_test.go @@ -41,21 +41,25 @@ func TestPhase2RateLimitOptionsHotReload(t *testing.T) { loginCookie := loginAsRoot(t, engine) - performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update", map[string]any{ - "key": "GlobalApiRateLimitNum", - "value": "450", - }) - performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update", map[string]any{ - "key": "GlobalApiRateLimitDuration", - "value": "240", - }) - performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update", map[string]any{ - "key": "CriticalRateLimitNum", - "value": "150", - }) - performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update", map[string]any{ - "key": "CriticalRateLimitDuration", - "value": "900", + performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update-batch", map[string]any{ + "options": []map[string]any{ + { + "key": "GlobalApiRateLimitNum", + "value": "450", + }, + { + "key": "GlobalApiRateLimitDuration", + "value": "240", + }, + { + "key": "CriticalRateLimitNum", + "value": "150", + }, + { + "key": "CriticalRateLimitDuration", + "value": "900", + }, + }, }) if common.GlobalApiRateLimitNum != 450 { @@ -88,6 +92,115 @@ func TestPhase2RateLimitOptionsHotReload(t *testing.T) { } } +func TestPhase2BatchOptionUpdateIsAtomic(t *testing.T) { + gin.SetMode(gin.TestMode) + common.RedisEnabled = false + setupTestDB(t) + model.InitOptionMap() + + oldGlobalAPI := common.GlobalApiRateLimitNum + t.Cleanup(func() { + common.GlobalApiRateLimitNum = oldGlobalAPI + }) + + engine := gin.New() + engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret")))) + router.SetApiRouter(engine) + + loginCookie := loginAsRoot(t, engine) + + payload, err := json.Marshal(map[string]any{ + "options": []map[string]any{ + { + "key": "GlobalApiRateLimitNum", + "value": "451", + }, + { + "key": "CriticalRateLimitDuration", + "value": "1800", + }, + }, + }) + if err != nil { + t.Fatalf("failed to marshal batch payload: %v", err) + } + + req := httptest.NewRequest(http.MethodPost, "/api/option/update-batch", bytes.NewReader(payload)) + req.Header.Set("Content-Type", "application/json") + req.AddCookie(loginCookie) + + recorder := httptest.NewRecorder() + engine.ServeHTTP(recorder, req) + if recorder.Code != http.StatusOK { + t.Fatalf("unexpected status %d: %s", recorder.Code, recorder.Body.String()) + } + + var resp apiResponse + if err = json.Unmarshal(recorder.Body.Bytes(), &resp); err != nil { + t.Fatalf("failed to unmarshal response: %v", err) + } + if resp.Success { + t.Fatal("expected invalid batch update to fail") + } + + if common.GlobalApiRateLimitNum != oldGlobalAPI { + t.Fatalf("expected GlobalApiRateLimitNum to remain %d after failed batch, got %d", oldGlobalAPI, common.GlobalApiRateLimitNum) + } + + resp = performSessionJSONRequest(t, engine, loginCookie, http.MethodGet, "/api/option/", nil) + var options []model.Option + decodeResponseData(t, resp, &options) + + optionMap := make(map[string]string, len(options)) + for _, option := range options { + optionMap[option.Key] = option.Value + } + + if optionMap["GlobalApiRateLimitNum"] == "451" { + t.Fatal("expected failed batch update to avoid persisting partial values") + } +} + +func TestPhase2BatchOptionUpdateValidatesMergedState(t *testing.T) { + gin.SetMode(gin.TestMode) + common.RedisEnabled = false + setupTestDB(t) + model.InitOptionMap() + + oldGitHubClientID := common.GitHubClientId + oldGitHubOAuthEnabled := common.GitHubOAuthEnabled + t.Cleanup(func() { + common.GitHubClientId = oldGitHubClientID + common.GitHubOAuthEnabled = oldGitHubOAuthEnabled + }) + + engine := gin.New() + engine.Use(sessions.Sessions("session", cookie.NewStore([]byte("test-secret")))) + router.SetApiRouter(engine) + + loginCookie := loginAsRoot(t, engine) + + performSessionJSONRequest(t, engine, loginCookie, http.MethodPost, "/api/option/update-batch", map[string]any{ + "options": []map[string]any{ + { + "key": "GitHubClientId", + "value": "client-id-from-batch", + }, + { + "key": "GitHubOAuthEnabled", + "value": "true", + }, + }, + }) + + if common.GitHubClientId != "client-id-from-batch" { + t.Fatalf("expected GitHubClientId to be updated from batch, got %q", common.GitHubClientId) + } + if !common.GitHubOAuthEnabled { + t.Fatal("expected GitHubOAuthEnabled to be enabled by merged batch state") + } +} + func loginAsRoot(t *testing.T, engine http.Handler) *http.Cookie { t.Helper() payload, err := json.Marshal(map[string]any{ diff --git a/openflare_server/web/features/performance/components/performance-page.tsx b/openflare_server/web/features/performance/components/performance-page.tsx index 51626bbc..2abb7ee0 100644 --- a/openflare_server/web/features/performance/components/performance-page.tsx +++ b/openflare_server/web/features/performance/components/performance-page.tsx @@ -12,7 +12,7 @@ import {useAuth} from '@/components/providers/auth-provider'; import {AppCard} from '@/components/ui/app-card'; import {StatusBadge} from '@/components/ui/status-badge'; import {getConfigVersionPreview} from '@/features/config-versions/api/config-versions'; -import {getOptions, updateOption} from '@/features/settings/api/settings'; +import {getOptions, updateOptions} from '@/features/settings/api/settings'; import type {OptionItem} from '@/features/settings/types'; import { CodeBlock, @@ -277,9 +277,7 @@ export function PerformancePage() { entries: Array<[string, string]>, successMessage: string, ) => { - for (const [key, value] of entries) { - await updateOption(key, value); - } + await updateOptions(entries.map(([key, value]) => ({key, value}))); await Promise.all([ queryClient.invalidateQueries({queryKey: settingsQueryKey}), diff --git a/openflare_server/web/features/settings/api/settings.ts b/openflare_server/web/features/settings/api/settings.ts index d4f1f394..5ee35b12 100644 --- a/openflare_server/web/features/settings/api/settings.ts +++ b/openflare_server/web/features/settings/api/settings.ts @@ -5,6 +5,7 @@ import type { DatabaseCleanupPayload, DatabaseCleanupResult, GeoIPLookupResult, + OptionBatchPayload, OptionItem, SettingsProfile, UpdateSelfPayload, @@ -21,6 +22,13 @@ export function updateOption(key: string, value: string) { }); } +export function updateOptions(options: OptionBatchPayload['options']) { + return apiRequest('/option/update-batch', { + method: 'POST', + body: JSON.stringify({ options }), + }); +} + export function lookupGeoIP(provider: string, ip: string) { return apiRequest('/option/geoip/lookup', { method: 'POST', diff --git a/openflare_server/web/features/settings/components/settings-page.tsx b/openflare_server/web/features/settings/components/settings-page.tsx index 395c6d97..b86eb2c1 100644 --- a/openflare_server/web/features/settings/components/settings-page.tsx +++ b/openflare_server/web/features/settings/components/settings-page.tsx @@ -25,7 +25,7 @@ import { getSettingsProfile, lookupGeoIP, rotateBootstrapToken, - updateOption, + updateOptions, updateSelf, } from '@/features/settings/api/settings'; import type { @@ -558,9 +558,7 @@ export function SettingsPage() { entries: Array<[string, string]>, successMessage: string, ) => { - for (const [key, value] of entries) { - await updateOption(key, value); - } + await updateOptions(entries.map(([key, value]) => ({ key, value }))); await queryClient.invalidateQueries({ queryKey: settingsQueryKey }); await queryClient.invalidateQueries({ queryKey: ['public-status'] }); diff --git a/openflare_server/web/features/settings/types.ts b/openflare_server/web/features/settings/types.ts index 924bd478..e7bc791a 100644 --- a/openflare_server/web/features/settings/types.ts +++ b/openflare_server/web/features/settings/types.ts @@ -5,6 +5,10 @@ export interface OptionItem { value: string; } +export interface OptionBatchPayload { + options: OptionItem[]; +} + export interface BootstrapTokenPayload { discovery_token: string; }