diff --git a/controller/settings.go b/controller/settings.go index 56c2319..9fcee12 100644 --- a/controller/settings.go +++ b/controller/settings.go @@ -6,6 +6,7 @@ import ( "log" "net/http" "strconv" + "strings" "time" "github.com/volcano-tts/tts-api/middleware" @@ -19,6 +20,9 @@ type SettingsResponse struct { APIKeySet bool `json:"api_key_set"` // 是否已设置(用于前端判断要不要提示必填) AuthKey string `json:"auth_key"` // 鉴权 key 打码(客户端访问 + admin 登录用) AuthKeySet bool `json:"auth_key_set"` + CORSAllowAll bool `json:"cors_allow_all"` // 允许所有来源(*) + CORSOrigins string `json:"cors_origins"` // 逗号分隔的白名单(原文,含大小写,trim 末尾 /) + CORSConfigured bool `json:"cors_configured"` // 是否配了 CORS(给 banner 用) DefaultResourceID string `json:"default_resource_id"` DefaultSpeaker string `json:"default_speaker"` DefaultFormat string `json:"default_format"` @@ -54,6 +58,9 @@ func SettingsGetHandler(w http.ResponseWriter, r *http.Request) { APIKeySet: all["api_key"] != "", AuthKey: maskAPIKeyField(all["auth_key"]), AuthKeySet: all["auth_key"] != "", + CORSAllowAll: all["cors_allow_all"] == "1" || all["cors_allow_all"] == "true", + CORSOrigins: all["cors_origins"], + CORSConfigured: all["cors_allow_all"] == "1" || all["cors_allow_all"] == "true" || all["cors_origins"] != "", DefaultResourceID: all["default_resource_id"], DefaultSpeaker: all["default_speaker"], DefaultFormat: all["default_format"], @@ -291,6 +298,99 @@ func SettingsAuthKeyHandler(w http.ResponseWriter, r *http.Request) { _ = json.NewEncoder(w).Encode(map[string]any{"ok": true}) } +// SettingsCORSRequest 是 PUT /api/settings/cors 的 body。 +// 两个字段都可选(至少给一个): +// - allow_all: true → 任意 Origin 都接受(*);设了之后 origins 失效 +// - origins: 一行一个 origin,后端 trim + lower + 去末尾 / +type SettingsCORSRequest struct { + AllowAll *bool `json:"allow_all,omitempty"` + Origins string `json:"origins,omitempty"` // 也接受 string 数组(任一形式) +} + +// SettingsCORSHandler PUT /api/settings/cors +// 鉴权: RequireAdmin。改完立即更新 setting.CORS(进程内生效,跨域请求从下个请求开始按新配置)。 +// 同源豁免由 middleware/cors.go 的 isSameOrigin 处理,不在这里管。 +func SettingsCORSHandler(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPut { + http.Error(w, "Method not allowed", http.StatusMethodNotAllowed) + return + } + s := GetAdminStore() + if s == nil { + middleware.SendJSONError(w, http.StatusServiceUnavailable, "database not ready", "configuration_error", "db_not_ready") + return + } + + r.Body = http.MaxBytesReader(w, r.Body, 1<<10) + var body SettingsCORSRequest + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + middleware.SendJSONError(w, http.StatusBadRequest, "invalid JSON body", "invalid_request_error", "bad_request") + return + } + if body.AllowAll == nil && trimAll(body.Origins) == "" && body.Origins != "" { + // 空 body 不算错误,用户可能是想"清空"(只清 origins 保留现状) + } + if body.AllowAll == nil && body.Origins == "" { + middleware.SendJSONError(w, http.StatusBadRequest, + "at least one of allow_all / origins required", + "invalid_request_error", "no_fields") + return + } + + updates := map[string]string{} + if body.AllowAll != nil { + updates["cors_allow_all"] = boolToStr(*body.AllowAll) + } + if body.Origins != "" { + // 校验每个 origin 至少像 http(s)://... (防止用户填空或填乱字符) + for _, line := range strings.Split(body.Origins, "\n") { + line = strings.TrimSpace(line) + if line == "" { + continue + } + low := strings.ToLower(line) + if !strings.HasPrefix(low, "http://") && !strings.HasPrefix(low, "https://") { + middleware.SendJSONError(w, http.StatusBadRequest, + fmt.Sprintf("invalid origin: %q (must start with http:// or https://)", line), + "invalid_request_error", "origin_invalid") + return + } + } + updates["cors_origins"] = body.Origins + } + if err := s.SettingsSetBatch(updates); err != nil { + log.Printf("[settings] cors set: %v", err) + middleware.SendJSONError(w, http.StatusInternalServerError, "write cors failed", "server_error", "db_write_failed") + return + } + + // 立即刷新 setting.CORS,跨域请求从下个请求开始按新配置生效 + // 复用 LoadRuntimeConfig 的解析逻辑(只取 cors 部分,避免覆盖其它运行时字段) + corsAllowAll, _ := s.SettingsGetBool("cors_allow_all", false) + originsStr := "" + if v, _, _ := s.SettingsGet("cors_origins"); v != "" { + originsStr = v + } + if corsAllowAll { + setting.CORS.AllowAll = true + setting.CORS.Origins = nil + } else if originsStr != "" { + setting.CORS.AllowAll = false + setting.CORS.Origins = setting.SplitOriginsForCORS(originsStr) + } else { + setting.CORS.AllowAll = false + setting.CORS.Origins = nil + } + log.Printf("[settings] cors updated (allow_all=%v origins=%q), runtime active", setting.CORS.AllowAll, originsStr) + w.Header().Set("Content-Type", "application/json; charset=utf-8") + _ = json.NewEncoder(w).Encode(map[string]any{ + "ok": true, + "allow_all": setting.CORS.AllowAll, + "origins": originsStr, + "cors_active": true, + }) +} + // maskAPIKeyField 复用 setting 包的打码风格(前 4 后 4 中间 ****)。 // 单独导出版本避免从 setting 包拉整个 APIKeyMask 之类的工具(那个是 unexported)。 func maskAPIKeyField(s string) string { diff --git a/middleware/cors.go b/middleware/cors.go index c453c51..8e3dc3c 100644 --- a/middleware/cors.go +++ b/middleware/cors.go @@ -6,6 +6,7 @@ import ( "strings" "github.com/volcano-tts/tts-api/common" + "github.com/volcano-tts/tts-api/installer" "github.com/volcano-tts/tts-api/setting" ) @@ -98,6 +99,15 @@ func splitOrigin(origin string) (host, scheme string) { func CORS(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // 安装模式完全跳过 CORS: + // - 用户首次装,不可能提前知道自己的访问域名来配 ALLOWED_ORIGINS + // - 装完进 normal 模式后,设的 CORS 才生效(从 DB 读) + // 这样 install 永远能成功,装完再通过 WebUI 配 CORS。 + if installer.GetMode() == installer.ModeSetup { + next.ServeHTTP(w, r) + return + } + origin := r.Header.Get("Origin") // 无 Origin 头:非跨域请求,跳过 CORS 处理 diff --git a/router/admin.html b/router/admin.html index 690fc1f..21d73ca 100644 --- a/router/admin.html +++ b/router/admin.html @@ -294,6 +294,42 @@ + +
+
跨域 CORS
+
+ 控制哪些前端域名能跨域调 /v1/audio/speech。同源始终放行。 +
+
+ +
勾上后下面白名单失效。仅测试用
+
+
+ + +
当前: {{ corsCurrentLabel }}
+
+
{{ corsErr }}
+
✓ 已保存,跨域请求从下个请求开始按新配置
+
+ +
+
+ + + +
+
+ ⚠ +
+
跨域 CORS 未配置
+
同源可用,跨域会 403。设置 → 跨域 CORS 配。
+
+ +
⚠ {{ actionErr }}
@@ -345,6 +381,18 @@ const settingsForm = ref({ default_resource_id: '', default_speaker: '', default_format: 'mp3', sample_rate: 24000, model: '' }); const apiKeyInput = ref(''); const authKeyInput = ref(''); + + // CORS 配置 (M3 follow-up) + const corsForm = ref({ allow_all: false, origins: '' }); + const corsErr = ref(''); + const corsOk = ref(false); + const savingCors = ref(false); + const corsCurrentLabel = computed(() => { + if (!settings.value) return '未配置'; + if (settings.value.cors_allow_all) return '允许所有来源(*)'; + if (settings.value.cors_origins) return settings.value.cors_origins; + return '未配置(同源可用,跨域会被 403)'; + }); const savingSettings = ref(false); const settingsErr = ref(''); const settingsOk = ref(false); @@ -399,9 +447,28 @@ sample_rate: r.data.sample_rate || 24000, model: r.data.model || '', }; + corsForm.value = { + allow_all: !!r.data.cors_allow_all, + origins: r.data.cors_origins || '', + }; settingsOk.value = false; } catch (e) { settingsErr.value = '加载设置失败: ' + (e.response?.data?.error?.message || e.message); } }; + const saveCors = async () => { + corsErr.value = ''; corsOk.value = false; + savingCors.value = true; + try { + await http.put('/settings/cors', { + allow_all: corsForm.value.allow_all, + origins: corsForm.value.origins, + }); + await loadSettings(); + corsOk.value = true; + setTimeout(() => corsOk.value = false, 3000); + } catch (e) { + corsErr.value = e.response?.data?.error?.message || e.message; + } finally { savingCors.value = false; } + }; const reloadAll = () => { loadOverview(); loadVoices(); }; const saveSettings = async () => { @@ -498,6 +565,7 @@ return { apiKey, keyInput, loginErr, login, logout, tab, overview, voices, actionErr, showAdd, form, addErr, adding, openAdd, submitAdd, toggle, remove, settings, settingsForm, apiKeyInput, authKeyInput, savingSettings, settingsErr, settingsOk, + corsForm, corsErr, corsOk, savingCors, corsCurrentLabel, saveCors, loadSettings, saveSettings, saveApiKey, saveAuthKey, resetSettingsForm, formatUptime, shortPath, reloadAll }; }, diff --git a/router/router.go b/router/router.go index fc301ce..c82781d 100644 --- a/router/router.go +++ b/router/router.go @@ -85,6 +85,7 @@ func Setup() *mux.Router { r.Handle("/api/settings", middleware.RequireAdmin(http.HandlerFunc(controller.SettingsUpdateHandler))).Methods("PUT") r.Handle("/api/settings/api-key", middleware.RequireAdmin(http.HandlerFunc(controller.SettingsAPIKeyHandler))).Methods("PUT") r.Handle("/api/settings/auth-key", middleware.RequireAdmin(http.HandlerFunc(controller.SettingsAuthKeyHandler))).Methods("PUT") + r.Handle("/api/settings/cors", middleware.RequireAdmin(http.HandlerFunc(controller.SettingsCORSHandler))).Methods("PUT") // 业务路由 r.HandleFunc("/v1/audio/speech", controller.OpenaiTTSHandler).Methods("POST", "OPTIONS") diff --git a/setting/config.go b/setting/config.go index 0fcfa28..98ae302 100644 --- a/setting/config.go +++ b/setting/config.go @@ -127,6 +127,30 @@ func normalizeOrigin(origin string) string { return strings.ToLower(origin) } +// SplitOriginsForCORS 解析逗号/换行/空格分隔的 origins 列表, +// 全部小写、trim 末尾 / 后面统一比较。导出供 controller 复用 +// (PUT /api/settings/cors 写完立即刷新 setting.CORS 用)。 +func SplitOriginsForCORS(s string) []string { + return splitAndLowerOrigins(s) +} + +// splitAndLowerOrigins 解析逗号/换行/空格分隔的 origins 列表, +// 全部小写、trim 末尾 / 后面统一比较(只在本包内用,外部用 SplitOriginsForCORS)。 +// 实现细节:用 strings.FieldsFunc 切分,首尾 trim,末尾去 /。 +func splitAndLowerOrigins(s string) []string { + parts := strings.FieldsFunc(s, func(r rune) bool { + return r == ',' || r == '\n' || r == ' ' || r == '\t' + }) + out := make([]string, 0, len(parts)) + for _, p := range parts { + p = strings.ToLower(strings.TrimRight(strings.TrimSpace(p), "/")) + if p != "" { + out = append(out, p) + } + } + return out +} + // LoadRuntimeConfig 从 store 加载 TTS 全局配置到 TTSOptions / TTSTimeout / Auth.APIKeys 内存。 // 启动期(master 模式)调一次,或 PUT /api/settings 后调一次(改完立即生效)。 // @@ -235,6 +259,32 @@ func LoadRuntimeConfig(s Store) error { Auth.APIKeys = nil } + // CORS 配置:DB > env + // cors_allow_all (bool): 允许所有来源(*) + // cors_origins (string): 逗号分隔白名单 + // 同源豁免由 middleware/cors.go 的 isSameOrigin 负责,DB 这里只管跨域名单 + corsAllowAll, _ := s.SettingsGetBool("cors_allow_all", false) + if !corsAllowAll { + // env 兜底 + if v := strings.ToLower(strings.TrimSpace(os.Getenv("CORS_ALLOW_ALL"))); v == "1" || v == "true" || v == "yes" { + corsAllowAll = true + } + } + originsStr := strings.TrimSpace(all["cors_origins"]) + if originsStr == "" { + originsStr = os.Getenv("ALLOWED_ORIGINS") + } + if corsAllowAll { + CORS.AllowAll = true + CORS.Origins = nil + } else if originsStr != "" { + CORS.AllowAll = false + CORS.Origins = SplitOriginsForCORS(originsStr) + } else { + CORS.AllowAll = false + CORS.Origins = nil + } + TTSConfigErr = nil return nil }