From 63611dc932e841e6e0b12af6282ade364771b805 Mon Sep 17 00:00:00 2001 From: MengMengCode <227010654+MengMengCode@users.noreply.github.com> Date: Tue, 9 Jun 2026 22:48:09 +0800 Subject: [PATCH] Support webSSH Origin Allowlist --- backend/internal/api/origins.go | 80 ++++++++++++++ backend/internal/api/websocket.go | 20 +--- backend/internal/config/config.go | 43 +++++--- backend/internal/config/origins.go | 136 ++++++++++++++++++++++++ backend/internal/config/store_sqlite.go | 5 + backend/internal/server/server.go | 36 +------ frontend/src/pages/Settings.tsx | 136 ++++++++++++++++++++---- frontend/src/services/api.ts | 11 ++ frontend/src/utils/i18n.ts | 7 ++ 9 files changed, 386 insertions(+), 88 deletions(-) create mode 100644 backend/internal/api/origins.go create mode 100644 backend/internal/config/origins.go diff --git a/backend/internal/api/origins.go b/backend/internal/api/origins.go new file mode 100644 index 0000000..43e8c94 --- /dev/null +++ b/backend/internal/api/origins.go @@ -0,0 +1,80 @@ +package api + +import ( + "encoding/json" + "net/http" + "strings" + + "clicd/internal/config" +) + +type webSSHOriginSettingsRequest struct { + Origins []string `json:"origins"` + WebSSHAllowedOrigins []string `json:"webssh_allowed_origins"` +} + +type webSSHOriginSettingsResponse struct { + Origins []string `json:"origins"` + CurrentOrigin string `json:"current_origin,omitempty"` +} + +func HandleWebSSHOriginSettings(w http.ResponseWriter, r *http.Request) { + switch r.Method { + case http.MethodGet: + jsonResponse(w, http.StatusOK, APIResponse{Success: true, Data: webSSHOriginSettingsStatus(r)}) + case http.MethodPut: + updateWebSSHOriginSettings(w, r) + default: + jsonResponse(w, http.StatusMethodNotAllowed, APIResponse{Success: false, Message: "Method not allowed"}) + } +} + +func updateWebSSHOriginSettings(w http.ResponseWriter, r *http.Request) { + var req webSSHOriginSettingsRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + jsonResponse(w, http.StatusBadRequest, APIResponse{Success: false, Message: "Invalid request body"}) + return + } + origins := req.Origins + if len(origins) == 0 && len(req.WebSSHAllowedOrigins) > 0 { + origins = req.WebSSHAllowedOrigins + } + normalized, err := config.NormalizeAllowedOrigins(origins) + if err != nil { + jsonResponse(w, http.StatusBadRequest, APIResponse{Success: false, Message: err.Error()}) + return + } + config.AppConfig.WebSSHAllowedOrigins = normalized + if err := config.SaveConfig(); err != nil { + jsonResponse(w, http.StatusInternalServerError, APIResponse{Success: false, Message: "Save Origin allowlist failed"}) + return + } + auditRequest(r, "settings.webssh_origins", "WebSSH Origin", "origins="+strings.Join(normalized, ","), true, "") + jsonResponse(w, http.StatusOK, APIResponse{Success: true, Message: "Origin allowlist saved", Data: webSSHOriginSettingsStatus(r)}) +} + +func webSSHOriginSettingsStatus(r *http.Request) webSSHOriginSettingsResponse { + origins := config.AppConfig.WebSSHAllowedOrigins + if origins == nil { + origins = []string{} + } + return webSSHOriginSettingsResponse{ + Origins: origins, + CurrentOrigin: requestOrigin(r), + } +} + +func requestOrigin(r *http.Request) string { + host := strings.TrimSpace(r.Host) + if host == "" { + return "" + } + scheme := "http" + if r.TLS != nil { + scheme = "https" + } + if forwarded := strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")); forwarded != "" { + scheme = strings.ToLower(strings.Split(forwarded, ",")[0]) + } + return scheme + "://" + host +} diff --git a/backend/internal/api/websocket.go b/backend/internal/api/websocket.go index 337cf98..7f0e097 100644 --- a/backend/internal/api/websocket.go +++ b/backend/internal/api/websocket.go @@ -1,10 +1,9 @@ package api import ( - "net" "net/http" - "net/url" - "strings" + + "clicd/internal/config" "github.com/gorilla/websocket" ) @@ -17,19 +16,6 @@ var upgrader = websocket.Upgrader{ if origin == "" { return true } - originURL, err := url.Parse(origin) - if err != nil { - return false - } - originHost := strings.ToLower(stripPort(originURL.Host)) - requestHost := strings.ToLower(stripPort(r.Host)) - return originHost != "" && originHost == requestHost + return config.IsOriginAllowed(origin, r.Host) }, } - -func stripPort(host string) string { - if parsedHost, _, err := net.SplitHostPort(host); err == nil { - return parsedHost - } - return strings.Trim(host, "[]") -} diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index a14bdc4..cd54f0e 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -363,6 +363,7 @@ type ClicdConfig struct { Snapshots []Snapshot `json:"snapshots"` PublicIPv4Pool []PublicIPv4Assignment `json:"public_ipv4_pool"` PublicIPv6Prefixes []PublicIPv6Prefix `json:"public_ipv6_prefixes"` + WebSSHAllowedOrigins []string `json:"webssh_allowed_origins"` SecurityAutoShutdown bool `json:"security_auto_shutdown"` Language string `json:"language"` SSL SSLConfig `json:"ssl"` @@ -480,23 +481,24 @@ func InitConfig() (*ClicdConfig, error) { } AppConfig = &ClicdConfig{ - AdminUser: adminUser, - AdminPassHash: string(hash), - JWTSecret: jwtSecret, - Port: 8999, - DataDir: dataDir, - Containers: []Container{}, - NextContainerID: 1, - NextVNCPort: 5900, - NextSSHPort: 22000, - SetupComplete: false, - SubUsers: []SubUser{}, - AuditLogs: []AuditLog{}, - Tasks: []SavedTask{}, - LoginLogs: []SavedLoginLog{}, - Snapshots: []Snapshot{}, - PublicIPv4Pool: []PublicIPv4Assignment{}, - PublicIPv6Prefixes: []PublicIPv6Prefix{}, + AdminUser: adminUser, + AdminPassHash: string(hash), + JWTSecret: jwtSecret, + Port: 8999, + DataDir: dataDir, + Containers: []Container{}, + NextContainerID: 1, + NextVNCPort: 5900, + NextSSHPort: 22000, + SetupComplete: false, + SubUsers: []SubUser{}, + AuditLogs: []AuditLog{}, + Tasks: []SavedTask{}, + LoginLogs: []SavedLoginLog{}, + Snapshots: []Snapshot{}, + PublicIPv4Pool: []PublicIPv4Assignment{}, + PublicIPv6Prefixes: []PublicIPv6Prefix{}, + WebSSHAllowedOrigins: []string{}, } if err := SaveConfig(); err != nil { @@ -555,6 +557,13 @@ func normalizeConfigDefaults(dataDir string) bool { AppConfig.PublicIPv6Prefixes = make([]PublicIPv6Prefix, 0) changed = true } + if AppConfig.WebSSHAllowedOrigins == nil { + AppConfig.WebSSHAllowedOrigins = make([]string, 0) + changed = true + } else if normalized, err := NormalizeAllowedOrigins(AppConfig.WebSSHAllowedOrigins); err == nil && strings.Join(normalized, "\n") != strings.Join(AppConfig.WebSSHAllowedOrigins, "\n") { + AppConfig.WebSSHAllowedOrigins = normalized + changed = true + } if AppConfig.SubUsers == nil { AppConfig.SubUsers = make([]SubUser, 0) changed = true diff --git a/backend/internal/config/origins.go b/backend/internal/config/origins.go new file mode 100644 index 0000000..846f8eb --- /dev/null +++ b/backend/internal/config/origins.go @@ -0,0 +1,136 @@ +package config + +import ( + "fmt" + "net" + "net/url" + "strings" +) + +// NormalizeAllowedOrigin accepts a browser Origin value such as +// https://www.example.com and returns a canonical form for exact matching. +func NormalizeAllowedOrigin(value string) (string, error) { + value = strings.TrimSpace(value) + if value == "" { + return "", nil + } + u, err := url.Parse(value) + if err != nil || u.Scheme == "" || u.Host == "" { + return "", fmt.Errorf("Origin must include scheme and host: %s", value) + } + scheme := strings.ToLower(u.Scheme) + if scheme != "http" && scheme != "https" { + return "", fmt.Errorf("Origin scheme must be http or https: %s", value) + } + if (u.Path != "" && u.Path != "/") || u.RawQuery != "" || u.Fragment != "" { + return "", fmt.Errorf("Origin must not include path, query, or fragment: %s", value) + } + host := normalizeOriginHostPort(u.Host, scheme) + if host == "" { + return "", fmt.Errorf("Origin host is required: %s", value) + } + return scheme + "://" + host, nil +} + +func NormalizeAllowedOrigins(values []string) ([]string, error) { + result := make([]string, 0, len(values)) + seen := map[string]bool{} + for _, value := range values { + origin, err := NormalizeAllowedOrigin(value) + if err != nil { + return nil, err + } + if origin == "" || seen[origin] { + continue + } + seen[origin] = true + result = append(result, origin) + } + return result, nil +} + +func IsOriginAllowed(origin string, requestHost string) bool { + origin = strings.TrimSpace(origin) + if origin == "" { + return true + } + if isSameRequestOrigin(origin, requestHost) { + return true + } + normalized, err := NormalizeAllowedOrigin(origin) + if err != nil { + return false + } + if AppConfig == nil { + return false + } + for _, allowed := range AppConfig.WebSSHAllowedOrigins { + allowed, err := NormalizeAllowedOrigin(allowed) + if err == nil && normalized == allowed { + return true + } + } + return false +} + +func isSameRequestOrigin(origin string, requestHost string) bool { + u, err := url.Parse(origin) + if err != nil || u.Host == "" { + return false + } + originHost := normalizeHostOnly(u.Hostname()) + host := normalizeHostOnly(requestHost) + if originHost == "" || host == "" { + return false + } + if originHost == host { + return true + } + return isLoopbackHost(originHost) && isLoopbackHost(host) +} + +func normalizeOriginHostPort(raw string, scheme string) string { + host := raw + port := "" + if h, p, err := net.SplitHostPort(raw); err == nil { + host = h + port = p + } + host = normalizeHostOnly(host) + if host == "" { + return "" + } + if (scheme == "https" && port == "443") || (scheme == "http" && port == "80") { + port = "" + } + if port != "" { + return net.JoinHostPort(host, port) + } + if strings.Contains(host, ":") && net.ParseIP(host) != nil { + return "[" + host + "]" + } + return host +} + +func normalizeHostOnly(raw string) string { + raw = strings.TrimSpace(raw) + if raw == "" { + return "" + } + if h, _, err := net.SplitHostPort(raw); err == nil { + raw = h + } + raw = strings.Trim(raw, "[]") + if ip := net.ParseIP(raw); ip != nil { + return strings.ToLower(ip.String()) + } + return strings.TrimSuffix(strings.ToLower(raw), ".") +} + +func isLoopbackHost(host string) bool { + if host == "localhost" { + return true + } + ip := net.ParseIP(host) + return ip != nil && ip.IsLoopback() +} diff --git a/backend/internal/config/store_sqlite.go b/backend/internal/config/store_sqlite.go index 8423007..e074b77 100644 --- a/backend/internal/config/store_sqlite.go +++ b/backend/internal/config/store_sqlite.go @@ -426,6 +426,9 @@ func loadConfigFromDB() (*ClicdConfig, bool, error) { if raw := strings.TrimSpace(meta["public_ipv6_prefixes"]); raw != "" { _ = json.Unmarshal([]byte(raw), &cfg.PublicIPv6Prefixes) } + if raw := strings.TrimSpace(meta["webssh_allowed_origins"]); raw != "" { + _ = json.Unmarshal([]byte(raw), &cfg.WebSSHAllowedOrigins) + } if cfg.Containers, err = loadContainers(); err != nil { return nil, false, err @@ -524,6 +527,7 @@ func saveMeta(tx *sql.Tx) error { sslCertificatesJSON, _ := json.Marshal(AppConfig.SSLCertificates) publicIPv4PoolJSON, _ := json.Marshal(AppConfig.PublicIPv4Pool) publicIPv6PrefixesJSON, _ := json.Marshal(AppConfig.PublicIPv6Prefixes) + webSSHAllowedOriginsJSON, _ := json.Marshal(AppConfig.WebSSHAllowedOrigins) values := map[string]string{ "admin_user": AppConfig.AdminUser, "admin_pass_hash": AppConfig.AdminPassHash, @@ -540,6 +544,7 @@ func saveMeta(tx *sql.Tx) error { "ssl_certificates": string(sslCertificatesJSON), "public_ipv4_pool": string(publicIPv4PoolJSON), "public_ipv6_prefixes": string(publicIPv6PrefixesJSON), + "webssh_allowed_origins": string(webSSHAllowedOriginsJSON), "schema_version": "1", "updated_at": time.Now().Format("2006-01-02 15:04:05"), } diff --git a/backend/internal/server/server.go b/backend/internal/server/server.go index 4513693..f41861b 100644 --- a/backend/internal/server/server.go +++ b/backend/internal/server/server.go @@ -4,9 +4,7 @@ import ( "crypto/tls" "fmt" "log" - "net" "net/http" - "net/url" "strings" "clicd/internal/api" @@ -19,7 +17,7 @@ var webFS http.FileSystem // corsMiddleware adds CORS headers func corsMiddleware(next http.HandlerFunc) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - if origin := r.Header.Get("Origin"); origin != "" && isAllowedOrigin(origin, r.Host) { + if origin := r.Header.Get("Origin"); origin != "" && config.IsOriginAllowed(origin, r.Host) { w.Header().Set("Access-Control-Allow-Origin", origin) w.Header().Set("Vary", "Origin") w.Header().Set("Access-Control-Allow-Credentials", "true") @@ -28,7 +26,7 @@ func corsMiddleware(next http.HandlerFunc) http.HandlerFunc { w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization, X-API-Key") if r.Method == http.MethodOptions { - if origin := r.Header.Get("Origin"); origin != "" && !isAllowedOrigin(origin, r.Host) { + if origin := r.Header.Get("Origin"); origin != "" && !config.IsOriginAllowed(origin, r.Host) { w.WriteHeader(http.StatusForbidden) return } @@ -40,34 +38,6 @@ func corsMiddleware(next http.HandlerFunc) http.HandlerFunc { } } -func isAllowedOrigin(origin string, requestHost string) bool { - u, err := url.Parse(origin) - if err != nil || u.Host == "" { - return false - } - originHost := normalizeHost(u.Host) - host := normalizeHost(requestHost) - if originHost == host { - return true - } - return isLoopbackHost(originHost) && isLoopbackHost(host) -} - -func normalizeHost(host string) string { - if h, _, err := net.SplitHostPort(host); err == nil { - return strings.ToLower(h) - } - return strings.ToLower(host) -} - -func isLoopbackHost(host string) bool { - if host == "localhost" { - return true - } - ip := net.ParseIP(host) - return ip != nil && ip.IsLoopback() -} - // setupRoutes configures API and static routes func setupRoutes(mux *http.ServeMux) { // API routes @@ -78,6 +48,7 @@ func setupRoutes(mux *http.ServeMux) { mux.HandleFunc("/api/change-username", corsMiddleware(api.AdminMiddleware(api.HandleAdminUsernameChange))) mux.HandleFunc("/api/login-logs", corsMiddleware(api.AdminMiddleware(api.HandleLoginLogs))) mux.HandleFunc("/api/ssl", corsMiddleware(api.AdminMiddleware(api.HandleSSLSettings))) + mux.HandleFunc("/api/webssh-origins", corsMiddleware(api.AdminMiddleware(api.HandleWebSSHOriginSettings))) mux.HandleFunc("/api/containers", corsMiddleware(api.AuthMiddleware(api.SubUserMiddleware(api.HandleContainers)))) mux.HandleFunc("/api/containers/list", corsMiddleware(api.AuthMiddleware(api.SubUserMiddleware(api.HandleContainerListAlias)))) mux.HandleFunc("/api/containers/", corsMiddleware(api.AuthMiddleware(api.SubUserMiddleware(api.HandleSingleContainer)))) @@ -148,6 +119,7 @@ func setupRoutes(mux *http.ServeMux) { mux.HandleFunc("/api/v1/audit-logs", corsMiddleware(api.AuthMiddleware(api.HandleAuditLogs))) mux.HandleFunc("/api/v1/login-logs", corsMiddleware(api.AuthMiddleware(api.HandleLoginLogs))) mux.HandleFunc("/api/v1/ssl", corsMiddleware(api.AdminMiddleware(api.HandleSSLSettings))) + mux.HandleFunc("/api/v1/webssh-origins", corsMiddleware(api.AdminMiddleware(api.HandleWebSSHOriginSettings))) mux.HandleFunc("/api/v1/security/alerts", corsMiddleware(api.AuthMiddleware(api.ScopeMiddleware("security:read", api.HandleSecurityAlerts)))) mux.HandleFunc("/api/v1/security/check", corsMiddleware(api.AuthMiddleware(api.ScopeMiddleware("security:check", api.HandleSecurityCheck)))) mux.HandleFunc("/api/v1/security/logs", corsMiddleware(api.AuthMiddleware(api.ScopeMiddleware("security:read", api.HandleSecurityLogs)))) diff --git a/frontend/src/pages/Settings.tsx b/frontend/src/pages/Settings.tsx index 6c059b5..088b9f4 100644 --- a/frontend/src/pages/Settings.tsx +++ b/frontend/src/pages/Settings.tsx @@ -1,13 +1,16 @@ import { Dispatch, SetStateAction, useCallback, useEffect, useState } from 'react' -import { Clock, Globe, Lock, LogIn, Monitor, RefreshCw, ShieldCheck, Upload, UserCog } from 'lucide-react' +import { Clock, Globe, Lock, LogIn, Monitor, RefreshCw, ShieldCheck, Terminal, Upload, UserCog } from 'lucide-react' import { changePassword, changeUsername, getLoginLogs, getSSLSettings, + getWebSSHOriginSettings, LoginLog, SSLSettings, updateSSLSettings, + updateWebSSHOriginSettings, + WebSSHOriginSettings, } from '../services/api' import { useDialog } from '../components/Dialog' import { useAuth } from '../contexts/AuthContext' @@ -33,6 +36,9 @@ export default function Settings() { const [keyPEM, setKeyPEM] = useState('') const [applyNow, setApplyNow] = useState(true) const [savingSSL, setSavingSSL] = useState(false) + const [webSSHOrigins, setWebSSHOrigins] = useState(null) + const [webSSHOriginsText, setWebSSHOriginsText] = useState('') + const [savingWebSSHOrigins, setSavingWebSSHOrigins] = useState(false) const fetchLogs = useCallback(async () => { try { @@ -60,12 +66,25 @@ export default function Settings() { } }, []) + const fetchWebSSHOrigins = useCallback(async () => { + try { + const res = await getWebSSHOriginSettings() + const data = res.data.data + if (!data) return + setWebSSHOrigins(data) + setWebSSHOriginsText((data.origins || []).join('\n')) + } catch (err) { + console.error(err) + } + }, []) + useEffect(() => { fetchLogs() fetchSSL() + fetchWebSSHOrigins() const timer = setInterval(fetchLogs, 15000) return () => clearInterval(timer) - }, [fetchLogs, fetchSSL]) + }, [fetchLogs, fetchSSL, fetchWebSSHOrigins]) const handleSSLModeChange = (mode: SSLSettings['mode']) => { setSSLMode(mode) @@ -101,6 +120,25 @@ export default function Settings() { } } + const handleSaveWebSSHOrigins = async () => { + setSavingWebSSHOrigins(true) + try { + const origins = webSSHOriginsText.split(/\r?\n/).map(item => item.trim()).filter(Boolean) + const res = await updateWebSSHOriginSettings(origins) + const data = res.data.data + if (data) { + setWebSSHOrigins(data) + setWebSSHOriginsText((data.origins || []).join('\n')) + } + dialog.alert('完成', 'Origin 白名单已保存') + } catch (err: unknown) { + const e = err as { response?: { data?: { message?: string } } } + dialog.alert('失败', e.response?.data?.message || 'Origin 白名单保存失败') + } finally { + setSavingWebSSHOrigins(false) + } + } + const handleSaveAccount = async () => { if (!oldPwd) { dialog.alert('提示', '请输入当前密码以确认修改') @@ -159,26 +197,37 @@ export default function Settings() {
- +
+ + + +

@@ -232,6 +281,49 @@ interface SSLCardProps { onSave: () => void } +interface WebSSHOriginCardProps { + settings: WebSSHOriginSettings | null + originsText: string + saving: boolean + onOriginsTextChange: (value: string) => void + onRefresh: () => void + onSave: () => void +} + +function WebSSHOriginCard(props: WebSSHOriginCardProps) { + return ( +
+
+

+ WebSSH Origin 白名单 +

+ +
+
+
+ +