diff --git a/openflare_agent/internal/wsclient/client.go b/openflare_agent/internal/wsclient/client.go index 132d2c7c..808294b9 100644 --- a/openflare_agent/internal/wsclient/client.go +++ b/openflare_agent/internal/wsclient/client.go @@ -2,165 +2,74 @@ package wsclient import ( "context" - "errors" - "log/slog" - "net" - "net/http" - "net/url" - "strings" "time" - "golang.org/x/net/websocket" - "openflare-agent/internal/protocol" + shared "openflare/utils/wsclient" ) type Client struct { - baseURL string - token string - timeout time.Duration + sharedClient *shared.Client } type Connection struct { - conn *websocket.Conn - url string - readTimeout time.Duration + sharedConn *shared.Connection } func New(baseURL string, token string, timeout time.Duration) *Client { return &Client{ - baseURL: strings.TrimRight(baseURL, "/"), - token: strings.TrimSpace(token), - timeout: timeout, + sharedClient: shared.New(shared.Config{ + BaseURL: baseURL, + Token: token, + Timeout: timeout, + HeaderKey: "X-Agent-Token", + WSPath: "/api/agent/ws", + }), } } func (c *Client) SetToken(token string) { - c.token = strings.TrimSpace(token) - slog.Debug("agent ws client token updated") + c.sharedClient.SetToken(token) } func (c *Client) URL() string { - wsURL, err := buildWebsocketURL(c.baseURL) - if err != nil { - return "" - } - return wsURL + return c.sharedClient.URL() } func (c *Client) Connect(ctx context.Context) (protocol.WebSocketConnection, error) { - wsURL, err := buildWebsocketURL(c.baseURL) + conn, err := c.sharedClient.Connect(ctx) if err != nil { return nil, err } - if strings.TrimSpace(c.token) == "" { - return nil, errors.New("agent ws token is empty") - } - origin := strings.TrimSpace(c.baseURL) - if origin == "" { - origin = "http://localhost" - } - config, err := websocket.NewConfig(wsURL, origin) - if err != nil { - return nil, err - } - config.Header = http.Header{} - config.Header.Set("X-Agent-Token", c.token) - if c.timeout > 0 { - config.Dialer = &net.Dialer{Timeout: c.timeout} - } - slog.Debug("agent ws dialing server", "url", wsURL) - conn, err := config.DialContext(ctx) - if err != nil { - return nil, err - } - slog.Debug("agent ws dial succeeded", "url", wsURL) - return &Connection{conn: conn, url: wsURL, readTimeout: websocketReadTimeout(c.timeout)}, nil -} - -func buildWebsocketURL(baseURL string) (string, error) { - parsed, err := url.Parse(strings.TrimRight(baseURL, "/")) - if err != nil { - return "", err - } - switch parsed.Scheme { - case "http": - parsed.Scheme = "ws" - case "https": - parsed.Scheme = "wss" - case "ws", "wss": - default: - return "", errors.New("server_url scheme must be http, https, ws, or wss") - } - parsed.Path = strings.TrimRight(parsed.Path, "/") + "/api/agent/ws" - parsed.RawQuery = "" - parsed.Fragment = "" - return parsed.String(), nil + return &Connection{sharedConn: conn}, nil } func (conn *Connection) URL() string { - if conn == nil { + if conn == nil || conn.sharedConn == nil { return "" } - return conn.url + return conn.sharedConn.URL } func (conn *Connection) SendStatus(payload protocol.NodePayload) error { - if conn == nil || conn.conn == nil { - return errors.New("agent ws connection is nil") - } - slog.Debug("agent ws sending status", - "node_id", payload.NodeID, - "current_version", payload.CurrentVersion, - "openresty_status", payload.OpenrestyStatus, - ) - return websocket.JSON.Send(conn.conn, protocol.WSOutboundMessage{ - Type: protocol.WSMessageTypeStatus, - Payload: payload, - }) + return conn.sharedConn.SendMessage(protocol.WSMessageTypeStatus, payload) } func (conn *Connection) SendPong() error { - if conn == nil || conn.conn == nil { - return errors.New("agent ws connection is nil") - } - slog.Debug("agent ws sending pong") - return websocket.JSON.Send(conn.conn, protocol.WSOutboundMessage{ - Type: protocol.WSMessageTypePong, - }) + return conn.sharedConn.SendMessage(protocol.WSMessageTypePong, nil) } func (conn *Connection) Receive() (protocol.WSMessage, error) { var message protocol.WSMessage - if conn == nil || conn.conn == nil { - return message, errors.New("agent ws connection is nil") - } - if conn.readTimeout > 0 { - _ = conn.conn.SetReadDeadline(time.Now().Add(conn.readTimeout)) - } - err := websocket.JSON.Receive(conn.conn, &message) - if err != nil { - var netErr net.Error - if errors.As(err, &netErr) && netErr.Timeout() { - slog.Debug("agent ws receive timeout waiting for server message", "timeout", conn.readTimeout) - } + if err := conn.sharedConn.Receive(&message); err != nil { return message, err } - slog.Debug("agent ws received message", "type", message.Type) return message, nil } -func websocketReadTimeout(requestTimeout time.Duration) time.Duration { - timeout := requestTimeout * 6 - if timeout < 75*time.Second { - return 75 * time.Second - } - return timeout -} - func (conn *Connection) Close() error { - if conn == nil || conn.conn == nil { + if conn == nil || conn.sharedConn == nil { return nil } - return conn.conn.Close() + return conn.sharedConn.Close() } diff --git a/openflare_relay/internal/wsclient/client.go b/openflare_relay/internal/wsclient/client.go index 4b7e1f55..434ee580 100644 --- a/openflare_relay/internal/wsclient/client.go +++ b/openflare_relay/internal/wsclient/client.go @@ -3,153 +3,66 @@ package wsclient import ( "context" "encoding/json" - "errors" - "log/slog" - "net" - "net/http" - "net/url" - "strings" "time" - "golang.org/x/net/websocket" "openflare/service" + shared "openflare/utils/wsclient" ) type Client struct { - baseURL string - token string - timeout time.Duration + sharedClient *shared.Client } type Connection struct { - conn *websocket.Conn - url string - readTimeout time.Duration + sharedConn *shared.Connection } func New(baseURL string, token string, timeout time.Duration) *Client { return &Client{ - baseURL: strings.TrimRight(baseURL, "/"), - token: strings.TrimSpace(token), - timeout: timeout, + sharedClient: shared.New(shared.Config{ + BaseURL: baseURL, + Token: token, + Timeout: timeout, + HeaderKey: "X-Agent-Token", + WSPath: "/api/relay/ws", + }), } } func (c *Client) SetToken(token string) { - c.token = strings.TrimSpace(token) - slog.Debug("relay ws client token updated") + c.sharedClient.SetToken(token) } func (c *Client) Connect(ctx context.Context) (*Connection, error) { - wsURL, err := buildWebsocketURL(c.baseURL) + conn, err := c.sharedClient.Connect(ctx) if err != nil { return nil, err } - if strings.TrimSpace(c.token) == "" { - return nil, errors.New("relay ws token is empty") - } - origin := strings.TrimSpace(c.baseURL) - if origin == "" { - origin = "http://localhost" - } - config, err := websocket.NewConfig(wsURL, origin) - if err != nil { - return nil, err - } - config.Header = http.Header{} - config.Header.Set("X-Agent-Token", c.token) - if c.timeout > 0 { - config.Dialer = &net.Dialer{Timeout: c.timeout} - } - slog.Debug("relay ws dialing server", "url", wsURL) - conn, err := config.DialContext(ctx) - if err != nil { - return nil, err - } - slog.Debug("relay ws dial succeeded", "url", wsURL) - return &Connection{conn: conn, url: wsURL, readTimeout: websocketReadTimeout(c.timeout)}, nil -} - -func buildWebsocketURL(baseURL string) (string, error) { - parsed, err := url.Parse(strings.TrimRight(baseURL, "/")) - if err != nil { - return "", err - } - switch parsed.Scheme { - case "http": - parsed.Scheme = "ws" - case "https": - parsed.Scheme = "wss" - case "ws", "wss": - default: - return "", errors.New("server_url scheme must be http, https, ws, or wss") - } - parsed.Path = strings.TrimRight(parsed.Path, "/") + "/api/relay/ws" - parsed.RawQuery = "" - parsed.Fragment = "" - return parsed.String(), nil + return &Connection{sharedConn: conn}, nil } func (conn *Connection) SendPing() error { - if conn == nil || conn.conn == nil { - return errors.New("relay ws connection is nil") - } - slog.Debug("relay ws sending ping") - _ = conn.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) - return websocket.JSON.Send(conn.conn, service.WSMessage{ - Type: "ping", - }) + return conn.sharedConn.SendMessage("ping", nil) } func (conn *Connection) SendPong() error { - if conn == nil || conn.conn == nil { - return errors.New("relay ws connection is nil") - } - slog.Debug("relay ws sending pong") - _ = conn.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) - return websocket.JSON.Send(conn.conn, service.WSMessage{ - Type: "pong", - }) + return conn.sharedConn.SendMessage("pong", nil) } func (conn *Connection) Receive() (service.WSMessage, error) { - var message service.WSMessage - if conn == nil || conn.conn == nil { - return message, errors.New("relay ws connection is nil") - } - if conn.readTimeout > 0 { - _ = conn.conn.SetReadDeadline(time.Now().Add(conn.readTimeout)) - } - // Use custom json unmarshaling to handle any type var raw struct { Type string `json:"type"` Payload json.RawMessage `json:"payload,omitempty"` } - err := websocket.JSON.Receive(conn.conn, &raw) - if err != nil { - var netErr net.Error - if errors.As(err, &netErr) && netErr.Timeout() { - slog.Debug("relay ws receive timeout waiting for server message", "timeout", conn.readTimeout) - } - return message, err + if err := conn.sharedConn.Receive(&raw); err != nil { + return service.WSMessage{}, err } - message.Type = raw.Type - message.Payload = raw.Payload - slog.Debug("relay ws received message", "type", message.Type) - return message, nil -} - -func websocketReadTimeout(requestTimeout time.Duration) time.Duration { - timeout := requestTimeout * 6 - if timeout < 75*time.Second { - return 75 * time.Second - } - return timeout + return service.WSMessage{ + Type: raw.Type, + Payload: raw.Payload, + }, nil } func (conn *Connection) Close() error { - if conn == nil || conn.conn == nil { - return nil - } - return conn.conn.Close() + return conn.sharedConn.Close() } diff --git a/openflare_server/utils/wsclient/client.go b/openflare_server/utils/wsclient/client.go new file mode 100644 index 00000000..fd4dc48f --- /dev/null +++ b/openflare_server/utils/wsclient/client.go @@ -0,0 +1,163 @@ +package wsclient + +import ( + "context" + "errors" + "log/slog" + "net" + "net/http" + "net/url" + "strings" + "time" + + "golang.org/x/net/websocket" +) + +type Config struct { + BaseURL string + Token string + Timeout time.Duration + HeaderKey string // e.g. "X-Agent-Token", "X-Tunnel-Token" + WSPath string // e.g. "/api/relay/ws", "/api/agent/ws", "/api/flared/ws" +} + +type Client struct { + cfg Config +} + +type Connection struct { + Conn *websocket.Conn + URL string + ReadTimeout time.Duration +} + +func New(cfg Config) *Client { + cfg.BaseURL = strings.TrimRight(cfg.BaseURL, "/") + cfg.Token = strings.TrimSpace(cfg.Token) + cfg.HeaderKey = strings.TrimSpace(cfg.HeaderKey) + cfg.WSPath = strings.TrimSpace(cfg.WSPath) + return &Client{ + cfg: cfg, + } +} + +func (c *Client) SetToken(token string) { + c.cfg.Token = strings.TrimSpace(token) +} + +func (c *Client) URL() string { + wsURL, err := c.BuildWebsocketURL() + if err != nil { + return "" + } + return wsURL +} + +func (c *Client) BuildWebsocketURL() (string, error) { + parsed, err := url.Parse(c.cfg.BaseURL) + if err != nil { + return "", err + } + switch parsed.Scheme { + case "http": + parsed.Scheme = "ws" + case "https": + parsed.Scheme = "wss" + case "ws", "wss": + default: + return "", errors.New("server_url scheme must be http, https, ws, or wss") + } + + wsPath := c.cfg.WSPath + if !strings.HasPrefix(wsPath, "/") { + wsPath = "/" + wsPath + } + parsed.Path = strings.TrimRight(parsed.Path, "/") + wsPath + parsed.RawQuery = "" + parsed.Fragment = "" + return parsed.String(), nil +} + +func (c *Client) Connect(ctx context.Context) (*Connection, error) { + wsURL, err := c.BuildWebsocketURL() + if err != nil { + return nil, err + } + if c.cfg.Token == "" { + return nil, errors.New("ws token is empty") + } + origin := c.cfg.BaseURL + if origin == "" { + origin = "http://localhost" + } + config, err := websocket.NewConfig(wsURL, origin) + if err != nil { + return nil, err + } + config.Header = http.Header{} + if c.cfg.HeaderKey != "" { + config.Header.Set(c.cfg.HeaderKey, c.cfg.Token) + } + if c.cfg.Timeout > 0 { + config.Dialer = &net.Dialer{Timeout: c.cfg.Timeout} + } + slog.Debug("ws dialing server", "url", wsURL) + conn, err := config.DialContext(ctx) + if err != nil { + return nil, err + } + slog.Debug("ws dial succeeded", "url", wsURL) + return &Connection{Conn: conn, URL: wsURL, ReadTimeout: websocketReadTimeout(c.cfg.Timeout)}, nil +} + +func (conn *Connection) SendMessage(msgType string, payload any) error { + if conn == nil || conn.Conn == nil { + return errors.New("ws connection is nil") + } + slog.Debug("ws sending message", "type", msgType) + + // Create the outbound message wrapper + message := struct { + Type string `json:"type"` + Payload any `json:"payload,omitempty"` + }{ + Type: msgType, + Payload: payload, + } + + _ = conn.Conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) + return websocket.JSON.Send(conn.Conn, message) +} + +func (conn *Connection) Receive(target any) error { + if conn == nil || conn.Conn == nil { + return errors.New("ws connection is nil") + } + if conn.ReadTimeout > 0 { + _ = conn.Conn.SetReadDeadline(time.Now().Add(conn.ReadTimeout)) + } + err := websocket.JSON.Receive(conn.Conn, target) + if err != nil { + var netErr net.Error + if errors.As(err, &netErr) && netErr.Timeout() { + slog.Debug("ws receive timeout waiting for server message", "timeout", conn.ReadTimeout) + } + return err + } + return nil +} + +func websocketReadTimeout(requestTimeout time.Duration) time.Duration { + timeout := requestTimeout * 6 + if timeout < 75*time.Second { + return 75 * time.Second + } + return timeout +} + +func (conn *Connection) Close() error { + if conn == nil || conn.Conn == nil { + return nil + } + return conn.Conn.Close() +} diff --git a/openflared/internal/wsclient/client.go b/openflared/internal/wsclient/client.go index 0a2477de..e567621d 100644 --- a/openflared/internal/wsclient/client.go +++ b/openflared/internal/wsclient/client.go @@ -3,142 +3,62 @@ package wsclient import ( "context" "encoding/json" - "errors" - "log/slog" - "net" - "net/http" - "net/url" - "strings" "time" - "golang.org/x/net/websocket" "openflare/service" + shared "openflare/utils/wsclient" ) type Client struct { - baseURL string - token string - timeout time.Duration + sharedClient *shared.Client } type Connection struct { - conn *websocket.Conn - url string - readTimeout time.Duration + sharedConn *shared.Connection } func New(baseURL string, token string, timeout time.Duration) *Client { return &Client{ - baseURL: strings.TrimRight(baseURL, "/"), - token: strings.TrimSpace(token), - timeout: timeout, + sharedClient: shared.New(shared.Config{ + BaseURL: baseURL, + Token: token, + Timeout: timeout, + HeaderKey: "X-Tunnel-Token", + WSPath: "/api/flared/ws", + }), } } func (c *Client) SetToken(token string) { - c.token = strings.TrimSpace(token) - slog.Debug("flared ws client token updated") + c.sharedClient.SetToken(token) } func (c *Client) Connect(ctx context.Context) (*Connection, error) { - wsURL, err := buildWebsocketURL(c.baseURL) + conn, err := c.sharedClient.Connect(ctx) if err != nil { return nil, err } - if strings.TrimSpace(c.token) == "" { - return nil, errors.New("flared ws token is empty") - } - origin := strings.TrimSpace(c.baseURL) - if origin == "" { - origin = "http://localhost" - } - config, err := websocket.NewConfig(wsURL, origin) - if err != nil { - return nil, err - } - config.Header = http.Header{} - config.Header.Set("X-Tunnel-Token", c.token) - if c.timeout > 0 { - config.Dialer = &net.Dialer{Timeout: c.timeout} - } - slog.Debug("flared ws dialing server", "url", wsURL) - conn, err := config.DialContext(ctx) - if err != nil { - return nil, err - } - slog.Debug("flared ws dial succeeded", "url", wsURL) - return &Connection{conn: conn, url: wsURL, readTimeout: websocketReadTimeout(c.timeout)}, nil -} - -func buildWebsocketURL(baseURL string) (string, error) { - parsed, err := url.Parse(strings.TrimRight(baseURL, "/")) - if err != nil { - return "", err - } - switch parsed.Scheme { - case "http": - parsed.Scheme = "ws" - case "https": - parsed.Scheme = "wss" - case "ws", "wss": - default: - return "", errors.New("server_url scheme must be http, https, ws, or wss") - } - parsed.Path = strings.TrimRight(parsed.Path, "/") + "/api/flared/ws" - parsed.RawQuery = "" - parsed.Fragment = "" - return parsed.String(), nil + return &Connection{sharedConn: conn}, nil } func (conn *Connection) SendPong() error { - if conn == nil || conn.conn == nil { - return errors.New("flared ws connection is nil") - } - slog.Debug("flared ws sending pong") - _ = conn.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)) - return websocket.JSON.Send(conn.conn, service.WSMessage{ - Type: "pong", - }) + return conn.sharedConn.SendMessage("pong", nil) } func (conn *Connection) Receive() (service.WSMessage, error) { - var message service.WSMessage - if conn == nil || conn.conn == nil { - return message, errors.New("flared ws connection is nil") - } - if conn.readTimeout > 0 { - _ = conn.conn.SetReadDeadline(time.Now().Add(conn.readTimeout)) - } - // Use custom json unmarshaling to handle any type var raw struct { Type string `json:"type"` Payload json.RawMessage `json:"payload,omitempty"` } - err := websocket.JSON.Receive(conn.conn, &raw) - if err != nil { - var netErr net.Error - if errors.As(err, &netErr) && netErr.Timeout() { - slog.Debug("flared ws receive timeout waiting for server message", "timeout", conn.readTimeout) - } - return message, err + if err := conn.sharedConn.Receive(&raw); err != nil { + return service.WSMessage{}, err } - message.Type = raw.Type - message.Payload = raw.Payload - slog.Debug("flared ws received message", "type", message.Type) - return message, nil -} - -func websocketReadTimeout(requestTimeout time.Duration) time.Duration { - timeout := requestTimeout * 6 - if timeout < 75*time.Second { - return 75 * time.Second - } - return timeout + return service.WSMessage{ + Type: raw.Type, + Payload: raw.Payload, + }, nil } func (conn *Connection) Close() error { - if conn == nil || conn.conn == nil { - return nil - } - return conn.conn.Close() + return conn.sharedConn.Close() }