mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +08:00
Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4566fc1f53 | |||
| 4e58bdd85b | |||
| c009b9e283 | |||
| 4e33e0e521 | |||
| 7252fb6285 |
@@ -12,6 +12,7 @@ import (
|
||||
"openflare-agent/internal/observability"
|
||||
"openflare-agent/internal/protocol"
|
||||
"openflare-agent/internal/state"
|
||||
"openflare-agent/internal/wsclient"
|
||||
)
|
||||
|
||||
type HeartbeatService interface {
|
||||
@@ -211,53 +212,75 @@ func (r *Runner) startWebSocket(ctx context.Context, nodeID string) (<-chan erro
|
||||
return done, nil
|
||||
}
|
||||
|
||||
type agentWSHandler struct {
|
||||
runner *Runner
|
||||
conn protocol.WebSocketConnection
|
||||
nodeID string
|
||||
statusTicker *time.Ticker
|
||||
}
|
||||
|
||||
func (h *agentWSHandler) OnConnect(ctx context.Context) error {
|
||||
return h.runner.sendWebSocketStatus(ctx, h.nodeID, h.conn)
|
||||
}
|
||||
|
||||
func (h *agentWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||
var payloadBytes []byte
|
||||
if msg.Payload != nil {
|
||||
payloadBytes = []byte(msg.Payload)
|
||||
}
|
||||
protoMsg := protocol.WSMessage{
|
||||
Type: msg.Type,
|
||||
Payload: payloadBytes,
|
||||
}
|
||||
changed, err := h.runner.handleWebSocketMessage(ctx, protoMsg, h.conn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if changed {
|
||||
h.statusTicker.Reset(h.runner.Config.HeartbeatInterval.Duration())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *agentWSHandler) OnClose(err error) {
|
||||
slog.Error("agent ws receive failed", "error", err)
|
||||
}
|
||||
|
||||
func (r *Runner) runWebSocket(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
|
||||
slog.Debug("agent ws connected", "url", conn.URL(), "node_id", nodeID)
|
||||
statusTicker := time.NewTicker(r.Config.HeartbeatInterval.Duration())
|
||||
defer statusTicker.Stop()
|
||||
|
||||
messages := make(chan protocol.WSMessage, 8)
|
||||
readDone := make(chan error, 1)
|
||||
childCtx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
// Start status ticker sender in background
|
||||
go func() {
|
||||
for {
|
||||
message, err := conn.Receive()
|
||||
if err != nil {
|
||||
readDone <- err
|
||||
return
|
||||
}
|
||||
select {
|
||||
case messages <- message:
|
||||
case <-ctx.Done():
|
||||
readDone <- ctx.Err()
|
||||
case <-childCtx.Done():
|
||||
return
|
||||
case <-statusTicker.C:
|
||||
if err := r.sendWebSocketStatus(childCtx, nodeID, conn); err != nil {
|
||||
slog.Error("agent ws send status failed", "error", err)
|
||||
_ = conn.Close()
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
if err := r.sendWebSocketStatus(ctx, nodeID, conn); err != nil {
|
||||
return err
|
||||
wsConn, ok := conn.(*wsclient.Connection)
|
||||
if !ok {
|
||||
return errors.New("invalid websocket connection type")
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case err := <-readDone:
|
||||
return err
|
||||
case <-statusTicker.C:
|
||||
if err := r.sendWebSocketStatus(ctx, nodeID, conn); err != nil {
|
||||
return err
|
||||
}
|
||||
case message := <-messages:
|
||||
changed, err := r.handleWebSocketMessage(ctx, message, conn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if changed {
|
||||
statusTicker.Reset(r.Config.HeartbeatInterval.Duration())
|
||||
}
|
||||
}
|
||||
}
|
||||
return wsConn.RunReceiveLoop(childCtx, &agentWSHandler{
|
||||
runner: r,
|
||||
conn: conn,
|
||||
nodeID: nodeID,
|
||||
statusTicker: statusTicker,
|
||||
})
|
||||
}
|
||||
|
||||
func (r *Runner) sendWebSocketStatus(ctx context.Context, nodeID string, conn protocol.WebSocketConnection) error {
|
||||
|
||||
@@ -2,165 +2,81 @@ 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 WSMessage = shared.WSMessage
|
||||
type MessageHandler = shared.MessageHandler
|
||||
|
||||
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) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error {
|
||||
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||
}
|
||||
|
||||
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()
|
||||
}
|
||||
|
||||
@@ -102,19 +102,32 @@ func (m *Manager) UpdateConfig(cfg *service.RelayConfig) {
|
||||
m.activeConfig.WebServerEnabled == cfg.WebServerEnabled {
|
||||
if m.cmd == nil && !m.stopping {
|
||||
slog.Warn("frps config unchanged but process is not running, restarting")
|
||||
if err := m.restartProcess(); err != nil {
|
||||
m.stopping = false
|
||||
m.generation++
|
||||
generation := m.generation
|
||||
if err := m.renderConfig(cfg); err != nil {
|
||||
slog.Error("failed to render frps config", "error", err)
|
||||
m.status = "unhealthy"
|
||||
m.lastError = err.Error()
|
||||
slog.Error("failed to restart frps with unchanged config", "error", err)
|
||||
return
|
||||
}
|
||||
go m.supervise(generation)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
m.activeConfig = cfg
|
||||
m.stopping = false
|
||||
m.generation++
|
||||
generation := m.generation
|
||||
slog.Info("relay config updated, reloading frps")
|
||||
|
||||
if m.cmd != nil && m.cmd.Process != nil {
|
||||
slog.Debug("stopping existing frps process")
|
||||
_ = m.cmd.Process.Kill()
|
||||
m.cmd = nil
|
||||
}
|
||||
|
||||
if err := m.renderConfig(cfg); err != nil {
|
||||
slog.Error("failed to render frps config", "error", err)
|
||||
m.status = "unhealthy"
|
||||
@@ -122,14 +135,7 @@ func (m *Manager) UpdateConfig(cfg *service.RelayConfig) {
|
||||
return
|
||||
}
|
||||
|
||||
if err := m.restartProcess(); err != nil {
|
||||
slog.Error("failed to restart frps", "error", err)
|
||||
m.status = "unhealthy"
|
||||
m.lastError = err.Error()
|
||||
} else {
|
||||
m.status = "healthy"
|
||||
m.lastError = ""
|
||||
}
|
||||
go m.supervise(generation)
|
||||
}
|
||||
|
||||
func (m *Manager) renderConfig(cfg *service.RelayConfig) error {
|
||||
@@ -166,63 +172,95 @@ func (m *Manager) renderConfig(cfg *service.RelayConfig) error {
|
||||
return os.WriteFile(m.configPath, buf.Bytes(), 0644)
|
||||
}
|
||||
|
||||
func (m *Manager) restartProcess() error {
|
||||
m.generation++
|
||||
generation := m.generation
|
||||
if m.cmd != nil && m.cmd.Process != nil {
|
||||
slog.Debug("stopping existing frps process")
|
||||
_ = m.cmd.Process.Kill()
|
||||
m.cmd = nil
|
||||
}
|
||||
return m.startProcessLocked(generation)
|
||||
}
|
||||
func (m *Manager) supervise(generation uint64) {
|
||||
backoff := 1 * time.Second
|
||||
const maxBackoff = 60 * time.Second
|
||||
|
||||
func (m *Manager) startProcessLocked(generation uint64) error {
|
||||
cmd := exec.Command(m.frpsPath, "-c", m.configPath)
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.cmd = cmd
|
||||
m.status = "healthy"
|
||||
m.lastError = ""
|
||||
|
||||
go func(c *exec.Cmd) {
|
||||
err := c.Wait()
|
||||
slog.Warn("frps process exited", "error", err)
|
||||
for {
|
||||
m.mu.Lock()
|
||||
if m.cmd == c {
|
||||
if m.stopping || m.generation != generation {
|
||||
m.mu.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
cmd := exec.Command(m.frpsPath, "-c", m.configPath)
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
err := cmd.Start()
|
||||
if err != nil {
|
||||
m.status = "unhealthy"
|
||||
m.lastError = fmt.Sprintf("failed to start: %v", err)
|
||||
slog.Error("failed to start frps", "error", err, "generation", generation)
|
||||
m.mu.Unlock()
|
||||
|
||||
if !m.sleepOrInterrupt(generation, backoff) {
|
||||
return
|
||||
}
|
||||
backoff = backoff * 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
m.cmd = cmd
|
||||
m.status = "healthy"
|
||||
m.lastError = ""
|
||||
m.mu.Unlock()
|
||||
|
||||
startedAt := time.Now()
|
||||
waitErr := cmd.Wait()
|
||||
|
||||
m.mu.Lock()
|
||||
if m.cmd == cmd {
|
||||
m.cmd = nil
|
||||
m.status = "unhealthy"
|
||||
if err != nil {
|
||||
m.lastError = err.Error()
|
||||
if waitErr != nil {
|
||||
m.lastError = fmt.Sprintf("exited with error: %v", waitErr)
|
||||
} else {
|
||||
m.lastError = "frps process exited"
|
||||
m.lastError = "exited unexpectedly"
|
||||
}
|
||||
slog.Warn("frps process exited unexpectedly", "error", waitErr, "generation", generation)
|
||||
}
|
||||
shouldContinue := !m.stopping && m.generation == generation
|
||||
m.mu.Unlock()
|
||||
|
||||
if !shouldContinue {
|
||||
return
|
||||
}
|
||||
|
||||
if time.Since(startedAt) >= 10*time.Second {
|
||||
backoff = 1 * time.Second
|
||||
}
|
||||
|
||||
if !m.sleepOrInterrupt(generation, backoff) {
|
||||
return
|
||||
}
|
||||
backoff = backoff * 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) sleepOrInterrupt(generation uint64, d time.Duration) bool {
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
deadline := time.Now().Add(d)
|
||||
for time.Now().Before(deadline) {
|
||||
select {
|
||||
case <-ticker.C:
|
||||
m.mu.RLock()
|
||||
interrupted := m.stopping || m.generation != generation
|
||||
m.mu.RUnlock()
|
||||
if interrupted {
|
||||
return false
|
||||
}
|
||||
}
|
||||
shouldRestart := !m.stopping && m.generation == generation
|
||||
m.mu.Unlock()
|
||||
if !shouldRestart {
|
||||
return
|
||||
}
|
||||
time.Sleep(2 * time.Second)
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.stopping || m.generation != generation {
|
||||
return
|
||||
}
|
||||
slog.Warn("restarting frps after unexpected exit")
|
||||
if err := m.startProcessLocked(generation); err != nil {
|
||||
m.status = "unhealthy"
|
||||
m.lastError = err.Error()
|
||||
slog.Error("failed to auto restart frps", "error", err)
|
||||
}
|
||||
}(cmd)
|
||||
|
||||
return nil
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (m *Manager) Stop() {
|
||||
|
||||
@@ -0,0 +1,295 @@
|
||||
package frps
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"openflare/service"
|
||||
)
|
||||
|
||||
// Helper to write control file for the dummy script
|
||||
func writeControl(t *testing.T, dir string, exitCode int, delaySeconds int) {
|
||||
controlPath := filepath.Join(dir, "control.txt")
|
||||
content := fmt.Sprintf("%d %d\n", exitCode, delaySeconds)
|
||||
err := os.WriteFile(controlPath, []byte(content), 0644)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write control file: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Setup a dummy executable script that reads control.txt to decide exit code and sleep duration
|
||||
func setupDummyScript(t *testing.T) (string, string) {
|
||||
dir := t.TempDir()
|
||||
scriptPath := filepath.Join(dir, "dummy_frps")
|
||||
|
||||
// On macOS/Linux, we write a shell script
|
||||
scriptContent := fmt.Sprintf(`#!/bin/sh
|
||||
control_file="%s/control.txt"
|
||||
EXIT_CODE=0
|
||||
DELAY=0
|
||||
if [ -f "$control_file" ]; then
|
||||
read -r EXIT_CODE DELAY < "$control_file"
|
||||
fi
|
||||
if [ -n "$DELAY" ] && [ "$DELAY" -gt 0 ] 2>/dev/null; then
|
||||
sleep "$DELAY"
|
||||
fi
|
||||
exit "${EXIT_CODE:-0}"
|
||||
`, dir)
|
||||
|
||||
err := os.WriteFile(scriptPath, []byte(scriptContent), 0755)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write dummy script: %v", err)
|
||||
}
|
||||
|
||||
return scriptPath, dir
|
||||
}
|
||||
|
||||
// Helper to poll for status to eliminate timing flakiness in tests
|
||||
func assertStatusEventually(t *testing.T, m *Manager, expectedStatus string, timeout time.Duration) {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
rt := m.GetRuntimeStatus()
|
||||
if rt.Status == expectedStatus {
|
||||
return
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
rt := m.GetRuntimeStatus()
|
||||
t.Fatalf("expected status eventually %s, got %s (err: %s)", expectedStatus, rt.Status, rt.LastError)
|
||||
}
|
||||
|
||||
func TestStartProcessSuccess(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 5) // exit code 0, sleep 5s
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
|
||||
rt := m.GetRuntimeStatus()
|
||||
if !rt.ProcessAlive {
|
||||
t.Error("expected process to be alive")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartProcessFailureAndBackoff(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
invalidScriptPath := filepath.Join(dir, "non_existent_frps")
|
||||
|
||||
m := NewManager(invalidScriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
rt := m.GetRuntimeStatus()
|
||||
if !strings.Contains(rt.LastError, "failed to start") {
|
||||
t.Errorf("expected error message containing 'failed to start', got %s", rt.LastError)
|
||||
}
|
||||
|
||||
// Correct the path to dummy script
|
||||
scriptPath, _ := setupDummyScript(t)
|
||||
writeControl(t, filepath.Dir(scriptPath), 0, 5)
|
||||
|
||||
m.mu.Lock()
|
||||
m.frpsPath = scriptPath
|
||||
m.mu.Unlock()
|
||||
|
||||
// Wait for backoff retry (1s backoff)
|
||||
assertStatusEventually(t, m, "healthy", 3*time.Second)
|
||||
|
||||
rt = m.GetRuntimeStatus()
|
||||
if !rt.ProcessAlive {
|
||||
t.Error("expected process to be alive now")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnexpectedExitAndAutorestart(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Start with immediate exit code 1
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
rt := m.GetRuntimeStatus()
|
||||
if !strings.Contains(rt.LastError, "exited with error") {
|
||||
t.Errorf("expected exit error, got %s", rt.LastError)
|
||||
}
|
||||
|
||||
// Change control to be healthy (runs for 5s, exit 0)
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
// Wait for the retry to fire (backoff was 1s)
|
||||
assertStatusEventually(t, m, "healthy", 3*time.Second)
|
||||
}
|
||||
|
||||
func TestBackoffReset(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Rapid exit to increase backoff
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
// Crashed once, backoff is 2s
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
// Now make it run successfully for 11 seconds (exit code 0, sleep 11s)
|
||||
writeControl(t, dir, 0, 11)
|
||||
|
||||
// Wait for next retry to start running
|
||||
assertStatusEventually(t, m, "healthy", 4*time.Second)
|
||||
|
||||
// Wait for process to run for 10.5 seconds to trigger backoff reset
|
||||
time.Sleep(10500 * time.Millisecond)
|
||||
|
||||
// Now make it crash again (exit code 1, sleep 0s)
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
// Wait for it to finish and crash
|
||||
assertStatusEventually(t, m, "unhealthy", 3*time.Second)
|
||||
|
||||
// It crashed. Since it ran for > 10s, backoff should have been reset to 1s.
|
||||
// We make it healthy again (exit code 0, sleep 5)
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
// Wait 1.5 seconds. If backoff was reset to 1s, it should be healthy now.
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
}
|
||||
|
||||
func TestImmediateRestartOnSameConfigDeadProcess(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Crashes immediately
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
// Let it crash
|
||||
assertStatusEventually(t, m, "unhealthy", 2*time.Second)
|
||||
|
||||
// Make it start successfully
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
// Send same config block to trigger immediate restart bypass of backoff sleep
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
// Check if it started immediately
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
}
|
||||
|
||||
func TestSupervisorGenerationInterrupt(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 10)
|
||||
|
||||
m := NewManager(scriptPath, dir, "agent-token")
|
||||
defer m.Stop()
|
||||
|
||||
cfg := &service.RelayConfig{
|
||||
BindPort: 7000,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
|
||||
m.UpdateConfig(cfg)
|
||||
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
|
||||
m.mu.Lock()
|
||||
gen1 := m.generation
|
||||
cmd1 := m.cmd
|
||||
m.mu.Unlock()
|
||||
|
||||
if cmd1 == nil {
|
||||
t.Fatal("expected active process")
|
||||
}
|
||||
|
||||
// Update configuration with new bind port to trigger new generation
|
||||
cfg2 := &service.RelayConfig{
|
||||
BindPort: 7001,
|
||||
VhostHTTPPort: 8080,
|
||||
AuthToken: "test-auth",
|
||||
WebServerEnabled: false,
|
||||
}
|
||||
m.UpdateConfig(cfg2)
|
||||
|
||||
assertStatusEventually(t, m, "healthy", 2*time.Second)
|
||||
|
||||
m.mu.Lock()
|
||||
gen2 := m.generation
|
||||
cmd2 := m.cmd
|
||||
m.mu.Unlock()
|
||||
|
||||
if gen2 <= gen1 {
|
||||
t.Errorf("expected generation incremented, got gen1=%d gen2=%d", gen1, gen2)
|
||||
}
|
||||
if cmd2 == cmd1 {
|
||||
t.Error("expected old process killed and new command started")
|
||||
}
|
||||
|
||||
// Verify old process is actually killed
|
||||
var cmd1Finished int32
|
||||
go func() {
|
||||
_ = cmd1.Wait()
|
||||
atomic.StoreInt32(&cmd1Finished, 1)
|
||||
}()
|
||||
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
if atomic.LoadInt32(&cmd1Finished) != 1 {
|
||||
t.Error("expected first process to be killed")
|
||||
}
|
||||
}
|
||||
@@ -51,66 +51,35 @@ func (r *Runner) Run(ctx context.Context) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
||||
// Send pings at 2× heartbeat interval to keep the server-side read deadline
|
||||
// from expiring (server closes the WS if no data arrives within ~30 s).
|
||||
pingInterval := r.Config.HeartbeatInterval.Duration() * 2
|
||||
pingTicker := time.NewTicker(pingInterval)
|
||||
defer pingTicker.Stop()
|
||||
type relayWSHandler struct {
|
||||
runner *Runner
|
||||
}
|
||||
|
||||
messages := make(chan service.WSMessage, 8)
|
||||
readDone := make(chan error, 1)
|
||||
go func() {
|
||||
for {
|
||||
msg, err := conn.Receive()
|
||||
if err != nil {
|
||||
readDone <- err
|
||||
return
|
||||
}
|
||||
select {
|
||||
case messages <- msg:
|
||||
case <-ctx.Done():
|
||||
readDone <- ctx.Err()
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
func (h *relayWSHandler) OnConnect(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case err := <-readDone:
|
||||
slog.Error("relay ws receive failed", "error", err)
|
||||
return
|
||||
case <-pingTicker.C:
|
||||
if err := conn.SendPing(); err != nil {
|
||||
slog.Error("relay ws send ping failed", "error", err)
|
||||
return
|
||||
}
|
||||
case msg := <-messages:
|
||||
switch msg.Type {
|
||||
case "ping":
|
||||
_ = conn.SendPong()
|
||||
case "pong":
|
||||
slog.Debug("relay ws pong received")
|
||||
case "relay_config":
|
||||
payloadBytes, ok := msg.Payload.(json.RawMessage)
|
||||
if !ok {
|
||||
slog.Error("invalid relay_config payload type")
|
||||
continue
|
||||
}
|
||||
var cfg service.RelayConfig
|
||||
if err := json.Unmarshal(payloadBytes, &cfg); err != nil {
|
||||
slog.Error("failed to unmarshal relay_config", "error", err)
|
||||
continue
|
||||
}
|
||||
r.FrpsManager.UpdateConfig(&cfg)
|
||||
default:
|
||||
slog.Debug("ignored unknown ws message type", "type", msg.Type)
|
||||
}
|
||||
func (h *relayWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||
switch msg.Type {
|
||||
case "relay_config":
|
||||
var cfg service.RelayConfig
|
||||
if err := json.Unmarshal(msg.Payload, &cfg); err != nil {
|
||||
slog.Error("failed to unmarshal relay_config", "error", err)
|
||||
return nil
|
||||
}
|
||||
h.runner.FrpsManager.UpdateConfig(&cfg)
|
||||
default:
|
||||
slog.Debug("ignored unknown ws message type", "type", msg.Type)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *relayWSHandler) OnClose(err error) {
|
||||
slog.Error("relay ws receive failed", "error", err)
|
||||
}
|
||||
|
||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
||||
_ = conn.RunReceiveLoop(ctx, &relayWSHandler{runner: r})
|
||||
}
|
||||
|
||||
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) {
|
||||
|
||||
@@ -3,151 +3,73 @@ 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 WSMessage = shared.WSMessage
|
||||
type MessageHandler = shared.MessageHandler
|
||||
|
||||
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")
|
||||
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")
|
||||
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
|
||||
return service.WSMessage{
|
||||
Type: raw.Type,
|
||||
Payload: raw.Payload,
|
||||
}, 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) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error {
|
||||
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||
}
|
||||
|
||||
func (conn *Connection) Close() error {
|
||||
if conn == nil || conn.conn == nil {
|
||||
return nil
|
||||
}
|
||||
return conn.conn.Close()
|
||||
return conn.sharedConn.Close()
|
||||
}
|
||||
|
||||
@@ -49,6 +49,8 @@ func main() {
|
||||
gin.SetMode(gin.ReleaseMode)
|
||||
}
|
||||
// Initialize SQL Database
|
||||
defer service.ShutdownWSHubs()
|
||||
|
||||
err := model.InitDB()
|
||||
if err != nil {
|
||||
slog.Error("initialize database failed", "error", err)
|
||||
|
||||
@@ -3,6 +3,7 @@ package service
|
||||
import (
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type WSMessage struct {
|
||||
@@ -65,13 +66,58 @@ type WSHub struct {
|
||||
name string
|
||||
mu sync.RWMutex
|
||||
clients map[string]*WSClient
|
||||
done chan struct{}
|
||||
}
|
||||
|
||||
func NewWSHub(name string) *WSHub {
|
||||
return &WSHub{
|
||||
h := &WSHub{
|
||||
name: name,
|
||||
clients: make(map[string]*WSClient),
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
go h.startPingLoop()
|
||||
return h
|
||||
}
|
||||
|
||||
func (h *WSHub) Close() {
|
||||
close(h.done)
|
||||
}
|
||||
|
||||
func (h *WSHub) startPingLoop() {
|
||||
ticker := time.NewTicker(10 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-h.done:
|
||||
return
|
||||
case <-ticker.C:
|
||||
h.mu.RLock()
|
||||
if len(h.clients) == 0 {
|
||||
h.mu.RUnlock()
|
||||
continue
|
||||
}
|
||||
clients := make([]*WSClient, 0, len(h.clients))
|
||||
for _, client := range h.clients {
|
||||
clients = append(clients, client)
|
||||
}
|
||||
h.mu.RUnlock()
|
||||
|
||||
for _, client := range clients {
|
||||
if !client.Send(WSMessage{
|
||||
Type: "ping",
|
||||
}) {
|
||||
slog.Warn("ws client send ping failed, queue full, disconnecting", "hub", h.name, "id", client.id)
|
||||
h.Disconnect(client.id)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func ShutdownWSHubs() {
|
||||
DefaultAgentWSHub.Close()
|
||||
DefaultFlaredWSHub.Close()
|
||||
DefaultRelayWSHub.Close()
|
||||
}
|
||||
|
||||
func (h *WSHub) Register(id string) *WSClient {
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
package wsclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"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 WSMessage struct {
|
||||
Type string `json:"type"`
|
||||
Payload json.RawMessage `json:"payload,omitempty"`
|
||||
}
|
||||
|
||||
type MessageHandler interface {
|
||||
OnConnect(ctx context.Context) error
|
||||
HandleMessage(ctx context.Context, msg WSMessage) error
|
||||
OnClose(err error)
|
||||
}
|
||||
|
||||
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) RunReceiveLoop(ctx context.Context, handler MessageHandler) error {
|
||||
doneChan := make(chan struct{})
|
||||
defer close(doneChan)
|
||||
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
_ = conn.Close()
|
||||
case <-doneChan:
|
||||
}
|
||||
}()
|
||||
|
||||
if err := handler.OnConnect(ctx); err != nil {
|
||||
handler.OnClose(err)
|
||||
return err
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
var raw WSMessage
|
||||
if err := conn.Receive(&raw); err != nil {
|
||||
handler.OnClose(err)
|
||||
return err
|
||||
}
|
||||
|
||||
switch raw.Type {
|
||||
case "ping":
|
||||
slog.Debug("ws received ping from server, replying with pong")
|
||||
if err := conn.SendMessage("pong", nil); err != nil {
|
||||
slog.Error("ws send pong response failed", "error", err)
|
||||
}
|
||||
case "pong":
|
||||
slog.Debug("ws received pong response from server")
|
||||
default:
|
||||
if err := handler.HandleMessage(ctx, raw); err != nil {
|
||||
slog.Error("ws handler failed to process message", "type", raw.Type, "error", err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (conn *Connection) Close() error {
|
||||
if conn == nil || conn.Conn == nil {
|
||||
return nil
|
||||
}
|
||||
return conn.Conn.Close()
|
||||
}
|
||||
@@ -49,31 +49,31 @@ func (r *Runner) Run(ctx context.Context) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
default:
|
||||
}
|
||||
type flaredWSHandler struct {
|
||||
runner *Runner
|
||||
}
|
||||
|
||||
msg, err := conn.Receive()
|
||||
if err != nil {
|
||||
slog.Error("flared ws receive failed", "error", err)
|
||||
return
|
||||
}
|
||||
func (h *flaredWSHandler) OnConnect(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
switch msg.Type {
|
||||
case "ping":
|
||||
_ = conn.SendPong()
|
||||
case "active_config":
|
||||
// Server notifies there is a new config available
|
||||
slog.Info("received config update notification from server")
|
||||
r.SyncService.Trigger()
|
||||
default:
|
||||
slog.Debug("ignored unknown ws message type", "type", msg.Type)
|
||||
}
|
||||
func (h *flaredWSHandler) HandleMessage(ctx context.Context, msg wsclient.WSMessage) error {
|
||||
switch msg.Type {
|
||||
case "active_config":
|
||||
slog.Info("received config update notification from server")
|
||||
h.runner.SyncService.Trigger()
|
||||
default:
|
||||
slog.Debug("ignored unknown ws message type", "type", msg.Type)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *flaredWSHandler) OnClose(err error) {
|
||||
slog.Error("flared ws receive failed", "error", err)
|
||||
}
|
||||
|
||||
func (r *Runner) handleConnection(ctx context.Context, conn *wsclient.Connection) {
|
||||
_ = conn.RunReceiveLoop(ctx, &flaredWSHandler{runner: r})
|
||||
}
|
||||
|
||||
func (r *Runner) sleepContext(ctx context.Context, d time.Duration) {
|
||||
|
||||
@@ -156,28 +156,57 @@ func (m *Manager) restartProcess(ctx context.Context, relayID string, configPath
|
||||
m.processes[relayID] = proc
|
||||
|
||||
go func() {
|
||||
backoff := 1 * time.Second
|
||||
const maxBackoff = 60 * time.Second
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-procCtx.Done():
|
||||
m.mu.Lock()
|
||||
if procCtx.Err() != nil {
|
||||
m.mu.Unlock()
|
||||
return
|
||||
default:
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
cmd := exec.CommandContext(procCtx, m.cfg.FrpcPath, "-c", configPath)
|
||||
|
||||
m.mu.Lock()
|
||||
proc.Cmd = cmd
|
||||
proc.Status = "running"
|
||||
m.mu.Unlock()
|
||||
|
||||
startedAt := time.Now()
|
||||
err := cmd.Run()
|
||||
|
||||
m.mu.Lock()
|
||||
if procCtx.Err() != nil {
|
||||
proc.Status = "stopped"
|
||||
m.mu.Unlock()
|
||||
return
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
if procCtx.Err() != nil {
|
||||
return
|
||||
}
|
||||
proc.LastError = err.Error()
|
||||
proc.Status = "error"
|
||||
slog.Error("frpc process exited unexpectedly", "relay_id", relayID, "error", err)
|
||||
time.Sleep(5 * time.Second) // backoff
|
||||
} else {
|
||||
proc.Status = "stopped"
|
||||
proc.LastError = "exited unexpectedly with code 0"
|
||||
slog.Warn("frpc process exited unexpectedly with code 0", "relay_id", relayID)
|
||||
}
|
||||
m.mu.Unlock()
|
||||
|
||||
if time.Since(startedAt) >= 10*time.Second {
|
||||
backoff = 1 * time.Second
|
||||
}
|
||||
|
||||
select {
|
||||
case <-procCtx.Done():
|
||||
return
|
||||
case <-time.After(backoff):
|
||||
backoff = backoff * 2
|
||||
if backoff > maxBackoff {
|
||||
backoff = maxBackoff
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -0,0 +1,267 @@
|
||||
package frpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"openflare-flared/internal/config"
|
||||
"openflare/service"
|
||||
)
|
||||
|
||||
// Helper to write control file for the dummy script
|
||||
func writeControl(t *testing.T, dir string, exitCode int, delaySeconds int) {
|
||||
controlPath := filepath.Join(dir, "control.txt")
|
||||
content := fmt.Sprintf("%d %d\n", exitCode, delaySeconds)
|
||||
err := os.WriteFile(controlPath, []byte(content), 0644)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write control file: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Setup a dummy executable script that reads control.txt to decide exit code and sleep duration
|
||||
func setupDummyScript(t *testing.T) (string, string) {
|
||||
dir := t.TempDir()
|
||||
scriptPath := filepath.Join(dir, "dummy_frpc")
|
||||
|
||||
// On macOS/Linux, we write a shell script
|
||||
scriptContent := fmt.Sprintf(`#!/bin/sh
|
||||
control_file="%s/control.txt"
|
||||
EXIT_CODE=0
|
||||
DELAY=0
|
||||
if [ -f "$control_file" ]; then
|
||||
read -r EXIT_CODE DELAY < "$control_file"
|
||||
fi
|
||||
if [ -n "$DELAY" ] && [ "$DELAY" -gt 0 ] 2>/dev/null; then
|
||||
sleep "$DELAY"
|
||||
fi
|
||||
exit "${EXIT_CODE:-0}"
|
||||
`, dir)
|
||||
|
||||
err := os.WriteFile(scriptPath, []byte(scriptContent), 0755)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to write dummy script: %v", err)
|
||||
}
|
||||
|
||||
return scriptPath, dir
|
||||
}
|
||||
|
||||
// Helper to poll for status to eliminate timing flakiness in tests
|
||||
func assertStatusEventually(t *testing.T, m *Manager, relayID string, expectedStatus string, timeout time.Duration) {
|
||||
deadline := time.Now().Add(timeout)
|
||||
for time.Now().Before(deadline) {
|
||||
m.mu.RLock()
|
||||
proc, ok := m.processes[relayID]
|
||||
m.mu.RUnlock()
|
||||
if ok && proc.Status == expectedStatus {
|
||||
return
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
m.mu.RLock()
|
||||
proc, ok := m.processes[relayID]
|
||||
var got string
|
||||
var errStr string
|
||||
if ok {
|
||||
got = proc.Status
|
||||
errStr = proc.LastError
|
||||
} else {
|
||||
got = "not_found"
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
t.Fatalf("expected status eventually %s, got %s (err: %s)", expectedStatus, got, errStr)
|
||||
}
|
||||
|
||||
func TestStartProcessSuccess(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
writeControl(t, dir, 0, 5) // exit code 0, sleep 5s
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
err := m.UpdateConfig(context.Background(), newConfig)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to UpdateConfig: %v", err)
|
||||
}
|
||||
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
m.mu.RUnlock()
|
||||
|
||||
proc.Cancel()
|
||||
assertStatusEventually(t, m, "relay-1", "stopped", 4*time.Second) // wait for clean stop
|
||||
}
|
||||
|
||||
func TestStartProcessFailureAndBackoff(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
invalidScriptPath := filepath.Join(dir, "non_existent_frpc")
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: invalidScriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
_ = m.UpdateConfig(context.Background(), newConfig)
|
||||
|
||||
assertStatusEventually(t, m, "relay-1", "error", 4*time.Second)
|
||||
|
||||
// Correct the path to dummy script
|
||||
scriptPath, _ := setupDummyScript(t)
|
||||
writeControl(t, filepath.Dir(scriptPath), 0, 5)
|
||||
|
||||
m.mu.Lock()
|
||||
m.cfg.FrpcPath = scriptPath
|
||||
m.mu.Unlock()
|
||||
|
||||
// Wait for backoff retry (1s backoff)
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
m.mu.RUnlock()
|
||||
proc.Cancel()
|
||||
}
|
||||
|
||||
func TestUnexpectedExit0CPUProtection(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Start with immediate exit code 0
|
||||
writeControl(t, dir, 0, 0)
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
_ = m.UpdateConfig(context.Background(), newConfig)
|
||||
|
||||
assertStatusEventually(t, m, "relay-1", "stopped", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
if !strings.Contains(proc.LastError, "exited unexpectedly with code 0") {
|
||||
t.Errorf("expected LastError to record exit status 0 warning, got %s", proc.LastError)
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
proc.Cancel()
|
||||
}
|
||||
|
||||
func TestBackoffReset(t *testing.T) {
|
||||
scriptPath, dir := setupDummyScript(t)
|
||||
// Rapid exit code 1 to increase backoff
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
cfg := &config.Config{
|
||||
ServerURL: "http://localhost:8080",
|
||||
TunnelToken: "test-token",
|
||||
FrpcPath: scriptPath,
|
||||
DataDir: dir,
|
||||
StatePath: filepath.Join(dir, "flared-state.json"),
|
||||
}
|
||||
|
||||
m := NewManager(cfg)
|
||||
newConfig := &service.FlaredTunnelConfigResponse{
|
||||
Version: "1",
|
||||
Checksum: "sum1",
|
||||
Relays: []service.FlaredRelayInfo{
|
||||
{
|
||||
RelayNodeID: "relay-1",
|
||||
Address: "127.0.0.1:7000",
|
||||
AuthToken: "auth-1",
|
||||
},
|
||||
},
|
||||
Proxies: nil,
|
||||
}
|
||||
|
||||
_ = m.UpdateConfig(context.Background(), newConfig)
|
||||
|
||||
// Wait to crash
|
||||
assertStatusEventually(t, m, "relay-1", "error", 4*time.Second)
|
||||
|
||||
// Now make it run successfully for 11 seconds (exit code 0, sleep 11s)
|
||||
writeControl(t, dir, 0, 11)
|
||||
|
||||
// Wait for next retry to start running
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
// Wait for process to run for 10.5 seconds to trigger backoff reset
|
||||
time.Sleep(10500 * time.Millisecond)
|
||||
|
||||
// Now make it crash again (exit code 1, sleep 0s)
|
||||
writeControl(t, dir, 1, 0)
|
||||
|
||||
// Wait for it to finish and crash
|
||||
assertStatusEventually(t, m, "relay-1", "error", 4*time.Second)
|
||||
|
||||
// It crashed. Since it ran for > 10s, backoff should have been reset to 1s.
|
||||
// We make it healthy again (exit code 0, sleep 5)
|
||||
writeControl(t, dir, 0, 5)
|
||||
|
||||
// Wait 1.5 seconds. If backoff was reset to 1s, it should be running now.
|
||||
assertStatusEventually(t, m, "relay-1", "running", 4*time.Second)
|
||||
|
||||
m.mu.RLock()
|
||||
proc := m.processes["relay-1"]
|
||||
m.mu.RUnlock()
|
||||
proc.Cancel()
|
||||
}
|
||||
@@ -3,141 +3,73 @@ 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 WSMessage = shared.WSMessage
|
||||
type MessageHandler = shared.MessageHandler
|
||||
|
||||
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
|
||||
return &Connection{sharedConn: conn}, 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) SendPing() error {
|
||||
return conn.sharedConn.SendMessage("ping", 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")
|
||||
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
|
||||
return service.WSMessage{
|
||||
Type: raw.Type,
|
||||
Payload: raw.Payload,
|
||||
}, 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) RunReceiveLoop(ctx context.Context, handler shared.MessageHandler) error {
|
||||
return conn.sharedConn.RunReceiveLoop(ctx, handler)
|
||||
}
|
||||
|
||||
func (conn *Connection) Close() error {
|
||||
if conn == nil || conn.conn == nil {
|
||||
return nil
|
||||
}
|
||||
return conn.conn.Close()
|
||||
return conn.sharedConn.Close()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user