discovery.go 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261
  1. package discovery
  2. import (
  3. "context"
  4. "encoding/json"
  5. "fmt"
  6. "net"
  7. "strings"
  8. "nettool/internal/config"
  9. "nettool/internal/deviceinfo"
  10. "nettool/internal/logger"
  11. "nettool/internal/model"
  12. )
  13. type Server struct {
  14. cfg config.Config
  15. log *logger.Logger
  16. deviceSvc *deviceinfo.Service
  17. }
  18. func New(cfg config.Config, log *logger.Logger, deviceSvc *deviceinfo.Service) *Server {
  19. return &Server{cfg: cfg, log: log, deviceSvc: deviceSvc}
  20. }
  21. func (s *Server) Run(ctx context.Context) error {
  22. addr, err := net.ResolveUDPAddr("udp4", fmt.Sprintf("%s:%d", s.cfg.UDPHost, s.cfg.UDPPort))
  23. if err != nil {
  24. return err
  25. }
  26. conn, err := net.ListenUDP("udp4", addr)
  27. if err != nil {
  28. return err
  29. }
  30. defer conn.Close()
  31. go func() {
  32. <-ctx.Done()
  33. _ = conn.Close()
  34. }()
  35. s.log.Info("udp discovery listening", "addr", conn.LocalAddr().String())
  36. if err := enableLocalAddrControl(conn); err != nil {
  37. s.log.Warn("udp discovery local interface detection is unavailable", "error", err.Error())
  38. }
  39. buf := make([]byte, 2048)
  40. for {
  41. n, remote, packetInfo, err := readFromUDP(conn, buf)
  42. if err != nil {
  43. if ctx.Err() != nil {
  44. return nil
  45. }
  46. return err
  47. }
  48. var req model.DiscoverRequest
  49. if err := json.Unmarshal(buf[:n], &req); err != nil || req.MessageType != "discover" {
  50. continue
  51. }
  52. lan2IP, mac := s.maintenanceEndpoint(packetInfo)
  53. if lan2IP == "" {
  54. s.log.Warn("skip discovery response because no 169.254 maintenance address was found")
  55. continue
  56. }
  57. device := s.deviceSvc.Get()
  58. resp := model.DiscoverResponse{
  59. ProtocolVersion: 1,
  60. MessageType: "discover_response",
  61. RequestID: req.RequestID,
  62. DeviceID: device.DeviceID,
  63. Hostname: device.Hostname,
  64. ServerVersion: device.ServerVersion,
  65. MAC: mac,
  66. LAN2IP: lan2IP,
  67. HTTPPort: s.cfg.HTTPPort,
  68. AuthRequired: true,
  69. }
  70. payload, _ := json.Marshal(resp)
  71. responseSourceIP := findResponseSourceIP(packetInfo.ifIndex, remote.IP)
  72. responseTarget := remote
  73. broadcastResponse := false
  74. if responseSourceIP == nil && packetInfo.ifIndex > 0 {
  75. responseSourceIP = net.ParseIP(lan2IP)
  76. responseTarget = broadcastResponseTarget(remote)
  77. broadcastResponse = true
  78. }
  79. if _, err := writeToUDP(conn, payload, responseTarget, packetInfo, responseSourceIP); err != nil {
  80. s.log.Warn("failed to send udp discovery response", "remote", remote.String(), "target", responseTarget.String(), "source_ip", responseSourceIP.String(), "maintenance_ip", lan2IP, "interface_index", packetInfo.ifIndex, "broadcast", broadcastResponse, "error", err.Error())
  81. }
  82. }
  83. }
  84. func broadcastResponseTarget(remote *net.UDPAddr) *net.UDPAddr {
  85. return &net.UDPAddr{IP: net.IPv4bcast, Port: remote.Port}
  86. }
  87. type udpPacketInfo struct {
  88. localIP net.IP
  89. ifIndex int
  90. }
  91. func (s *Server) maintenanceEndpoint(packetInfo udpPacketInfo) (string, string) {
  92. if packetInfo.ifIndex > 0 {
  93. lan2IP, mac := findLinkLocalEndpointByInterfaceIndex(packetInfo.ifIndex)
  94. if lan2IP != "" {
  95. return lan2IP, mac
  96. }
  97. return "", ""
  98. }
  99. if packetInfo.localIP != nil {
  100. lan2IP, mac := findLinkLocalEndpointByInterfaceIP(packetInfo.localIP.String())
  101. if lan2IP != "" {
  102. return lan2IP, mac
  103. }
  104. }
  105. if s.cfg.MaintenanceIP != "" {
  106. mac := findMACByIP(s.cfg.MaintenanceIP)
  107. if mac != "" {
  108. return s.cfg.MaintenanceIP, mac
  109. }
  110. }
  111. return findFirstLinkLocalEndpoint()
  112. }
  113. func findResponseSourceIP(interfaceIndex int, remoteIP net.IP) net.IP {
  114. if interfaceIndex <= 0 {
  115. return nil
  116. }
  117. iface, err := net.InterfaceByIndex(interfaceIndex)
  118. if err != nil {
  119. return nil
  120. }
  121. addresses, err := iface.Addrs()
  122. if err != nil {
  123. return nil
  124. }
  125. return bestSourceIPForRemote(addresses, remoteIP)
  126. }
  127. func bestSourceIPForRemote(addresses []net.Addr, remoteIP net.IP) net.IP {
  128. remoteIPv4 := remoteIP.To4()
  129. if remoteIPv4 == nil {
  130. return nil
  131. }
  132. var best net.IP
  133. bestPrefix := -1
  134. for _, address := range addresses {
  135. ipNet, ok := address.(*net.IPNet)
  136. if !ok || ipNet.IP.To4() == nil || !ipNet.Contains(remoteIPv4) {
  137. continue
  138. }
  139. prefix, bits := ipNet.Mask.Size()
  140. if bits != 32 || prefix <= bestPrefix {
  141. continue
  142. }
  143. best = ipNet.IP.To4()
  144. bestPrefix = prefix
  145. }
  146. return best
  147. }
  148. func findLinkLocalEndpointByInterfaceIndex(index int) (string, string) {
  149. iface, err := net.InterfaceByIndex(index)
  150. if err != nil {
  151. return "", ""
  152. }
  153. return findLinkLocalEndpointOnInterface(*iface)
  154. }
  155. func findLinkLocalEndpointByInterfaceIP(ip string) (string, string) {
  156. ifaces, err := net.Interfaces()
  157. if err != nil {
  158. return "", ""
  159. }
  160. for _, iface := range ifaces {
  161. lan2IP, mac := findLinkLocalEndpointOnInterface(iface)
  162. if lan2IP == ip {
  163. return lan2IP, mac
  164. }
  165. }
  166. return "", ""
  167. }
  168. func findFirstLinkLocalEndpoint() (string, string) {
  169. ifaces, err := net.Interfaces()
  170. if err != nil {
  171. return "", ""
  172. }
  173. for _, iface := range ifaces {
  174. lan2IP, mac := findLinkLocalEndpointOnInterface(iface)
  175. if lan2IP != "" {
  176. return lan2IP, mac
  177. }
  178. }
  179. return "", ""
  180. }
  181. func findLinkLocalEndpointOnInterface(iface net.Interface) (string, string) {
  182. if iface.Flags&net.FlagLoopback != 0 || iface.Flags&net.FlagUp == 0 || len(iface.HardwareAddr) == 0 {
  183. return "", ""
  184. }
  185. addrs, err := iface.Addrs()
  186. if err != nil {
  187. return "", ""
  188. }
  189. for _, addr := range addrs {
  190. current := ipv4FromAddr(addr)
  191. if current == nil || !strings.HasPrefix(current.String(), "169.254.") {
  192. continue
  193. }
  194. return current.String(), iface.HardwareAddr.String()
  195. }
  196. return "", ""
  197. }
  198. func findMACByIP(ip string) string {
  199. ifaces, err := net.Interfaces()
  200. if err != nil {
  201. return ""
  202. }
  203. var fallback string
  204. for _, iface := range ifaces {
  205. if iface.Flags&net.FlagLoopback != 0 || len(iface.HardwareAddr) == 0 {
  206. continue
  207. }
  208. addrs, err := iface.Addrs()
  209. if err != nil {
  210. continue
  211. }
  212. for _, addr := range addrs {
  213. current := ipv4FromAddr(addr)
  214. if current == nil {
  215. continue
  216. }
  217. if current.String() == ip {
  218. return iface.HardwareAddr.String()
  219. }
  220. if fallback == "" && strings.HasPrefix(current.String(), "169.254.") {
  221. fallback = iface.HardwareAddr.String()
  222. }
  223. }
  224. }
  225. return fallback
  226. }
  227. func ipv4FromAddr(addr net.Addr) net.IP {
  228. var current net.IP
  229. switch value := addr.(type) {
  230. case *net.IPNet:
  231. current = value.IP
  232. case *net.IPAddr:
  233. current = value.IP
  234. }
  235. if current == nil {
  236. return nil
  237. }
  238. return current.To4()
  239. }