fix: correct tunnel protocol handling for KCP cleanup and diagnosis

- Fix KCP tunnel not reclaimed after deletion: service name was
  hardcoded to {id}_tls but services were created as {id}_kcp,
  causing DeleteService to never find the actual KCP service.
  Now reads tunnel.Protocol from DB and derives correct name.

- Fix KCP diagnosis using TCP ping instead of UDP ping:
  tunnel.Protocol was always hardcoded to 'tls' at creation,
  so isUDPBasedProtocol() never matched kcp tunnels. Now
  stores the actual protocol from entry node configuration.

- Fix addTunnelServiceOnNode to extract service name from
  serviceData instead of hardcoding _tls suffix.

- Fix rollbackTunnelRuntime to accept protocol parameter
  so retry cleanup uses correct service name.

- UpdateTunnelTx now persists protocol on tunnel updates.
This commit is contained in:
sagitchu
2026-04-21 18:48:11 +08:00
parent 5107f59d94
commit c431d79403
2 changed files with 33 additions and 7 deletions
+31 -6
View File
@@ -657,11 +657,15 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
if trimmed := strings.TrimSpace(inIP); trimmed != "" { if trimmed := strings.TrimSpace(inIP); trimmed != "" {
tunnelInIP = sql.NullString{String: trimmed, Valid: true} tunnelInIP = sql.NullString{String: trimmed, Valid: true}
} }
tunnelProtocol := "tls"
if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
tunnelProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
}
tunnel := model.Tunnel{ tunnel := model.Tunnel{
Name: name, Name: name,
TrafficRatio: trafficRatio, TrafficRatio: trafficRatio,
Type: typeVal, Type: typeVal,
Protocol: "tls", Protocol: tunnelProtocol,
Flow: flow, Flow: flow,
CreatedTime: now, CreatedTime: now,
UpdatedTime: now, UpdatedTime: now,
@@ -702,7 +706,7 @@ func (h *Handler) tunnelCreate(w http.ResponseWriter, r *http.Request) {
if typeVal == 2 { if typeVal == 2 {
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState) createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
if applyErr != nil { if applyErr != nil {
h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID) h.rollbackTunnelRuntime(createdChains, createdServices, tunnelID, tunnelProtocol)
h.releaseFederationRuntimeRefs(federationReleaseRefs) h.releaseFederationRuntimeRefs(federationReleaseRefs)
_ = h.deleteTunnelByID(tunnelID) _ = h.deleteTunnelByID(tunnelID)
response.WriteJSON(w, response.ErrDefault(applyErr.Error())) response.WriteJSON(w, response.ErrDefault(applyErr.Error()))
@@ -722,7 +726,11 @@ func (h *Handler) cleanupTunnelRuntime(tunnelID int64) {
return return
} }
serviceName := fmt.Sprintf("%d_tls", tunnelID) protocol := strings.TrimSpace(tunnel.Protocol)
if protocol == "" {
protocol = "tls"
}
serviceName := fmt.Sprintf("%d_%s", tunnelID, protocol)
chainName := fmt.Sprintf("chains_%d", tunnelID) chainName := fmt.Sprintf("chains_%d", tunnelID)
for _, row := range chainRows { for _, row := range chainRows {
@@ -816,6 +824,10 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
} }
defer func() { tx.Rollback() }() defer func() { tx.Rollback() }()
updateProtocol := "tls"
if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
updateProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
}
if err := h.repo.UpdateTunnelTx( if err := h.repo.UpdateTunnelTx(
tx, tx,
id, id,
@@ -826,6 +838,7 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
asInt(req["status"], 1), asInt(req["status"], 1),
inIp, inIp,
ipPreference, ipPreference,
updateProtocol,
now, now,
); err != nil { ); err != nil {
response.WriteJSON(w, response.Err(-2, err.Error())) response.WriteJSON(w, response.Err(-2, err.Error()))
@@ -873,7 +886,11 @@ func (h *Handler) tunnelUpdate(w http.ResponseWriter, r *http.Request) {
if typeVal == 2 { if typeVal == 2 {
createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState) createdChains, createdServices, applyErr := h.applyTunnelRuntime(runtimeState)
if applyErr != nil { if applyErr != nil {
h.rollbackTunnelRuntime(createdChains, createdServices, id) updateProtocol := "tls"
if len(runtimeState.InNodes) > 0 && strings.TrimSpace(runtimeState.InNodes[0].Protocol) != "" {
updateProtocol = strings.TrimSpace(runtimeState.InNodes[0].Protocol)
}
h.rollbackTunnelRuntime(createdChains, createdServices, id, updateProtocol)
h.releaseFederationRuntimeRefs(federationReleaseRefs) h.releaseFederationRuntimeRefs(federationReleaseRefs)
_ = h.repo.DeleteFederationTunnelBindingsByTunnel(id) _ = h.repo.DeleteFederationTunnelBindingsByTunnel(id)
if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) { if len(federationReleaseRefs) == 0 && shouldDeferTunnelRuntimeApplyError(applyErr) {
@@ -3362,6 +3379,11 @@ func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []m
return errors.New("invalid tunnel service context") return errors.New("invalid tunnel service context")
} }
serviceName := fmt.Sprintf("%d_tls", tunnelID) serviceName := fmt.Sprintf("%d_tls", tunnelID)
if len(serviceData) > 0 {
if name, ok := serviceData[0]["name"].(string); ok && strings.TrimSpace(name) != "" {
serviceName = strings.TrimSpace(name)
}
}
return retryTunnelServiceAddWithCleanup( return retryTunnelServiceAddWithCleanup(
func() error { func() error {
_, err := h.sendNodeCommand(nodeID, "AddService", serviceData, true, false) _, err := h.sendNodeCommand(nodeID, "AddService", serviceData, true, false)
@@ -3375,12 +3397,15 @@ func (h *Handler) addTunnelServiceOnNode(nodeID, tunnelID int64, serviceData []m
) )
} }
func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tunnelID int64) { func (h *Handler) rollbackTunnelRuntime(chainNodeIDs, serviceNodeIDs []int64, tunnelID int64, protocol string) {
if h == nil || tunnelID <= 0 { if h == nil || tunnelID <= 0 {
return return
} }
if protocol == "" {
protocol = "tls"
}
seenServices := make(map[int64]struct{}) seenServices := make(map[int64]struct{})
serviceName := fmt.Sprintf("%d_tls", tunnelID) serviceName := fmt.Sprintf("%d_%s", tunnelID, protocol)
for i := len(serviceNodeIDs) - 1; i >= 0; i-- { for i := len(serviceNodeIDs) - 1; i >= 0; i-- {
nodeID := serviceNodeIDs[i] nodeID := serviceNodeIDs[i]
if _, ok := seenServices[nodeID]; ok { if _, ok := seenServices[nodeID]; ok {
@@ -394,7 +394,7 @@ func (r *Repository) UpdateTunnelOrder(tunnelID int64, inx int, now int64) {
Updates(map[string]interface{}{"inx": inx, "updated_time": now}).Error Updates(map[string]interface{}{"inx": inx, "updated_time": now}).Error
} }
func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, typeVal int, flow int64, trafficRatio float64, status int, inIP, ipPreference string, now int64) error { func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, typeVal int, flow int64, trafficRatio float64, status int, inIP, ipPreference string, protocol string, now int64) error {
if tx == nil { if tx == nil {
return errors.New("database unavailable") return errors.New("database unavailable")
} }
@@ -408,6 +408,7 @@ func (r *Repository) UpdateTunnelTx(tx *gorm.DB, tunnelID int64, name string, ty
"status": status, "status": status,
"in_ip": nullStringFromInterface(inIP), "in_ip": nullStringFromInterface(inIP),
"ip_preference": ipPreference, "ip_preference": ipPreference,
"protocol": protocol,
"updated_time": now, "updated_time": now,
}).Error }).Error
} }