[优化] 修复隧道相关配置支持

This commit is contained in:
ryan
2026-06-01 20:34:02 +08:00
parent 2635a47d29
commit 3b979eb5d5
10 changed files with 245 additions and 24 deletions
+52 -2
View File
@@ -89,6 +89,10 @@ type snapshotRoute struct {
BasicAuthUsername string `json:"basic_auth_username,omitempty"`
BasicAuthPassword string `json:"basic_auth_password,omitempty"`
Remark string `json:"remark,omitempty"`
UpstreamType string `json:"upstream_type,omitempty"`
TunnelNodeID *uint `json:"tunnel_node_id,omitempty"`
TunnelTargetAddr string `json:"tunnel_target_addr,omitempty"`
TunnelTargetProto string `json:"tunnel_target_protocol,omitempty"`
}
type snapshotWAFRuleGroup struct {
@@ -500,10 +504,22 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
if err != nil {
return nil, fmt.Errorf("路由 %s 自定义请求头无效", route.Domain)
}
upstreamType := normalizeUpstreamType(route.UpstreamType)
originURL := route.OriginURL
upstreams, err := decodeStoredUpstreams(route.Upstreams, route.OriginURL)
if err != nil {
return nil, fmt.Errorf("路由 %s 上游配置无效", route.Domain)
}
var tunnelNodeID *uint
var tunnelTargetAddr string
var tunnelTargetProtocol string
if upstreamType == "tunnel" {
originURL = resolveTunnelOpenRestyUpstreamURL()
upstreams = []string{originURL}
tunnelNodeID = route.TunnelNodeID
tunnelTargetAddr = strings.TrimSpace(route.TunnelTargetAddr)
tunnelTargetProtocol = normalizeTunnelTargetProtocol(route.TunnelTargetProtocol)
}
cacheRules, err := decodeStoredCacheRules(route.CacheRules)
if err != nil {
return nil, fmt.Errorf("路由 %s 缓存规则无效", route.Domain)
@@ -520,7 +536,7 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
SiteName: normalizeProxyRouteSiteNameInput(route, route.SiteName, domains[0]),
Domain: domains[0],
Domains: domains,
OriginURL: route.OriginURL,
OriginURL: originURL,
OriginHost: route.OriginHost,
Upstreams: upstreams,
Enabled: route.Enabled,
@@ -542,11 +558,29 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) {
BasicAuthUsername: route.BasicAuthUsername,
BasicAuthPassword: route.BasicAuthPassword,
Remark: route.Remark,
UpstreamType: upstreamType,
TunnelNodeID: tunnelNodeID,
TunnelTargetAddr: tunnelTargetAddr,
TunnelTargetProto: tunnelTargetProtocol,
})
}
return items, nil
}
func resolveTunnelOpenRestyUpstreamURL() string {
port := 8080
relayNodes, err := model.ListNodesByType("tunnel_relay")
if err == nil {
for _, node := range relayNodes {
if node != nil && node.RelayVhostHTTPPort > 0 {
port = node.RelayVhostHTTPPort
break
}
}
}
return fmt.Sprintf("http://127.0.0.1:%d", port)
}
func buildSnapshotWAFDocument(routes []*model.ProxyRoute) (snapshotWAFDocument, error) {
if err := EnsureDefaultWAFRuleGroup(); err != nil {
return snapshotWAFDocument{}, err
@@ -778,6 +812,15 @@ func normalizeSnapshotRoutes(routes []snapshotRoute) []snapshotRoute {
routes[index].BasicAuthUsername = ""
routes[index].BasicAuthPassword = ""
}
routes[index].UpstreamType = normalizeUpstreamType(routes[index].UpstreamType)
if routes[index].UpstreamType == "tunnel" {
routes[index].TunnelTargetAddr = strings.TrimSpace(routes[index].TunnelTargetAddr)
routes[index].TunnelTargetProto = normalizeTunnelTargetProtocol(routes[index].TunnelTargetProto)
} else {
routes[index].TunnelNodeID = nil
routes[index].TunnelTargetAddr = ""
routes[index].TunnelTargetProto = ""
}
}
return routes
}
@@ -803,7 +846,7 @@ func flattenSnapshotRoutesByDomain(routes []snapshotRoute) map[string]snapshotRo
}
func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool {
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP || left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || left.PoWEnabled != right.PoWEnabled || left.BasicAuthEnabled != right.BasicAuthEnabled || left.BasicAuthUsername != right.BasicAuthUsername || left.BasicAuthPassword != right.BasicAuthPassword || !uintSliceEqual(left.CertIDs, right.CertIDs) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) {
if left.SiteName != right.SiteName || left.Domain != right.Domain || left.OriginURL != right.OriginURL || left.OriginHost != right.OriginHost || left.EnableHTTPS != right.EnableHTTPS || left.RedirectHTTP != right.RedirectHTTP || left.LimitConnPerServer != right.LimitConnPerServer || left.LimitConnPerIP != right.LimitConnPerIP || left.LimitRate != right.LimitRate || left.CacheEnabled != right.CacheEnabled || left.CachePolicy != right.CachePolicy || left.PoWEnabled != right.PoWEnabled || left.BasicAuthEnabled != right.BasicAuthEnabled || left.BasicAuthUsername != right.BasicAuthUsername || left.BasicAuthPassword != right.BasicAuthPassword || left.UpstreamType != right.UpstreamType || !uintPtrEqual(left.TunnelNodeID, right.TunnelNodeID) || left.TunnelTargetAddr != right.TunnelTargetAddr || left.TunnelTargetProto != right.TunnelTargetProto || !uintSliceEqual(left.CertIDs, right.CertIDs) || !uintSliceEqual(left.DomainCertIDs, right.DomainCertIDs) {
return false
}
if len(left.Domains) != len(right.Domains) {
@@ -1165,6 +1208,13 @@ func uintSliceEqual(left []uint, right []uint) bool {
return true
}
func uintPtrEqual(left *uint, right *uint) bool {
if left == nil || right == nil {
return left == nil && right == nil
}
return *left == *right
}
func nextVersionNumber(now time.Time) (string, error) {
prefix := now.Format("20060102")
var latest model.ConfigVersion
+46 -1
View File
@@ -63,6 +63,7 @@ type ProxyRouteInput struct {
Remark string `json:"remark"`
UpstreamType string `json:"upstream_type"`
TunnelNodeID *uint `json:"tunnel_node_id"`
TunnelID *uint `json:"tunnel_id"`
TunnelTargetAddr string `json:"tunnel_target_addr"`
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
}
@@ -102,6 +103,7 @@ type ProxyRouteView struct {
Remark string `json:"remark"`
UpstreamType string `json:"upstream_type"`
TunnelNodeID *uint `json:"tunnel_node_id"`
TunnelID *uint `json:"tunnel_id"`
TunnelTargetAddr string `json:"tunnel_target_addr"`
TunnelTargetProtocol string `json:"tunnel_target_protocol"`
CreatedAt time.Time `json:"created_at"`
@@ -327,7 +329,14 @@ func buildProxyRoute(route *model.ProxyRoute, input ProxyRouteInput) (*model.Pro
route.Remark = remark
route.UpstreamType = upstreamType
if upstreamType == "tunnel" {
route.TunnelNodeID = input.TunnelNodeID
tunnelNodeID, err := normalizeTunnelNodeID(input.TunnelNodeID, input.TunnelID)
if err != nil {
return nil, err
}
if err := validateTunnelRouteInput(tunnelNodeID, input.TunnelTargetAddr, input.TunnelTargetProtocol); err != nil {
return nil, err
}
route.TunnelNodeID = tunnelNodeID
route.TunnelTargetAddr = strings.TrimSpace(input.TunnelTargetAddr)
route.TunnelTargetProtocol = normalizeTunnelTargetProtocol(input.TunnelTargetProtocol)
} else {
@@ -422,6 +431,7 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) {
Remark: route.Remark,
UpstreamType: route.UpstreamType,
TunnelNodeID: route.TunnelNodeID,
TunnelID: route.TunnelNodeID,
TunnelTargetAddr: route.TunnelTargetAddr,
TunnelTargetProtocol: route.TunnelTargetProtocol,
CreatedAt: route.CreatedAt,
@@ -429,6 +439,41 @@ func buildProxyRouteView(route *model.ProxyRoute) (*ProxyRouteView, error) {
}, nil
}
func normalizeTunnelNodeID(tunnelNodeID *uint, legacyTunnelID *uint) (*uint, error) {
if tunnelNodeID != nil && *tunnelNodeID != 0 {
return tunnelNodeID, nil
}
if legacyTunnelID != nil && *legacyTunnelID != 0 {
return legacyTunnelID, nil
}
return nil, errors.New("tunnel_node_id is required for tunnel upstream")
}
func validateTunnelRouteInput(tunnelNodeID *uint, targetAddr string, targetProtocol string) error {
if tunnelNodeID == nil || *tunnelNodeID == 0 {
return errors.New("tunnel_node_id is required for tunnel upstream")
}
tunnelNode, err := model.GetNodeByID(*tunnelNodeID)
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New("tunnel client node does not exist")
}
return err
}
if tunnelNode.NodeType != "tunnel_client" {
return errors.New("tunnel_node_id must reference a tunnel_client node")
}
if strings.TrimSpace(targetAddr) == "" {
return errors.New("tunnel_target_addr is required for tunnel upstream")
}
switch strings.ToLower(strings.TrimSpace(targetProtocol)) {
case "", "http", "https":
return nil
default:
return errors.New("tunnel_target_protocol must be http or https")
}
}
func normalizeProxyRouteSiteNameInput(route *model.ProxyRoute, raw string, primaryDomain string) string {
siteName := strings.TrimSpace(raw)
if siteName != "" {
+27 -4
View File
@@ -4,8 +4,10 @@ import (
"errors"
"fmt"
"log/slog"
"net"
"openflare/common"
"openflare/model"
"strconv"
"strings"
"time"
@@ -369,10 +371,7 @@ func GetFlaredTunnelConfig(node *model.Node) (*FlaredTunnelConfigResponse, error
relays := make([]FlaredRelayInfo, 0, len(relayNodes))
for _, node := range relayNodes {
if node.RelayStatus == "healthy" || node.Status == NodeStatusOnline {
addr := strings.TrimSpace(node.RelayClientAccessAddr)
if addr == "" {
addr = fmt.Sprintf("%s:%d", strings.TrimSpace(node.IP), node.RelayBindPort)
}
addr := relayClientAddress(node)
relays = append(relays, FlaredRelayInfo{
RelayNodeID: node.NodeID,
Address: addr,
@@ -413,6 +412,30 @@ func GetFlaredTunnelConfig(node *model.Node) (*FlaredTunnelConfigResponse, error
}, nil
}
func relayClientAddress(node *model.Node) string {
if node == nil {
return ""
}
port := node.RelayBindPort
if port <= 0 {
port = 7000
}
addr := strings.TrimSpace(node.RelayClientAccessAddr)
if addr == "" {
addr = strings.TrimSpace(node.IP)
}
if addr == "" {
return fmt.Sprintf("127.0.0.1:%d", port)
}
if _, _, err := net.SplitHostPort(addr); err == nil {
return addr
}
if strings.Contains(addr, ":") && strings.Count(addr, ":") > 1 {
return net.JoinHostPort(addr, strconv.Itoa(port))
}
return fmt.Sprintf("%s:%d", addr, port)
}
func parseTunnelTargetAddr(addr string) (string, int) {
addr = strings.TrimSpace(addr)
if addr == "" {
+78
View File
@@ -3,6 +3,7 @@ package service
import (
"errors"
"openflare/model"
"strings"
"testing"
"time"
@@ -295,3 +296,80 @@ func TestGetFlaredTunnelConfigRequiresActiveVersion(t *testing.T) {
t.Logf("GetFlaredTunnelConfig returned wrapped error: %v", err)
}
}
func TestTunnelRoutePublishAndFlaredConfigUseRelayPorts(t *testing.T) {
setupServiceTestDB(t)
relayNode := &model.Node{
NodeID: "node-relay-ports",
Name: "relay-ports",
IP: "85.235.64.179",
AccessToken: "relay-token-ports",
Status: NodeStatusOnline,
NodeType: "tunnel_relay",
RelayStatus: "healthy",
RelayBindPort: 17000,
RelayVhostHTTPPort: 18080,
RelayAuthToken: "relay-auth-token",
RelayClientAccessAddr: "de-e",
}
if err := relayNode.Insert(); err != nil {
t.Fatalf("failed to seed relay node: %v", err)
}
tunnelNode := &model.Node{
NodeID: "node-flared-ports",
Name: "flared-ports",
IP: "",
AccessToken: "tunnel-token-ports",
Status: NodeStatusOnline,
NodeType: "tunnel_client",
Version: "v0.2.0",
}
if err := tunnelNode.Insert(); err != nil {
t.Fatalf("failed to seed tunnel client node: %v", err)
}
route, err := CreateProxyRoute(ProxyRouteInput{
Domain: "flared.example.com",
UpstreamType: "tunnel",
TunnelID: &tunnelNode.ID,
TunnelTargetAddr: "10.0.0.8:8080",
TunnelTargetProtocol: "http",
Enabled: true,
})
if err != nil {
t.Fatalf("CreateProxyRoute failed: %v", err)
}
if route.TunnelNodeID == nil || *route.TunnelNodeID != tunnelNode.ID {
t.Fatalf("expected legacy tunnel_id to bind tunnel_node_id, got %+v", route.TunnelNodeID)
}
result, err := PublishConfigVersion("root", false)
if err != nil {
t.Fatalf("PublishConfigVersion failed: %v", err)
}
if !strings.Contains(result.Version.RenderedConfig, "server 127.0.0.1:18080 max_fails=3 fail_timeout=10s;") {
t.Fatalf("expected rendered OpenResty upstream to use relay vhost port, got:\n%s", result.Version.RenderedConfig)
}
config, err := GetFlaredTunnelConfig(tunnelNode)
if err != nil {
t.Fatalf("GetFlaredTunnelConfig failed: %v", err)
}
if len(config.Relays) != 1 {
t.Fatalf("expected one relay, got %+v", config.Relays)
}
if config.Relays[0].Address != "de-e:17000" {
t.Fatalf("expected relay client address to include bind port, got %q", config.Relays[0].Address)
}
if len(config.Proxies) != 1 {
t.Fatalf("expected one proxy, got %+v", config.Proxies)
}
proxy := config.Proxies[0]
if proxy.LocalAddr != "10.0.0.8" || proxy.LocalPort != 8080 {
t.Fatalf("unexpected proxy target: %+v", proxy)
}
if len(proxy.CustomDomains) != 1 || proxy.CustomDomains[0] != "flared.example.com" {
t.Fatalf("unexpected proxy domains: %+v", proxy.CustomDomains)
}
}
@@ -367,9 +367,10 @@ export function NodesPage() {
)}
/>
) : node.node_type === 'tunnel_client' ? (
<span className="text-sm text-[var(--foreground-secondary)]">
-
</span>
<StatusBadge
label={node.status === 'online' ? '运行中' : '未知'}
variant={node.status === 'online' ? 'success' : 'warning'}
/>
) : (
<StatusBadge
label={getOpenrestyStatusLabel(
@@ -383,16 +384,20 @@ export function NodesPage() {
</div>
</td>
<td className="px-3 py-4 text-[var(--foreground-secondary)]">
{node.current_version || '未应用'}
{node.current_version || (node.node_type === 'tunnel_relay' ? '实时配置' : '未应用')}
</td>
<td className="px-3 py-4">
<div className="space-y-2">
<StatusBadge
label={getApplyLabel(node.latest_apply_result)}
variant={getApplyVariant(
node.latest_apply_result,
)}
/>
{node.node_type === 'tunnel_relay' ? (
<span className="text-sm text-[var(--foreground-secondary)]">—</span>
) : (
<StatusBadge
label={getApplyLabel(node.latest_apply_result)}
variant={getApplyVariant(
node.latest_apply_result,
)}
/>
)}
</div>
</td>
<td className="px-3 py-4 text-[var(--foreground-secondary)]">
@@ -561,7 +561,7 @@ function ReverseProxySection({
upstream_type: route.upstream_type || 'direct',
origin_urls_text: route.upstream_list.join('\n'),
origin_host: route.origin_host || '',
tunnel_id: route.tunnel_id ? String(route.tunnel_id) : '',
tunnel_id: route.tunnel_node_id ? String(route.tunnel_node_id) : '',
tunnel_target_addr: route.tunnel_target_addr || '',
tunnel_target_protocol: (route.tunnel_target_protocol as 'http' | 'https') || 'http',
custom_headers_text: customHeadersToText(route.custom_header_list),
@@ -574,7 +574,7 @@ function ReverseProxySection({
upstream_type: route.upstream_type || 'direct',
origin_urls_text: route.upstream_list.join('\n'),
origin_host: route.origin_host || '',
tunnel_id: route.tunnel_id ? String(route.tunnel_id) : '',
tunnel_id: route.tunnel_node_id ? String(route.tunnel_node_id) : '',
tunnel_target_addr: route.tunnel_target_addr || '',
tunnel_target_protocol: (route.tunnel_target_protocol as 'http' | 'https') || 'http',
custom_headers_text: customHeadersToText(route.custom_header_list),
@@ -632,7 +632,7 @@ function ReverseProxySection({
custom_headers: headers,
remark: values.remark.trim(),
upstream_type: values.upstream_type,
tunnel_id: values.upstream_type === 'tunnel' && values.tunnel_id ? Number(values.tunnel_id) : null,
tunnel_node_id: values.upstream_type === 'tunnel' && values.tunnel_id ? Number(values.tunnel_id) : null,
tunnel_target_addr: values.upstream_type === 'tunnel' ? values.tunnel_target_addr : '',
tunnel_target_protocol: values.upstream_type === 'tunnel' ? values.tunnel_target_protocol : '',
}),
@@ -249,7 +249,7 @@ export function ProxyRouteCreateDrawer({
basic_auth_enabled: false,
remark: values.remark.trim(),
upstream_type: values.upstream_type,
tunnel_id: values.upstream_type === 'tunnel' && values.tunnel_id ? Number(values.tunnel_id) : null,
tunnel_node_id: values.upstream_type === 'tunnel' && values.tunnel_id ? Number(values.tunnel_id) : null,
tunnel_target_addr: values.upstream_type === 'tunnel' ? values.tunnel_target_addr : '',
tunnel_target_protocol: values.upstream_type === 'tunnel' ? values.tunnel_target_protocol : '',
});
@@ -300,11 +300,20 @@ export function buildPayloadFromRoute(
basic_auth_enabled: route.basic_auth_enabled,
basic_auth_username: route.basic_auth_username,
basic_auth_password: route.basic_auth_password,
upstream_type: route.upstream_type,
tunnel_node_id: route.tunnel_node_id ?? route.tunnel_id ?? null,
tunnel_target_addr: route.tunnel_target_addr || '',
tunnel_target_protocol: route.tunnel_target_protocol || '',
...overrides,
};
}
export function getUpstreamSummary(route: ProxyRouteItem) {
if (route.upstream_type === 'tunnel') {
const protocol = route.tunnel_target_protocol || 'http';
const target = route.tunnel_target_addr || '未配置目标';
return `Tunnel → ${protocol}://${target}`;
}
if (route.upstream_list.length <= 1) {
return route.origin_url;
}
@@ -54,6 +54,7 @@ export interface ProxyRouteItem {
basic_auth_password: string;
remark: string;
upstream_type: 'direct' | 'tunnel';
tunnel_node_id?: number | null;
tunnel_id?: number | null;
tunnel_target_addr?: string;
tunnel_target_protocol?: string;
@@ -93,6 +94,7 @@ export interface ProxyRouteMutationPayload {
basic_auth_password?: string;
remark: string;
upstream_type?: 'direct' | 'tunnel';
tunnel_node_id?: number | null;
tunnel_id?: number | null;
tunnel_target_addr?: string;
tunnel_target_protocol?: string;
+12 -3
View File
@@ -6,6 +6,7 @@ import (
"encoding/json"
"fmt"
"log/slog"
"net"
"os"
"os/exec"
"path/filepath"
@@ -210,9 +211,17 @@ auth.token = "%s"
}
func parseAddr(addr string) (string, string) {
parts := strings.Split(addr, ":")
if len(parts) == 2 {
return parts[0], parts[1]
addr = strings.TrimSpace(addr)
if addr == "" {
return "127.0.0.1", "7000"
}
host, port, err := net.SplitHostPort(addr)
if err == nil {
return strings.Trim(host, "[]"), port
}
lastColon := strings.LastIndex(addr, ":")
if lastColon > 0 && strings.Count(addr, ":") == 1 {
return addr[:lastColon], addr[lastColon+1:]
}
return addr, "7000"
}