mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-28 05:46:36 +08:00
实现代理节点功能,包括主程序入口、配置加载、心跳检测和状态管理,添加相关服务和API响应结构
This commit is contained in:
@@ -1,3 +1,48 @@
|
||||
package main
|
||||
|
||||
func main() {}
|
||||
import (
|
||||
"context"
|
||||
"flag"
|
||||
"log"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
|
||||
"atsflare-agent/internal/agent"
|
||||
"atsflare-agent/internal/config"
|
||||
"atsflare-agent/internal/heartbeat"
|
||||
"atsflare-agent/internal/httpclient"
|
||||
"atsflare-agent/internal/nginx"
|
||||
"atsflare-agent/internal/state"
|
||||
syncservice "atsflare-agent/internal/sync"
|
||||
)
|
||||
|
||||
func main() {
|
||||
configPath := flag.String("config", "./agent.json", "agent config path")
|
||||
flag.Parse()
|
||||
|
||||
cfg, err := config.Load(*configPath)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
client := httpclient.New(cfg.ServerURL, cfg.AgentToken, cfg.RequestTimeout)
|
||||
stateStore := state.NewStore(cfg.StatePath)
|
||||
runner := &agent.Runner{
|
||||
Config: cfg,
|
||||
StateStore: stateStore,
|
||||
HeartbeatService: heartbeat.New(client),
|
||||
SyncService: syncservice.New(client, &nginx.Manager{
|
||||
RouteConfigPath: cfg.RouteConfigPath,
|
||||
Executor: &nginx.ShellExecutor{
|
||||
Binary: cfg.NginxBinary,
|
||||
},
|
||||
}, stateStore),
|
||||
}
|
||||
|
||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
|
||||
defer stop()
|
||||
|
||||
if err = runner.Run(ctx); err != nil && err != context.Canceled {
|
||||
log.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
package agent
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"atsflare-agent/internal/config"
|
||||
"atsflare-agent/internal/heartbeat"
|
||||
"atsflare-agent/internal/protocol"
|
||||
"atsflare-agent/internal/state"
|
||||
syncservice "atsflare-agent/internal/sync"
|
||||
)
|
||||
|
||||
type Runner struct {
|
||||
Config *config.Config
|
||||
StateStore *state.Store
|
||||
HeartbeatService *heartbeat.Service
|
||||
SyncService *syncservice.Service
|
||||
}
|
||||
|
||||
func (r *Runner) Run(ctx context.Context) error {
|
||||
nodeID, err := r.StateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err = r.HeartbeatService.Register(ctx, r.nodePayload(nodeID)); err != nil {
|
||||
return err
|
||||
}
|
||||
if err = r.SyncService.SyncOnce(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
heartbeatTicker := time.NewTicker(r.Config.HeartbeatInterval)
|
||||
defer heartbeatTicker.Stop()
|
||||
syncTicker := time.NewTicker(r.Config.SyncInterval)
|
||||
defer syncTicker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-heartbeatTicker.C:
|
||||
if err = r.HeartbeatService.Heartbeat(ctx, r.nodePayload(nodeID)); err != nil {
|
||||
return err
|
||||
}
|
||||
case <-syncTicker.C:
|
||||
if err = r.SyncService.SyncOnce(ctx); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) nodePayload(nodeID string) protocol.NodePayload {
|
||||
snapshot, _ := r.StateStore.Load()
|
||||
return protocol.NodePayload{
|
||||
NodeID: nodeID,
|
||||
Name: r.Config.NodeName,
|
||||
IP: r.Config.NodeIP,
|
||||
AgentVersion: r.Config.AgentVersion,
|
||||
NginxVersion: r.Config.NginxVersion,
|
||||
CurrentVersion: snapshot.CurrentVersion,
|
||||
LastError: snapshot.LastError,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
ServerURL string `json:"server_url"`
|
||||
AgentToken string `json:"agent_token"`
|
||||
NodeName string `json:"node_name"`
|
||||
NodeIP string `json:"node_ip"`
|
||||
AgentVersion string `json:"agent_version"`
|
||||
NginxVersion string `json:"nginx_version"`
|
||||
RouteConfigPath string `json:"route_config_path"`
|
||||
StatePath string `json:"state_path"`
|
||||
NginxBinary string `json:"nginx_binary"`
|
||||
HeartbeatInterval time.Duration `json:"heartbeat_interval"`
|
||||
SyncInterval time.Duration `json:"sync_interval"`
|
||||
RequestTimeout time.Duration `json:"request_timeout"`
|
||||
}
|
||||
|
||||
func Load(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cfg := &Config{}
|
||||
if err = json.Unmarshal(data, cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
applyDefaults(cfg)
|
||||
if err = validate(cfg); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func applyDefaults(cfg *Config) {
|
||||
if cfg.AgentVersion == "" {
|
||||
cfg.AgentVersion = "dev"
|
||||
}
|
||||
if cfg.RouteConfigPath == "" {
|
||||
cfg.RouteConfigPath = filepath.Clean("./atsflare_routes.conf")
|
||||
}
|
||||
if cfg.StatePath == "" {
|
||||
cfg.StatePath = filepath.Clean("./atsf_agent_state.json")
|
||||
}
|
||||
if cfg.NginxBinary == "" {
|
||||
cfg.NginxBinary = "nginx"
|
||||
}
|
||||
if cfg.HeartbeatInterval <= 0 {
|
||||
cfg.HeartbeatInterval = 30 * time.Second
|
||||
}
|
||||
if cfg.SyncInterval <= 0 {
|
||||
cfg.SyncInterval = 30 * time.Second
|
||||
}
|
||||
if cfg.RequestTimeout <= 0 {
|
||||
cfg.RequestTimeout = 10 * time.Second
|
||||
}
|
||||
}
|
||||
|
||||
func validate(cfg *Config) error {
|
||||
if cfg.ServerURL == "" {
|
||||
return errors.New("server_url 不能为空")
|
||||
}
|
||||
if cfg.AgentToken == "" {
|
||||
return errors.New("agent_token 不能为空")
|
||||
}
|
||||
if cfg.NodeName == "" {
|
||||
return errors.New("node_name 不能为空")
|
||||
}
|
||||
if cfg.NodeIP == "" {
|
||||
return errors.New("node_ip 不能为空")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package heartbeat
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"atsflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
type Client interface {
|
||||
RegisterNode(ctx context.Context, payload protocol.NodePayload) error
|
||||
Heartbeat(ctx context.Context, payload protocol.NodePayload) error
|
||||
}
|
||||
|
||||
type Service struct {
|
||||
client Client
|
||||
}
|
||||
|
||||
func New(client Client) *Service {
|
||||
return &Service{client: client}
|
||||
}
|
||||
|
||||
func (s *Service) Register(ctx context.Context, payload protocol.NodePayload) error {
|
||||
return s.client.RegisterNode(ctx, payload)
|
||||
}
|
||||
|
||||
func (s *Service) Heartbeat(ctx context.Context, payload protocol.NodePayload) error {
|
||||
return s.client.Heartbeat(ctx, payload)
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package httpclient
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"atsflare-agent/internal/protocol"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
baseURL string
|
||||
token string
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func New(baseURL string, token string, timeout time.Duration) *Client {
|
||||
return &Client{
|
||||
baseURL: strings.TrimRight(baseURL, "/"),
|
||||
token: token,
|
||||
httpClient: &http.Client{
|
||||
Timeout: timeout,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Client) RegisterNode(ctx context.Context, payload protocol.NodePayload) error {
|
||||
return c.postJSON(ctx, "/api/agent/nodes/register", payload, nil)
|
||||
}
|
||||
|
||||
func (c *Client) Heartbeat(ctx context.Context, payload protocol.NodePayload) error {
|
||||
return c.postJSON(ctx, "/api/agent/nodes/heartbeat", payload, nil)
|
||||
}
|
||||
|
||||
func (c *Client) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) {
|
||||
resp := protocol.APIResponse[protocol.ActiveConfigResponse]{}
|
||||
if err := c.getJSON(ctx, "/api/agent/config-versions/active", &resp); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !resp.Success {
|
||||
return nil, errors.New(resp.Message)
|
||||
}
|
||||
return &resp.Data, nil
|
||||
}
|
||||
|
||||
func (c *Client) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
|
||||
return c.postJSON(ctx, "/api/agent/apply-logs", payload, nil)
|
||||
}
|
||||
|
||||
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 {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("X-Agent-Token", c.token)
|
||||
return c.do(req, target)
|
||||
}
|
||||
|
||||
func (c *Client) postJSON(ctx context.Context, path string, body any, target any) error {
|
||||
data, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.baseURL+path, bytes.NewReader(data))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-Agent-Token", c.token)
|
||||
return c.do(req, target)
|
||||
}
|
||||
|
||||
func (c *Client) do(req *http.Request, target any) error {
|
||||
res, err := c.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
if res.StatusCode != http.StatusOK {
|
||||
return errors.New(res.Status)
|
||||
}
|
||||
if target == nil {
|
||||
var wrapper protocol.APIResponse[json.RawMessage]
|
||||
if err = json.NewDecoder(res.Body).Decode(&wrapper); err != nil {
|
||||
return err
|
||||
}
|
||||
if !wrapper.Success {
|
||||
return errors.New(wrapper.Message)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return json.NewDecoder(res.Body).Decode(target)
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package nginx
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
type Executor interface {
|
||||
Test(ctx context.Context) error
|
||||
Reload(ctx context.Context) error
|
||||
}
|
||||
|
||||
type ShellExecutor struct {
|
||||
Binary string
|
||||
}
|
||||
|
||||
func (e *ShellExecutor) Test(ctx context.Context) error {
|
||||
cmd := exec.CommandContext(ctx, e.Binary, "-t")
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("nginx -t failed: %w: %s", err, string(output))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *ShellExecutor) Reload(ctx context.Context) error {
|
||||
cmd := exec.CommandContext(ctx, e.Binary, "-s", "reload")
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return fmt.Errorf("nginx reload failed: %w: %s", err, string(output))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type Manager struct {
|
||||
RouteConfigPath string
|
||||
Executor Executor
|
||||
}
|
||||
|
||||
func (m *Manager) Apply(ctx context.Context, content string) error {
|
||||
backupPath, hadExisting, err := m.backup()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err = os.WriteFile(m.RouteConfigPath, []byte(content), 0o644); err != nil {
|
||||
return err
|
||||
}
|
||||
if err = m.Executor.Test(ctx); err != nil {
|
||||
_ = m.restore(backupPath, hadExisting)
|
||||
return err
|
||||
}
|
||||
if err = m.Executor.Reload(ctx); err != nil {
|
||||
_ = m.restore(backupPath, hadExisting)
|
||||
return err
|
||||
}
|
||||
if backupPath != "" {
|
||||
_ = os.Remove(backupPath)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) backup() (string, bool, error) {
|
||||
if m.RouteConfigPath == "" {
|
||||
return "", false, errors.New("route config path 不能为空")
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(m.RouteConfigPath), 0o755); err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
data, err := os.ReadFile(m.RouteConfigPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", false, nil
|
||||
}
|
||||
return "", false, err
|
||||
}
|
||||
backupPath := m.RouteConfigPath + ".bak"
|
||||
if err = os.WriteFile(backupPath, data, 0o644); err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return backupPath, true, nil
|
||||
}
|
||||
|
||||
func (m *Manager) restore(backupPath string, hadExisting bool) error {
|
||||
if hadExisting {
|
||||
data, err := os.ReadFile(backupPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(m.RouteConfigPath, data, 0o644)
|
||||
}
|
||||
if err := os.Remove(m.RouteConfigPath); err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -1,5 +1,11 @@
|
||||
package protocol
|
||||
|
||||
type APIResponse[T any] struct {
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
Data T `json:"data"`
|
||||
}
|
||||
|
||||
type NodePayload struct {
|
||||
NodeID string `json:"node_id"`
|
||||
Name string `json:"name"`
|
||||
|
||||
@@ -0,0 +1,96 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type Snapshot struct {
|
||||
NodeID string `json:"node_id"`
|
||||
CurrentVersion string `json:"current_version"`
|
||||
CurrentChecksum string `json:"current_checksum"`
|
||||
LastError string `json:"last_error"`
|
||||
}
|
||||
|
||||
type Store struct {
|
||||
path string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewStore(path string) *Store {
|
||||
return &Store{path: filepath.Clean(path)}
|
||||
}
|
||||
|
||||
func (s *Store) Load() (*Snapshot, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.loadUnlocked()
|
||||
}
|
||||
|
||||
func (s *Store) EnsureNodeID() (string, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
snapshot, err := s.loadUnlocked()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if snapshot.NodeID != "" {
|
||||
return snapshot.NodeID, nil
|
||||
}
|
||||
snapshot.NodeID, err = newNodeID()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err = s.saveUnlocked(snapshot); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return snapshot.NodeID, nil
|
||||
}
|
||||
|
||||
func (s *Store) Save(snapshot *Snapshot) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.saveUnlocked(snapshot)
|
||||
}
|
||||
|
||||
func (s *Store) loadUnlocked() (*Snapshot, error) {
|
||||
data, err := os.ReadFile(s.path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return &Snapshot{}, nil
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
snapshot := &Snapshot{}
|
||||
if len(data) == 0 {
|
||||
return snapshot, nil
|
||||
}
|
||||
if err = json.Unmarshal(data, snapshot); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
|
||||
func (s *Store) saveUnlocked(snapshot *Snapshot) error {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(snapshot, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(s.path, data, 0o644)
|
||||
}
|
||||
|
||||
func newNodeID() (string, error) {
|
||||
buf := make([]byte, 8)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "node-" + hex.EncodeToString(buf), nil
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package state
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestEnsureNodeIDPersists(t *testing.T) {
|
||||
store := NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
nodeID1, err := store.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
nodeID2, err := store.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID second call failed: %v", err)
|
||||
}
|
||||
if nodeID1 == "" || nodeID1 != nodeID2 {
|
||||
t.Fatal("expected node id to persist across calls")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package sync
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"atsflare-agent/internal/nginx"
|
||||
"atsflare-agent/internal/protocol"
|
||||
"atsflare-agent/internal/state"
|
||||
)
|
||||
|
||||
const (
|
||||
ApplyResultSuccess = "success"
|
||||
ApplyResultFailed = "failed"
|
||||
)
|
||||
|
||||
type ConfigClient interface {
|
||||
GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error)
|
||||
ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error
|
||||
}
|
||||
|
||||
type Service struct {
|
||||
client ConfigClient
|
||||
nginxManager *nginx.Manager
|
||||
stateStore *state.Store
|
||||
}
|
||||
|
||||
func New(client ConfigClient, nginxManager *nginx.Manager, stateStore *state.Store) *Service {
|
||||
return &Service{
|
||||
client: client,
|
||||
nginxManager: nginxManager,
|
||||
stateStore: stateStore,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) SyncOnce(ctx context.Context) error {
|
||||
snapshot, err := s.stateStore.Load()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
config, err := s.client.GetActiveConfig(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if snapshot.CurrentVersion == config.Version && snapshot.CurrentChecksum == config.Checksum {
|
||||
return nil
|
||||
}
|
||||
if err = s.nginxManager.Apply(ctx, config.RenderedConfig); err != nil {
|
||||
snapshot.LastError = err.Error()
|
||||
_ = s.stateStore.Save(snapshot)
|
||||
reportErr := s.client.ReportApplyLog(ctx, protocol.ApplyLogPayload{
|
||||
NodeID: snapshot.NodeID,
|
||||
Version: config.Version,
|
||||
Result: ApplyResultFailed,
|
||||
Message: err.Error(),
|
||||
})
|
||||
if reportErr != nil {
|
||||
return reportErr
|
||||
}
|
||||
return err
|
||||
}
|
||||
snapshot.CurrentVersion = config.Version
|
||||
snapshot.CurrentChecksum = config.Checksum
|
||||
snapshot.LastError = ""
|
||||
if err = s.stateStore.Save(snapshot); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.client.ReportApplyLog(ctx, protocol.ApplyLogPayload{
|
||||
NodeID: snapshot.NodeID,
|
||||
Version: config.Version,
|
||||
Result: ApplyResultSuccess,
|
||||
Message: "apply success",
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package sync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"atsflare-agent/internal/nginx"
|
||||
"atsflare-agent/internal/protocol"
|
||||
"atsflare-agent/internal/state"
|
||||
)
|
||||
|
||||
type fakeExecutor struct {
|
||||
testErr error
|
||||
reloadErr error
|
||||
}
|
||||
|
||||
type fakeClient struct {
|
||||
config protocol.ActiveConfigResponse
|
||||
reports []protocol.ApplyLogPayload
|
||||
}
|
||||
|
||||
func (f *fakeExecutor) Test(ctx context.Context) error {
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
func (f *fakeExecutor) Reload(ctx context.Context) error {
|
||||
return f.reloadErr
|
||||
}
|
||||
|
||||
func (f *fakeClient) GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error) {
|
||||
return &f.config, nil
|
||||
}
|
||||
|
||||
func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
|
||||
f.reports = append(f.reports, payload)
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestSyncOnceSuccess(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
config: protocol.ActiveConfigResponse{
|
||||
Version: "20260309-001",
|
||||
Checksum: "checksum-1",
|
||||
RenderedConfig: "server { listen 80; }",
|
||||
CreatedAt: time.Now().Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(t.TempDir(), "state.json"))
|
||||
nodeID, err := stateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
snapshot, _ := stateStore.Load()
|
||||
snapshot.NodeID = nodeID
|
||||
if err = stateStore.Save(snapshot); err != nil {
|
||||
t.Fatalf("failed to save initial state: %v", err)
|
||||
}
|
||||
|
||||
routePath := filepath.Join(t.TempDir(), "routes.conf")
|
||||
service := New(client, &nginx.Manager{
|
||||
RouteConfigPath: routePath,
|
||||
Executor: &fakeExecutor{},
|
||||
}, stateStore)
|
||||
|
||||
if err = service.SyncOnce(context.Background()); err != nil {
|
||||
t.Fatalf("SyncOnce failed: %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(routePath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read route config: %v", err)
|
||||
}
|
||||
if string(data) != "server { listen 80; }" {
|
||||
t.Fatal("expected rendered config to be written to route file")
|
||||
}
|
||||
snapshot, err = stateStore.Load()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to load state: %v", err)
|
||||
}
|
||||
if snapshot.CurrentVersion != "20260309-001" || snapshot.CurrentChecksum != "checksum-1" {
|
||||
t.Fatal("expected state store to persist current version and checksum")
|
||||
}
|
||||
if len(client.reports) != 1 || client.reports[0].Result != ApplyResultSuccess {
|
||||
t.Fatal("expected successful apply report to be sent")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncOnceRollbackOnNginxFailure(t *testing.T) {
|
||||
client := &fakeClient{
|
||||
config: protocol.ActiveConfigResponse{
|
||||
Version: "20260309-002",
|
||||
Checksum: "checksum-2",
|
||||
RenderedConfig: "server { listen 81; }",
|
||||
CreatedAt: time.Now().Format(time.RFC3339),
|
||||
},
|
||||
}
|
||||
|
||||
tempDir := t.TempDir()
|
||||
routePath := filepath.Join(tempDir, "routes.conf")
|
||||
if err := os.WriteFile(routePath, []byte("server { listen 80; }"), 0o644); err != nil {
|
||||
t.Fatalf("failed to seed route file: %v", err)
|
||||
}
|
||||
|
||||
stateStore := state.NewStore(filepath.Join(tempDir, "state.json"))
|
||||
nodeID, err := stateStore.EnsureNodeID()
|
||||
if err != nil {
|
||||
t.Fatalf("EnsureNodeID failed: %v", err)
|
||||
}
|
||||
if err = stateStore.Save(&state.Snapshot{
|
||||
NodeID: nodeID,
|
||||
CurrentVersion: "20260309-001",
|
||||
CurrentChecksum: "checksum-1",
|
||||
}); err != nil {
|
||||
t.Fatalf("failed to seed state: %v", err)
|
||||
}
|
||||
|
||||
service := New(client, &nginx.Manager{
|
||||
RouteConfigPath: routePath,
|
||||
Executor: &fakeExecutor{
|
||||
testErr: context.DeadlineExceeded,
|
||||
},
|
||||
}, stateStore)
|
||||
|
||||
err = service.SyncOnce(context.Background())
|
||||
if err == nil {
|
||||
t.Fatal("expected SyncOnce to fail when nginx test fails")
|
||||
}
|
||||
|
||||
data, readErr := os.ReadFile(routePath)
|
||||
if readErr != nil {
|
||||
t.Fatalf("failed to read route file after rollback: %v", readErr)
|
||||
}
|
||||
if string(data) != "server { listen 80; }" {
|
||||
t.Fatal("expected original route config to be restored after rollback")
|
||||
}
|
||||
snapshot, loadErr := stateStore.Load()
|
||||
if loadErr != nil {
|
||||
t.Fatalf("failed to load state: %v", loadErr)
|
||||
}
|
||||
if snapshot.CurrentVersion != "20260309-001" {
|
||||
t.Fatal("expected failed sync not to overwrite current version")
|
||||
}
|
||||
if len(client.reports) != 1 || client.reports[0].Result != ApplyResultFailed {
|
||||
t.Fatal("expected failed apply report to be sent")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user