mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-06 18:06:36 +08:00
fix(security): patch SSRF and config info disclosure vulnerabilities
This commit is contained in:
@@ -350,6 +350,12 @@ func (h *Handler) getConfigByName(w http.ResponseWriter, r *http.Request) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
switch req.Name {
|
||||||
|
case "license_key", "cloudflare_secret_key", "jwt_secret":
|
||||||
|
response.WriteJSON(w, response.Err(403, "禁止访问敏感配置"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
cfg, err := h.repo.GetConfigByName(req.Name)
|
cfg, err := h.repo.GetConfigByName(req.Name)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
@@ -374,6 +380,15 @@ func (h *Handler) getConfigs(w http.ResponseWriter, r *http.Request) {
|
|||||||
response.WriteJSON(w, response.Err(-2, err.Error()))
|
response.WriteJSON(w, response.Err(-2, err.Error()))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Filter sensitive config keys if the user is not an admin
|
||||||
|
ctxClaims := r.Context().Value(middleware.ClaimsContextKey)
|
||||||
|
if claims, ok := ctxClaims.(auth.Claims); !ok || claims.RoleID != 0 {
|
||||||
|
delete(cfgMap, "license_key")
|
||||||
|
delete(cfgMap, "cloudflare_secret_key")
|
||||||
|
delete(cfgMap, "jwt_secret")
|
||||||
|
}
|
||||||
|
|
||||||
response.WriteJSON(w, response.OK(cfgMap))
|
response.WriteJSON(w, response.OK(cfgMap))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -301,6 +301,10 @@ func (h *Handler) nodeCreate(w http.ResponseWriter, r *http.Request) {
|
|||||||
response.WriteJSON(w, response.ErrDefault("节点名称和地址不能为空"))
|
response.WriteJSON(w, response.ErrDefault("节点名称和地址不能为空"))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if err := IsValidNodeAddress(serverIP); err != nil {
|
||||||
|
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
now := time.Now().UnixMilli()
|
now := time.Now().UnixMilli()
|
||||||
inx := h.repo.NextIndex("node")
|
inx := h.repo.NextIndex("node")
|
||||||
@@ -365,6 +369,15 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
|||||||
newHTTP := asInt(req["http"], currentHTTP)
|
newHTTP := asInt(req["http"], currentHTTP)
|
||||||
newTLS := asInt(req["tls"], currentTLS)
|
newTLS := asInt(req["tls"], currentTLS)
|
||||||
newSocks := asInt(req["socks"], currentSocks)
|
newSocks := asInt(req["socks"], currentSocks)
|
||||||
|
|
||||||
|
serverIP := asString(req["serverIp"])
|
||||||
|
if serverIP != "" {
|
||||||
|
if err := IsValidNodeAddress(serverIP); err != nil {
|
||||||
|
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if currentStatus == 1 && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
|
if currentStatus == 1 && (newHTTP != currentHTTP || newTLS != currentTLS || newSocks != currentSocks) {
|
||||||
if err := h.applyNodeProtocolChange(id, newHTTP, newTLS, newSocks); err != nil {
|
if err := h.applyNodeProtocolChange(id, newHTTP, newTLS, newSocks); err != nil {
|
||||||
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
response.WriteJSON(w, response.ErrDefault(err.Error()))
|
||||||
@@ -375,7 +388,7 @@ func (h *Handler) nodeUpdate(w http.ResponseWriter, r *http.Request) {
|
|||||||
now := time.Now().UnixMilli()
|
now := time.Now().UnixMilli()
|
||||||
if err := h.repo.UpdateNode(id,
|
if err := h.repo.UpdateNode(id,
|
||||||
asString(req["name"]),
|
asString(req["name"]),
|
||||||
asString(req["serverIp"]),
|
serverIP,
|
||||||
nullableText(asString(req["serverIpV4"])),
|
nullableText(asString(req["serverIpV4"])),
|
||||||
nullableText(asString(req["serverIpV6"])),
|
nullableText(asString(req["serverIpV6"])),
|
||||||
defaultString(asString(req["port"]), "1000-65535"),
|
defaultString(asString(req["port"]), "1000-65535"),
|
||||||
@@ -1680,6 +1693,10 @@ func (h *Handler) forwardCreate(w http.ResponseWriter, r *http.Request) {
|
|||||||
response.WriteJSON(w, response.Err(-1, "普通用户无法设置限速规则"))
|
response.WriteJSON(w, response.Err(-1, "普通用户无法设置限速规则"))
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||||
|
response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络或保留地址"))
|
||||||
|
return
|
||||||
|
}
|
||||||
}
|
}
|
||||||
speedID := asAnyToInt64Ptr(req["speedId"])
|
speedID := asAnyToInt64Ptr(req["speedId"])
|
||||||
speedID, err = h.normalizeSpeedLimitReference(speedID)
|
speedID, err = h.normalizeSpeedLimitReference(speedID)
|
||||||
@@ -1795,6 +1812,12 @@ func (h *Handler) forwardUpdate(w http.ResponseWriter, r *http.Request) {
|
|||||||
if remoteAddr == "" {
|
if remoteAddr == "" {
|
||||||
remoteAddr = forward.RemoteAddr
|
remoteAddr = forward.RemoteAddr
|
||||||
}
|
}
|
||||||
|
if actorRole != 0 {
|
||||||
|
if err := IsSafeRemoteAddr(remoteAddr); err != nil {
|
||||||
|
response.WriteJSON(w, response.Err(403, "禁止将目标地址设置为内部网络或保留地址"))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
strategy := strings.TrimSpace(asString(req["strategy"]))
|
strategy := strings.TrimSpace(asString(req["strategy"]))
|
||||||
if strategy == "" {
|
if strategy == "" {
|
||||||
strategy = forward.Strategy
|
strategy = forward.Strategy
|
||||||
|
|||||||
@@ -0,0 +1,57 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// IsSafeRemoteAddr checks if a given address is safe to connect to (prevents SSRF/Open Proxy).
|
||||||
|
// It resolves domains to IPs to prevent DNS rebinding attacks pointing to internal networks.
|
||||||
|
func IsSafeRemoteAddr(addr string) error {
|
||||||
|
host, _, err := net.SplitHostPort(addr)
|
||||||
|
if err != nil {
|
||||||
|
// If there is no port, try to treat the whole string as host
|
||||||
|
if strings.Contains(err.Error(), "missing port in address") {
|
||||||
|
host = addr
|
||||||
|
} else {
|
||||||
|
return fmt.Errorf("invalid address format: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
ips, err := net.LookupIP(host)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("could not resolve address: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, ip := range ips {
|
||||||
|
if ip.IsLoopback() || ip.IsPrivate() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() || ip.IsUnspecified() || ip.IsMulticast() {
|
||||||
|
return fmt.Errorf("address resolves to internal or reserved IP: %s", ip.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsValidNodeAddress ensures the address is strictly a host or host:port.
|
||||||
|
// It explicitly denies schemes (http://, https://), paths (/...), and query params (?).
|
||||||
|
func IsValidNodeAddress(addr string) error {
|
||||||
|
addr = strings.TrimSpace(addr)
|
||||||
|
if strings.Contains(addr, "://") {
|
||||||
|
return fmt.Errorf("address must not contain scheme (e.g. http://)")
|
||||||
|
}
|
||||||
|
if strings.ContainsAny(addr, "/?") {
|
||||||
|
return fmt.Errorf("address must not contain path or query parameters")
|
||||||
|
}
|
||||||
|
|
||||||
|
// Try parsing as host:port
|
||||||
|
_, _, err := net.SplitHostPort(addr)
|
||||||
|
if err != nil {
|
||||||
|
if strings.Contains(err.Error(), "missing port in address") {
|
||||||
|
// It's just a host, which is fine
|
||||||
|
} else {
|
||||||
|
return fmt.Errorf("invalid address format")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user