mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 17:56:37 +08:00
[优化] 重构 WebSocket 客户端,整合共享连接逻辑并简化代码
This commit is contained in:
@@ -2,165 +2,74 @@ package wsclient
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
|
||||||
"log/slog"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/url"
|
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"golang.org/x/net/websocket"
|
|
||||||
|
|
||||||
"openflare-agent/internal/protocol"
|
"openflare-agent/internal/protocol"
|
||||||
|
shared "openflare/utils/wsclient"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Client struct {
|
type Client struct {
|
||||||
baseURL string
|
sharedClient *shared.Client
|
||||||
token string
|
|
||||||
timeout time.Duration
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type Connection struct {
|
type Connection struct {
|
||||||
conn *websocket.Conn
|
sharedConn *shared.Connection
|
||||||
url string
|
|
||||||
readTimeout time.Duration
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||||
return &Client{
|
return &Client{
|
||||||
baseURL: strings.TrimRight(baseURL, "/"),
|
sharedClient: shared.New(shared.Config{
|
||||||
token: strings.TrimSpace(token),
|
BaseURL: baseURL,
|
||||||
timeout: timeout,
|
Token: token,
|
||||||
|
Timeout: timeout,
|
||||||
|
HeaderKey: "X-Agent-Token",
|
||||||
|
WSPath: "/api/agent/ws",
|
||||||
|
}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) SetToken(token string) {
|
func (c *Client) SetToken(token string) {
|
||||||
c.token = strings.TrimSpace(token)
|
c.sharedClient.SetToken(token)
|
||||||
slog.Debug("agent ws client token updated")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) URL() string {
|
func (c *Client) URL() string {
|
||||||
wsURL, err := buildWebsocketURL(c.baseURL)
|
return c.sharedClient.URL()
|
||||||
if err != nil {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
return wsURL
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) Connect(ctx context.Context) (protocol.WebSocketConnection, error) {
|
func (c *Client) Connect(ctx context.Context) (protocol.WebSocketConnection, error) {
|
||||||
wsURL, err := buildWebsocketURL(c.baseURL)
|
conn, err := c.sharedClient.Connect(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(c.token) == "" {
|
return &Connection{sharedConn: conn}, nil
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) URL() string {
|
func (conn *Connection) URL() string {
|
||||||
if conn == nil {
|
if conn == nil || conn.sharedConn == nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
return conn.url
|
return conn.sharedConn.URL
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) SendStatus(payload protocol.NodePayload) error {
|
func (conn *Connection) SendStatus(payload protocol.NodePayload) error {
|
||||||
if conn == nil || conn.conn == nil {
|
return conn.sharedConn.SendMessage(protocol.WSMessageTypeStatus, payload)
|
||||||
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,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) SendPong() error {
|
func (conn *Connection) SendPong() error {
|
||||||
if conn == nil || conn.conn == nil {
|
return conn.sharedConn.SendMessage(protocol.WSMessageTypePong, 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,
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) Receive() (protocol.WSMessage, error) {
|
func (conn *Connection) Receive() (protocol.WSMessage, error) {
|
||||||
var message protocol.WSMessage
|
var message protocol.WSMessage
|
||||||
if conn == nil || conn.conn == nil {
|
if err := conn.sharedConn.Receive(&message); err != 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)
|
|
||||||
}
|
|
||||||
return message, err
|
return message, err
|
||||||
}
|
}
|
||||||
slog.Debug("agent ws received message", "type", message.Type)
|
|
||||||
return message, nil
|
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 {
|
func (conn *Connection) Close() error {
|
||||||
if conn == nil || conn.conn == nil {
|
if conn == nil || conn.sharedConn == nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return conn.conn.Close()
|
return conn.sharedConn.Close()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -3,153 +3,66 @@ package wsclient
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
|
||||||
"log/slog"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/url"
|
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"golang.org/x/net/websocket"
|
|
||||||
"openflare/service"
|
"openflare/service"
|
||||||
|
shared "openflare/utils/wsclient"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Client struct {
|
type Client struct {
|
||||||
baseURL string
|
sharedClient *shared.Client
|
||||||
token string
|
|
||||||
timeout time.Duration
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type Connection struct {
|
type Connection struct {
|
||||||
conn *websocket.Conn
|
sharedConn *shared.Connection
|
||||||
url string
|
|
||||||
readTimeout time.Duration
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||||
return &Client{
|
return &Client{
|
||||||
baseURL: strings.TrimRight(baseURL, "/"),
|
sharedClient: shared.New(shared.Config{
|
||||||
token: strings.TrimSpace(token),
|
BaseURL: baseURL,
|
||||||
timeout: timeout,
|
Token: token,
|
||||||
|
Timeout: timeout,
|
||||||
|
HeaderKey: "X-Agent-Token",
|
||||||
|
WSPath: "/api/relay/ws",
|
||||||
|
}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) SetToken(token string) {
|
func (c *Client) SetToken(token string) {
|
||||||
c.token = strings.TrimSpace(token)
|
c.sharedClient.SetToken(token)
|
||||||
slog.Debug("relay ws client token updated")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||||
wsURL, err := buildWebsocketURL(c.baseURL)
|
conn, err := c.sharedClient.Connect(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(c.token) == "" {
|
return &Connection{sharedConn: conn}, nil
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) SendPing() error {
|
func (conn *Connection) SendPing() error {
|
||||||
if conn == nil || conn.conn == nil {
|
return conn.sharedConn.SendMessage("ping", 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",
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) SendPong() error {
|
func (conn *Connection) SendPong() error {
|
||||||
if conn == nil || conn.conn == nil {
|
return conn.sharedConn.SendMessage("pong", 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",
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) Receive() (service.WSMessage, error) {
|
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 {
|
var raw struct {
|
||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
Payload json.RawMessage `json:"payload,omitempty"`
|
Payload json.RawMessage `json:"payload,omitempty"`
|
||||||
}
|
}
|
||||||
err := websocket.JSON.Receive(conn.conn, &raw)
|
if err := conn.sharedConn.Receive(&raw); err != nil {
|
||||||
if err != nil {
|
return service.WSMessage{}, err
|
||||||
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
|
|
||||||
}
|
}
|
||||||
message.Type = raw.Type
|
return service.WSMessage{
|
||||||
message.Payload = raw.Payload
|
Type: raw.Type,
|
||||||
slog.Debug("relay ws received message", "type", message.Type)
|
Payload: raw.Payload,
|
||||||
return message, nil
|
}, 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 {
|
func (conn *Connection) Close() error {
|
||||||
if conn == nil || conn.conn == nil {
|
return conn.sharedConn.Close()
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return conn.conn.Close()
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -3,142 +3,62 @@ package wsclient
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
|
||||||
"log/slog"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/url"
|
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"golang.org/x/net/websocket"
|
|
||||||
"openflare/service"
|
"openflare/service"
|
||||||
|
shared "openflare/utils/wsclient"
|
||||||
)
|
)
|
||||||
|
|
||||||
type Client struct {
|
type Client struct {
|
||||||
baseURL string
|
sharedClient *shared.Client
|
||||||
token string
|
|
||||||
timeout time.Duration
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type Connection struct {
|
type Connection struct {
|
||||||
conn *websocket.Conn
|
sharedConn *shared.Connection
|
||||||
url string
|
|
||||||
readTimeout time.Duration
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||||
return &Client{
|
return &Client{
|
||||||
baseURL: strings.TrimRight(baseURL, "/"),
|
sharedClient: shared.New(shared.Config{
|
||||||
token: strings.TrimSpace(token),
|
BaseURL: baseURL,
|
||||||
timeout: timeout,
|
Token: token,
|
||||||
|
Timeout: timeout,
|
||||||
|
HeaderKey: "X-Tunnel-Token",
|
||||||
|
WSPath: "/api/flared/ws",
|
||||||
|
}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) SetToken(token string) {
|
func (c *Client) SetToken(token string) {
|
||||||
c.token = strings.TrimSpace(token)
|
c.sharedClient.SetToken(token)
|
||||||
slog.Debug("flared ws client token updated")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
func (c *Client) Connect(ctx context.Context) (*Connection, error) {
|
||||||
wsURL, err := buildWebsocketURL(c.baseURL)
|
conn, err := c.sharedClient.Connect(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(c.token) == "" {
|
return &Connection{sharedConn: conn}, nil
|
||||||
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
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) SendPong() error {
|
func (conn *Connection) SendPong() error {
|
||||||
if conn == nil || conn.conn == nil {
|
return conn.sharedConn.SendMessage("pong", 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",
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func (conn *Connection) Receive() (service.WSMessage, error) {
|
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 {
|
var raw struct {
|
||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
Payload json.RawMessage `json:"payload,omitempty"`
|
Payload json.RawMessage `json:"payload,omitempty"`
|
||||||
}
|
}
|
||||||
err := websocket.JSON.Receive(conn.conn, &raw)
|
if err := conn.sharedConn.Receive(&raw); err != nil {
|
||||||
if err != nil {
|
return service.WSMessage{}, err
|
||||||
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
|
|
||||||
}
|
}
|
||||||
message.Type = raw.Type
|
return service.WSMessage{
|
||||||
message.Payload = raw.Payload
|
Type: raw.Type,
|
||||||
slog.Debug("flared ws received message", "type", message.Type)
|
Payload: raw.Payload,
|
||||||
return message, nil
|
}, 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 {
|
func (conn *Connection) Close() error {
|
||||||
if conn == nil || conn.conn == nil {
|
return conn.sharedConn.Close()
|
||||||
return nil
|
|
||||||
}
|
|
||||||
return conn.conn.Close()
|
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user