mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 00:26:37 +08:00
feat: Implement node management features including creation, updating, and deletion
- Added CreateNode, UpdateNode, and DeleteNode functions in the controller for managing nodes. - Introduced NodeInput struct for input validation during node creation and updates. - Enhanced the Node model to include DiscoveryToken and Pending status. - Updated the AgentRegister and AgentHeartbeat functions to utilize the new node management logic. - Refactored the API router to include new routes for node management. - Improved the frontend Node component to support node creation, editing, and deletion with appropriate UI feedback. - Added tests to ensure the new functionality works as expected.
This commit is contained in:
@@ -2,7 +2,9 @@ package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"atsflare-agent/internal/config"
|
||||
@@ -11,8 +13,9 @@ import (
|
||||
)
|
||||
|
||||
type HeartbeatService interface {
|
||||
Register(ctx context.Context, payload protocol.NodePayload) error
|
||||
Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error)
|
||||
Heartbeat(ctx context.Context, payload protocol.NodePayload) error
|
||||
SetToken(token string)
|
||||
}
|
||||
|
||||
type SyncService interface {
|
||||
@@ -33,21 +36,20 @@ func (r *Runner) Run(ctx context.Context) error {
|
||||
return err
|
||||
}
|
||||
log.Printf("agent runner started: node_id=%s node=%s ip=%s", nodeID, r.Config.NodeName, r.Config.NodeIP)
|
||||
if err = r.HeartbeatService.Register(ctx, r.nodePayload(nodeID)); err != nil {
|
||||
log.Printf("agent register failed: %v", err)
|
||||
} else {
|
||||
log.Printf("agent register succeeded: node_id=%s", nodeID)
|
||||
}
|
||||
if err = r.SyncService.SyncOnStartup(ctx); err != nil {
|
||||
r.recordSyncError(err)
|
||||
log.Printf("agent startup sync failed: %v", err)
|
||||
} else {
|
||||
log.Printf("agent startup sync completed")
|
||||
}
|
||||
if err = r.HeartbeatService.Heartbeat(ctx, r.nodePayload(nodeID)); err != nil {
|
||||
log.Printf("agent startup heartbeat failed: %v", err)
|
||||
} else {
|
||||
log.Printf("agent startup heartbeat succeeded: node_id=%s", nodeID)
|
||||
if r.hasAgentToken() {
|
||||
if err = r.SyncService.SyncOnStartup(ctx); err != nil {
|
||||
r.recordSyncError(err)
|
||||
log.Printf("agent startup sync failed: %v", err)
|
||||
} else {
|
||||
log.Printf("agent startup sync completed")
|
||||
}
|
||||
if err = r.HeartbeatService.Heartbeat(ctx, r.nodePayload(nodeID)); err != nil {
|
||||
log.Printf("agent startup heartbeat failed: %v", err)
|
||||
} else {
|
||||
log.Printf("agent startup heartbeat succeeded: node_id=%s", nodeID)
|
||||
}
|
||||
} else if err = r.tryRegister(ctx, &nodeID); err != nil {
|
||||
log.Printf("agent initial discovery register failed: %v", err)
|
||||
}
|
||||
|
||||
heartbeatTicker := time.NewTicker(r.Config.HeartbeatInterval)
|
||||
@@ -61,10 +63,19 @@ func (r *Runner) Run(ctx context.Context) error {
|
||||
log.Printf("agent runner shutting down: %v", ctx.Err())
|
||||
return ctx.Err()
|
||||
case <-heartbeatTicker.C:
|
||||
if !r.hasAgentToken() {
|
||||
if err = r.tryRegister(ctx, &nodeID); err != nil {
|
||||
log.Printf("agent discovery register failed: %v", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err = r.HeartbeatService.Heartbeat(ctx, r.nodePayload(nodeID)); err != nil {
|
||||
log.Printf("agent heartbeat failed: %v", err)
|
||||
}
|
||||
case <-syncTicker.C:
|
||||
if !r.hasAgentToken() {
|
||||
continue
|
||||
}
|
||||
log.Printf("agent sync tick: node_id=%s", nodeID)
|
||||
if err = r.SyncService.SyncOnce(ctx); err != nil {
|
||||
r.recordSyncError(err)
|
||||
@@ -76,6 +87,47 @@ func (r *Runner) Run(ctx context.Context) error {
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) hasAgentToken() bool {
|
||||
return strings.TrimSpace(r.Config.AgentToken) != ""
|
||||
}
|
||||
|
||||
func (r *Runner) tryRegister(ctx context.Context, nodeID *string) error {
|
||||
if strings.TrimSpace(r.Config.DiscoveryToken) == "" {
|
||||
return errors.New("agent_token 为空且未配置 discovery_token")
|
||||
}
|
||||
log.Printf("agent discovery registration started")
|
||||
response, err := r.HeartbeatService.Register(ctx, r.nodePayload(*nodeID))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if response == nil || strings.TrimSpace(response.AgentToken) == "" || strings.TrimSpace(response.NodeID) == "" {
|
||||
return errors.New("discovery register response 缺少 node_id 或 agent_token")
|
||||
}
|
||||
snapshot, err := r.StateStore.Load()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
snapshot.NodeID = response.NodeID
|
||||
if err = r.StateStore.Save(snapshot); err != nil {
|
||||
return err
|
||||
}
|
||||
r.Config.AgentToken = response.AgentToken
|
||||
r.Config.DiscoveryToken = ""
|
||||
if err = r.Config.Save(); err != nil {
|
||||
return err
|
||||
}
|
||||
r.HeartbeatService.SetToken(response.AgentToken)
|
||||
*nodeID = response.NodeID
|
||||
log.Printf("agent discovery registration succeeded: node_id=%s", response.NodeID)
|
||||
if err = r.SyncService.SyncOnStartup(ctx); err != nil {
|
||||
r.recordSyncError(err)
|
||||
log.Printf("agent post-register startup sync failed: %v", err)
|
||||
} else {
|
||||
log.Printf("agent post-register startup sync completed")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Runner) recordSyncError(err error) {
|
||||
if err == nil || r.StateStore == nil {
|
||||
return
|
||||
|
||||
@@ -3,6 +3,7 @@ package agent
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
@@ -18,15 +19,17 @@ type fakeHeartbeatService struct {
|
||||
registerCalls int
|
||||
heartbeatCalls int
|
||||
registerErr error
|
||||
registerResp *protocol.RegisterNodeResponse
|
||||
heartbeatErrs []error
|
||||
onHeartbeat func(int)
|
||||
lastToken string
|
||||
}
|
||||
|
||||
func (f *fakeHeartbeatService) Register(ctx context.Context, payload protocol.NodePayload) error {
|
||||
func (f *fakeHeartbeatService) Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.registerCalls++
|
||||
return f.registerErr
|
||||
return f.registerResp, f.registerErr
|
||||
}
|
||||
|
||||
func (f *fakeHeartbeatService) Heartbeat(ctx context.Context, payload protocol.NodePayload) error {
|
||||
@@ -45,6 +48,12 @@ func (f *fakeHeartbeatService) Heartbeat(ctx context.Context, payload protocol.N
|
||||
return err
|
||||
}
|
||||
|
||||
func (f *fakeHeartbeatService) SetToken(token string) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.lastToken = token
|
||||
}
|
||||
|
||||
type fakeSyncService struct {
|
||||
mu sync.Mutex
|
||||
startupErr error
|
||||
@@ -90,6 +99,7 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) {
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
AgentToken: "agent-token",
|
||||
NodeName: "edge-01",
|
||||
NodeIP: "10.0.0.8",
|
||||
AgentVersion: "0.1.0",
|
||||
@@ -106,8 +116,8 @@ func TestRunnerKeepsHeartbeatWhenStartupSyncFails(t *testing.T) {
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
if heartbeatService.registerCalls != 1 {
|
||||
t.Fatalf("expected 1 register call, got %d", heartbeatService.registerCalls)
|
||||
if heartbeatService.registerCalls != 0 {
|
||||
t.Fatalf("expected no discovery register call, got %d", heartbeatService.registerCalls)
|
||||
}
|
||||
if heartbeatService.heartbeatCalls < 2 {
|
||||
t.Fatalf("expected heartbeat loop to continue, got %d heartbeat calls", heartbeatService.heartbeatCalls)
|
||||
@@ -140,6 +150,7 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
AgentToken: "agent-token",
|
||||
NodeName: "edge-01",
|
||||
NodeIP: "10.0.0.8",
|
||||
AgentVersion: "0.1.0",
|
||||
@@ -156,8 +167,8 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
if heartbeatService.registerCalls != 1 {
|
||||
t.Fatalf("expected register attempt, got %d", heartbeatService.registerCalls)
|
||||
if heartbeatService.registerCalls != 0 {
|
||||
t.Fatalf("expected no register attempt, got %d", heartbeatService.registerCalls)
|
||||
}
|
||||
if syncService.syncOnceCalls == 0 {
|
||||
t.Fatal("expected sync loop to continue after heartbeat/register errors")
|
||||
@@ -170,3 +181,72 @@ func TestRunnerDoesNotExitOnHeartbeatOrSyncError(t *testing.T) {
|
||||
t.Fatalf("expected sync error to be recorded, got %q", snapshot.LastError)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerDiscoveryRegisterUpdatesTokenAndNodeID(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
heartbeatService := &fakeHeartbeatService{
|
||||
registerResp: &protocol.RegisterNodeResponse{
|
||||
NodeID: "node-server-assigned",
|
||||
AgentToken: "agent-token-issued",
|
||||
Name: "edge-01",
|
||||
},
|
||||
onHeartbeat: func(callCount int) {
|
||||
if callCount >= 1 {
|
||||
cancel()
|
||||
}
|
||||
},
|
||||
}
|
||||
syncService := &fakeSyncService{}
|
||||
configPath := filepath.Join(t.TempDir(), "agent.json")
|
||||
if err := os.WriteFile(configPath, []byte(`{"server_url":"http://127.0.0.1:3000","discovery_token":"discovery-token","node_name":"edge-01","node_ip":"10.0.0.8"}`), 0o644); err != nil {
|
||||
t.Fatalf("failed to seed config file: %v", err)
|
||||
}
|
||||
cfg, err := config.Load(configPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load config: %v", err)
|
||||
}
|
||||
runner := &Runner{
|
||||
Config: &config.Config{
|
||||
ServerURL: cfg.ServerURL,
|
||||
DiscoveryToken: cfg.DiscoveryToken,
|
||||
NodeName: cfg.NodeName,
|
||||
NodeIP: cfg.NodeIP,
|
||||
AgentVersion: "0.1.0",
|
||||
NginxVersion: "1.25.5",
|
||||
HeartbeatInterval: 10 * time.Millisecond,
|
||||
SyncInterval: 20 * time.Millisecond,
|
||||
},
|
||||
StateStore: stateStore,
|
||||
HeartbeatService: heartbeatService,
|
||||
SyncService: syncService,
|
||||
}
|
||||
runner.Config = cfg
|
||||
runner.Config.AgentVersion = "0.1.0"
|
||||
runner.Config.NginxVersion = "1.25.5"
|
||||
runner.Config.HeartbeatInterval = 10 * time.Millisecond
|
||||
runner.Config.SyncInterval = 20 * time.Millisecond
|
||||
|
||||
err = runner.Run(ctx)
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("expected context cancellation, got %v", err)
|
||||
}
|
||||
if heartbeatService.registerCalls == 0 {
|
||||
t.Fatal("expected discovery register to be attempted")
|
||||
}
|
||||
if heartbeatService.lastToken != "agent-token-issued" {
|
||||
t.Fatalf("expected client token to be updated, got %q", heartbeatService.lastToken)
|
||||
}
|
||||
snapshot, loadErr := stateStore.Load()
|
||||
if loadErr != nil {
|
||||
t.Fatalf("failed to load state: %v", loadErr)
|
||||
}
|
||||
if snapshot.NodeID != "node-server-assigned" {
|
||||
t.Fatalf("expected node id to be replaced, got %q", snapshot.NodeID)
|
||||
}
|
||||
if runner.Config.AgentToken != "agent-token-issued" || runner.Config.DiscoveryToken != "" {
|
||||
t.Fatal("expected config token rotation to complete")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package config
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net"
|
||||
"os"
|
||||
pathpkg "path"
|
||||
"path/filepath"
|
||||
@@ -20,6 +21,7 @@ const (
|
||||
type Config struct {
|
||||
ServerURL string `json:"server_url"`
|
||||
AgentToken string `json:"agent_token"`
|
||||
DiscoveryToken string `json:"discovery_token"`
|
||||
NodeName string `json:"node_name"`
|
||||
NodeIP string `json:"node_ip"`
|
||||
AgentVersion string `json:"agent_version"`
|
||||
@@ -36,6 +38,7 @@ type Config struct {
|
||||
HeartbeatInterval time.Duration `json:"heartbeat_interval"`
|
||||
SyncInterval time.Duration `json:"sync_interval"`
|
||||
RequestTimeout time.Duration `json:"request_timeout"`
|
||||
configPath string
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
@@ -47,6 +50,7 @@ func Load(path string) (*Config, error) {
|
||||
if err = json.Unmarshal(data, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cfg.configPath = path
|
||||
applyDefaults(cfg, filepath.Dir(path))
|
||||
if err = validate(cfg); err != nil {
|
||||
return nil, err
|
||||
@@ -71,6 +75,12 @@ func applyDefaults(cfg *Config, baseDir string) {
|
||||
if cfg.DataDir == "" {
|
||||
cfg.DataDir = filepath.Join(baseDir, "data")
|
||||
}
|
||||
if cfg.NodeName == "" {
|
||||
cfg.NodeName = detectHostname()
|
||||
}
|
||||
if cfg.NodeIP == "" {
|
||||
cfg.NodeIP = detectNodeIP()
|
||||
}
|
||||
if cfg.NginxPath == "" {
|
||||
cfg.RouteConfigPath = joinManagedPath(cfg.DataDir, defaultDockerRouteConfigRelativePath)
|
||||
cfg.StatePath = joinManagedPath(cfg.DataDir, defaultDockerStateRelativePath)
|
||||
@@ -137,8 +147,8 @@ func validate(cfg *Config) error {
|
||||
if cfg.ServerURL == "" {
|
||||
return errors.New("server_url 不能为空")
|
||||
}
|
||||
if cfg.AgentToken == "" {
|
||||
return errors.New("agent_token 不能为空")
|
||||
if strings.TrimSpace(cfg.AgentToken) == "" && strings.TrimSpace(cfg.DiscoveryToken) == "" {
|
||||
return errors.New("agent_token 和 discovery_token 不能同时为空")
|
||||
}
|
||||
if cfg.NodeName == "" {
|
||||
return errors.New("node_name 不能为空")
|
||||
@@ -148,3 +158,52 @@ func validate(cfg *Config) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cfg *Config) Save() error {
|
||||
if cfg == nil {
|
||||
return errors.New("config 不能为空")
|
||||
}
|
||||
if cfg.configPath == "" {
|
||||
return errors.New("config path 未初始化")
|
||||
}
|
||||
data, err := json.MarshalIndent(cfg, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(cfg.configPath, data, 0o644)
|
||||
}
|
||||
|
||||
func detectHostname() string {
|
||||
host, err := os.Hostname()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(host)
|
||||
}
|
||||
|
||||
func detectNodeIP() string {
|
||||
interfaces, err := net.Interfaces()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
for _, iface := range interfaces {
|
||||
if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 {
|
||||
continue
|
||||
}
|
||||
addrs, err := iface.Addrs()
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, addr := range addrs {
|
||||
ipNet, ok := addr.(*net.IPNet)
|
||||
if !ok || ipNet.IP == nil || ipNet.IP.IsLoopback() {
|
||||
continue
|
||||
}
|
||||
ipv4 := ipNet.IP.To4()
|
||||
if ipv4 != nil {
|
||||
return ipv4.String()
|
||||
}
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
@@ -7,8 +7,9 @@ import (
|
||||
)
|
||||
|
||||
type Client interface {
|
||||
RegisterNode(ctx context.Context, payload protocol.NodePayload) error
|
||||
RegisterNode(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error)
|
||||
Heartbeat(ctx context.Context, payload protocol.NodePayload) error
|
||||
SetToken(token string)
|
||||
}
|
||||
|
||||
type Service struct {
|
||||
@@ -19,10 +20,14 @@ func New(client Client) *Service {
|
||||
return &Service{client: client}
|
||||
}
|
||||
|
||||
func (s *Service) Register(ctx context.Context, payload protocol.NodePayload) error {
|
||||
func (s *Service) Register(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
|
||||
return s.client.RegisterNode(ctx, payload)
|
||||
}
|
||||
|
||||
func (s *Service) Heartbeat(ctx context.Context, payload protocol.NodePayload) error {
|
||||
return s.client.Heartbeat(ctx, payload)
|
||||
}
|
||||
|
||||
func (s *Service) SetToken(token string) {
|
||||
s.client.SetToken(token)
|
||||
}
|
||||
|
||||
@@ -29,9 +29,17 @@ func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload) error {
|
||||
func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload) (*protocol.RegisterNodeResponse, error) {
|
||||
log.Printf("http register node request: node_id=%s current_version=%s", payload.NodeID, payload.CurrentVersion)
|
||||
return c.postJSON(ctx, "/api/agent/nodes/register", payload, nil)
|
||||
resp := protocol.APIResponse[protocol.RegisterNodeResponse]{}
|
||||
if err := c.postJSON(ctx, "/api/agent/nodes/register", payload, &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.Success {
|
||||
return nil, errors.New(resp.Message)
|
||||
}
|
||||
log.Printf("http register node response: node_id=%s", resp.Data.NodeID)
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) error {
|
||||
@@ -56,6 +64,11 @@ func (c *Client) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPa
|
||||
return c.postJSON(ctx, "/api/agent/apply-logs", payload, nil)
|
||||
}
|
||||
|
||||
func (c *Client) SetToken(token string) {
|
||||
c.token = strings.TrimSpace(token)
|
||||
log.Printf("http client token updated")
|
||||
}
|
||||
|
||||
func (c *Client) getJSON(ctx context.Context, path string, target any) error {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.baseURL+path, nil)
|
||||
if err != nil {
|
||||
|
||||
@@ -16,6 +16,12 @@ type NodePayload struct {
|
||||
LastError string `json:"last_error"`
|
||||
}
|
||||
|
||||
type RegisterNodeResponse struct {
|
||||
NodeID string `json:"node_id"`
|
||||
AgentToken string `json:"agent_token"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
|
||||
type ApplyLogPayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Version string `json:"version"`
|
||||
|
||||
Reference in New Issue
Block a user