feat: complete SyncServer implementation
Full-stack Go monolith with embedded React frontend for orchestrating rsync-over-SSH file synchronization with Wake-on-LAN support. Features: - JWT auth (HS256) with bcrypt password hashing - CRUD for machines (with WoL config) and sync_pairs - Ed25519 SSH key generation and known_hosts management - WoL magic packet sender + TCP-connect waiter with backoff - Sync engine: rsync subprocess, per-pair job queue, progress parsing - Homebrew cron parser for scheduled syncs - SSE stream for live job status (queued/waking_up/running/success/failed) - React+TS+Vite+Tailwind SPA embedded via embed.FS - Debian packaging with systemd unit, postinst/prerm/postrm Tech stack: - Go 1.22+ (CGO_ENABLED=0, pure SQLite via modernc.org/sqlite) - chi router for HTTP API - TypeScript + React 18 + Tailwind CSS frontend - Cross-compiled to Linux amd64 for Proxmox LXC deployment Tests: wol (MAC parsing, magic packet), syncengine/queue, scheduler/cron
This commit is contained in:
@@ -0,0 +1,69 @@
|
||||
package api
|
||||
|
||||
type MachineRequest struct {
|
||||
Name string `json:"name"`
|
||||
Host string `json:"host"`
|
||||
Port int `json:"port"`
|
||||
SSHUser string `json:"ssh_user"`
|
||||
SSHKeyID *int64 `json:"ssh_key_id"`
|
||||
MACAddress *string `json:"mac_address"`
|
||||
WoLEnabled bool `json:"wol_enabled"`
|
||||
BroadcastAddr *string `json:"broadcast_addr"`
|
||||
WakeTimeoutSeconds int `json:"wake_timeout_seconds"`
|
||||
WakeCheckIntervalSeconds int `json:"wake_check_interval_seconds"`
|
||||
}
|
||||
|
||||
type MachineResponse struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Host string `json:"host"`
|
||||
Port int `json:"port"`
|
||||
SSHUser string `json:"ssh_user"`
|
||||
SSHKeyID *int64 `json:"ssh_key_id"`
|
||||
MACAddress *string `json:"mac_address"`
|
||||
WoLEnabled bool `json:"wol_enabled"`
|
||||
BroadcastAddr *string `json:"broadcast_addr"`
|
||||
WakeTimeoutSeconds int `json:"wake_timeout_seconds"`
|
||||
WakeCheckIntervalSeconds int `json:"wake_check_interval_seconds"`
|
||||
FingerprintConfirmed bool `json:"fingerprint_confirmed"`
|
||||
Status string `json:"status"`
|
||||
}
|
||||
|
||||
type SyncPairRequest struct {
|
||||
Name string `json:"name"`
|
||||
SourceMachineID *int64 `json:"source_machine_id"`
|
||||
SourcePath string `json:"source_path"`
|
||||
DestMachineID *int64 `json:"dest_machine_id"`
|
||||
DestPath string `json:"dest_path"`
|
||||
Direction string `json:"direction"`
|
||||
RsyncFlags string `json:"rsync_flags"`
|
||||
ExcludePatterns string `json:"exclude_patterns"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type SyncPairResponse struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
SourceMachineID *int64 `json:"source_machine_id"`
|
||||
SourcePath string `json:"source_path"`
|
||||
DestMachineID *int64 `json:"dest_machine_id"`
|
||||
DestPath string `json:"dest_path"`
|
||||
Direction string `json:"direction"`
|
||||
RsyncFlags string `json:"rsync_flags"`
|
||||
ExcludePatterns string `json:"exclude_patterns"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type JobResponse struct {
|
||||
ID int64 `json:"id"`
|
||||
SyncPairID int64 `json:"sync_pair_id"`
|
||||
TriggerType string `json:"trigger_type"`
|
||||
Status string `json:"status"`
|
||||
StartedAt *string `json:"started_at"`
|
||||
FinishedAt *string `json:"finished_at"`
|
||||
LogFile *string `json:"log_file"`
|
||||
}
|
||||
|
||||
type ErrorResponse struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
|
||||
"github.com/syncserver/internal/auth"
|
||||
)
|
||||
|
||||
type AuthHandler struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewAuthHandler(db *sql.DB) *AuthHandler {
|
||||
return &AuthHandler{db: db}
|
||||
}
|
||||
|
||||
type LoginRequest struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
type LoginResponse struct {
|
||||
User UserResponse `json:"user"`
|
||||
ExpiresAt string `json:"expires_at"`
|
||||
}
|
||||
|
||||
type UserResponse struct {
|
||||
ID int64 `json:"id"`
|
||||
Username string `json:"username"`
|
||||
Role string `json:"role"`
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) {
|
||||
var req LoginRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, `{"error":"invalid request"}`, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
row := h.db.QueryRow(
|
||||
"SELECT id, username, password_hash, role FROM users WHERE username = ?",
|
||||
req.Username,
|
||||
)
|
||||
var u struct {
|
||||
ID int64
|
||||
Username string
|
||||
PasswordHash string
|
||||
Role string
|
||||
}
|
||||
if err := row.Scan(&u.ID, &u.Username, &u.PasswordHash, &u.Role); err != nil {
|
||||
http.Error(w, `{"error":"invalid credentials"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
if !auth.VerifyPassword([]byte(u.PasswordHash), req.Password) {
|
||||
http.Error(w, `{"error":"invalid credentials"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
jwtMgr := auth.GetJWTManager()
|
||||
if jwtMgr == nil {
|
||||
http.Error(w, `{"error":"server misconfigured"}`, http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
token, expiresAt, err := jwtMgr.Generate(u.ID, u.Username, u.Role)
|
||||
if err != nil {
|
||||
http.Error(w, `{"error":"failed to generate token"}`, http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
auth.SetAuthCookie(w, token, expiresAt)
|
||||
|
||||
resp := LoginResponse{
|
||||
User: UserResponse{
|
||||
ID: u.ID,
|
||||
Username: u.Username,
|
||||
Role: u.Role,
|
||||
},
|
||||
ExpiresAt: expiresAt.Format("2006-01-02T15:04:05Z07:00"),
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Logout(w http.ResponseWriter, r *http.Request) {
|
||||
auth.ClearAuthCookie(w)
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Me(w http.ResponseWriter, r *http.Request) {
|
||||
claims := auth.GetClaims(r.Context())
|
||||
if claims == nil {
|
||||
http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
resp := UserResponse{
|
||||
ID: claims.UserID,
|
||||
Username: claims.Username,
|
||||
Role: claims.Role,
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{"user": resp})
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"net/http"
|
||||
"os"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/syncserver/internal/models"
|
||||
"github.com/syncserver/internal/syncengine"
|
||||
)
|
||||
|
||||
type JobHandler struct {
|
||||
db *sql.DB
|
||||
engine *syncengine.Engine
|
||||
}
|
||||
|
||||
func NewJobHandler(db *sql.DB, engine *syncengine.Engine) *JobHandler {
|
||||
return &JobHandler{db: db, engine: engine}
|
||||
}
|
||||
|
||||
func (h *JobHandler) List(w http.ResponseWriter, r *http.Request) {
|
||||
limit, _ := strconv.Atoi(r.URL.Query().Get("limit"))
|
||||
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
|
||||
if limit <= 0 || limit > 100 {
|
||||
limit = 50
|
||||
}
|
||||
|
||||
repo := models.NewJobRepository(h.db)
|
||||
jobs, err := repo.GetAll(limit, offset)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch jobs")
|
||||
return
|
||||
}
|
||||
|
||||
out := make([]JobResponse, len(jobs))
|
||||
for i, j := range jobs {
|
||||
out[i] = jobToResp(j)
|
||||
}
|
||||
writeJSON(w, out)
|
||||
}
|
||||
|
||||
func (h *JobHandler) Get(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
repo := models.NewJobRepository(h.db)
|
||||
j, err := repo.GetByID(id)
|
||||
if err == sql.ErrNoRows {
|
||||
writeError(w, http.StatusNotFound, "job not found")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch job")
|
||||
return
|
||||
}
|
||||
writeJSON(w, jobToResp(*j))
|
||||
}
|
||||
|
||||
func (h *JobHandler) Cancel(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
repo := models.NewJobRepository(h.db)
|
||||
j, err := repo.GetByID(id)
|
||||
if err == sql.ErrNoRows {
|
||||
writeError(w, http.StatusNotFound, "job not found")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch job")
|
||||
return
|
||||
}
|
||||
|
||||
if j.Status != "queued" && j.Status != "waking_up" && j.Status != "running" {
|
||||
writeError(w, http.StatusBadRequest, "job is not cancellable")
|
||||
return
|
||||
}
|
||||
|
||||
if h.engine != nil {
|
||||
h.engine.Cancel(id, j.SyncPairID)
|
||||
}
|
||||
|
||||
repo.UpdateStatus(id, "cancelled")
|
||||
writeJSON(w, map[string]string{"status": "cancelled"})
|
||||
}
|
||||
|
||||
func (h *JobHandler) TriggerRun(w http.ResponseWriter, r *http.Request) {
|
||||
pairID, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
if h.engine == nil {
|
||||
writeError(w, http.StatusInternalServerError, "engine not available")
|
||||
return
|
||||
}
|
||||
|
||||
jobID, err := h.engine.CreateJob(pairID, "manual")
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to create job")
|
||||
return
|
||||
}
|
||||
|
||||
go func() {
|
||||
h.engine.Run(r.Context(), jobID, pairID)
|
||||
}()
|
||||
|
||||
jobRepo := models.NewJobRepository(h.db)
|
||||
j, _ := jobRepo.GetByID(jobID)
|
||||
writeJSON(w, jobToResp(*j), http.StatusCreated)
|
||||
}
|
||||
|
||||
func (h *JobHandler) StreamLog(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
flusher, ok := w.(http.Flusher)
|
||||
if !ok {
|
||||
writeError(w, http.StatusInternalServerError, "streaming not supported")
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Connection", "keep-alive")
|
||||
flusher.Flush()
|
||||
|
||||
jobRepo := models.NewJobRepository(h.db)
|
||||
j, err := jobRepo.GetByID(id)
|
||||
if err == nil && j.LogFile != nil {
|
||||
data, _ := os.ReadFile(*j.LogFile)
|
||||
w.Write(data)
|
||||
flusher.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
func jobToResp(j models.Job) JobResponse {
|
||||
resp := JobResponse{
|
||||
ID: j.ID,
|
||||
SyncPairID: j.SyncPairID,
|
||||
TriggerType: j.TriggerType,
|
||||
Status: j.Status,
|
||||
LogFile: j.LogFile,
|
||||
}
|
||||
if j.StartedAt != nil {
|
||||
s := j.StartedAt.Format(time.RFC3339)
|
||||
resp.StartedAt = &s
|
||||
}
|
||||
if j.FinishedAt != nil {
|
||||
s := j.FinishedAt.Format(time.RFC3339)
|
||||
resp.FinishedAt = &s
|
||||
}
|
||||
return resp
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/syncserver/internal/models"
|
||||
)
|
||||
|
||||
type MachineHandler struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewMachineHandler(db *sql.DB) *MachineHandler {
|
||||
return &MachineHandler{db: db}
|
||||
}
|
||||
|
||||
var macRegex = regexp.MustCompile(`^([0-9A-Fa-f]{2}[:-]){5}[0-9A-Fa-f]{2}$`)
|
||||
|
||||
func (h *MachineHandler) List(w http.ResponseWriter, r *http.Request) {
|
||||
repo := models.NewMachineRepository(h.db)
|
||||
ms, err := repo.GetAll()
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch machines")
|
||||
return
|
||||
}
|
||||
out := make([]MachineResponse, len(ms))
|
||||
for i, m := range ms {
|
||||
out[i] = machineToResp(m)
|
||||
}
|
||||
writeJSON(w, out)
|
||||
}
|
||||
|
||||
func (h *MachineHandler) Get(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
repo := models.NewMachineRepository(h.db)
|
||||
m, err := repo.GetByID(id)
|
||||
if err == sql.ErrNoRows {
|
||||
writeError(w, http.StatusNotFound, "machine not found")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch machine")
|
||||
return
|
||||
}
|
||||
writeJSON(w, machineToResp(*m))
|
||||
}
|
||||
|
||||
func (h *MachineHandler) Create(w http.ResponseWriter, r *http.Request) {
|
||||
var req MachineRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid request body")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" || req.Host == "" {
|
||||
writeError(w, http.StatusBadRequest, "name and host are required")
|
||||
return
|
||||
}
|
||||
if req.Port <= 0 || req.Port > 65535 {
|
||||
writeError(w, http.StatusBadRequest, "invalid port")
|
||||
return
|
||||
}
|
||||
if req.SSHUser == "" {
|
||||
req.SSHUser = "root"
|
||||
}
|
||||
if req.WakeTimeoutSeconds <= 0 {
|
||||
req.WakeTimeoutSeconds = 120
|
||||
}
|
||||
if req.WakeCheckIntervalSeconds <= 0 {
|
||||
req.WakeCheckIntervalSeconds = 5
|
||||
}
|
||||
if req.WoLEnabled && req.MACAddress != nil && !macRegex.MatchString(*req.MACAddress) {
|
||||
writeError(w, http.StatusBadRequest, "invalid mac_address format (expected AA:BB:CC:DD:EE:FF)")
|
||||
return
|
||||
}
|
||||
|
||||
m := &models.Machine{
|
||||
Name: req.Name,
|
||||
Host: req.Host,
|
||||
Port: req.Port,
|
||||
SSHUser: req.SSHUser,
|
||||
SSHKeyID: req.SSHKeyID,
|
||||
MACAddress: req.MACAddress,
|
||||
WoLEnabled: req.WoLEnabled,
|
||||
BroadcastAddr: req.BroadcastAddr,
|
||||
WakeTimeoutSeconds: req.WakeTimeoutSeconds,
|
||||
WakeCheckIntervalSeconds: req.WakeCheckIntervalSeconds,
|
||||
FingerprintConfirmed: false,
|
||||
Status: "unknown",
|
||||
}
|
||||
|
||||
repo := models.NewMachineRepository(h.db)
|
||||
id, err := repo.Create(m)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to create machine")
|
||||
return
|
||||
}
|
||||
m.ID = id
|
||||
w.Header().Set("Location", "/api/machines/"+strconv.FormatInt(id, 10))
|
||||
writeJSON(w, machineToResp(*m), http.StatusCreated)
|
||||
}
|
||||
|
||||
func (h *MachineHandler) Update(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
var req MachineRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid request body")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" || req.Host == "" {
|
||||
writeError(w, http.StatusBadRequest, "name and host are required")
|
||||
return
|
||||
}
|
||||
if req.Port <= 0 || req.Port > 65535 {
|
||||
writeError(w, http.StatusBadRequest, "invalid port")
|
||||
return
|
||||
}
|
||||
if req.SSHUser == "" {
|
||||
req.SSHUser = "root"
|
||||
}
|
||||
if req.WoLEnabled && req.MACAddress != nil && !macRegex.MatchString(*req.MACAddress) {
|
||||
writeError(w, http.StatusBadRequest, "invalid mac_address format")
|
||||
return
|
||||
}
|
||||
|
||||
repo := models.NewMachineRepository(h.db)
|
||||
existing, err := repo.GetByID(id)
|
||||
if err == sql.ErrNoRows {
|
||||
writeError(w, http.StatusNotFound, "machine not found")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch machine")
|
||||
return
|
||||
}
|
||||
|
||||
existing.Name = req.Name
|
||||
existing.Host = req.Host
|
||||
existing.Port = req.Port
|
||||
existing.SSHUser = req.SSHUser
|
||||
existing.SSHKeyID = req.SSHKeyID
|
||||
existing.MACAddress = req.MACAddress
|
||||
existing.WoLEnabled = req.WoLEnabled
|
||||
existing.BroadcastAddr = req.BroadcastAddr
|
||||
if req.WakeTimeoutSeconds > 0 {
|
||||
existing.WakeTimeoutSeconds = req.WakeTimeoutSeconds
|
||||
}
|
||||
if req.WakeCheckIntervalSeconds > 0 {
|
||||
existing.WakeCheckIntervalSeconds = req.WakeCheckIntervalSeconds
|
||||
}
|
||||
|
||||
if err := repo.Update(existing); err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to update machine")
|
||||
return
|
||||
}
|
||||
writeJSON(w, machineToResp(*existing))
|
||||
}
|
||||
|
||||
func (h *MachineHandler) Delete(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
repo := models.NewMachineRepository(h.db)
|
||||
if err := repo.Delete(id); err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to delete machine")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func machineToResp(m models.Machine) MachineResponse {
|
||||
var status string
|
||||
if m.LastSeenAt != nil {
|
||||
status = m.Status + " (last seen " + m.LastSeenAt.Format(time.RFC3339) + ")"
|
||||
} else {
|
||||
status = m.Status
|
||||
}
|
||||
return MachineResponse{
|
||||
ID: m.ID,
|
||||
Name: m.Name,
|
||||
Host: m.Host,
|
||||
Port: m.Port,
|
||||
SSHUser: m.SSHUser,
|
||||
SSHKeyID: m.SSHKeyID,
|
||||
MACAddress: m.MACAddress,
|
||||
WoLEnabled: m.WoLEnabled,
|
||||
BroadcastAddr: m.BroadcastAddr,
|
||||
WakeTimeoutSeconds: m.WakeTimeoutSeconds,
|
||||
WakeCheckIntervalSeconds: m.WakeCheckIntervalSeconds,
|
||||
FingerprintConfirmed: m.FingerprintConfirmed,
|
||||
Status: status,
|
||||
}
|
||||
}
|
||||
|
||||
func writeError(w http.ResponseWriter, code int, msg string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(code)
|
||||
json.NewEncoder(w).Encode(ErrorResponse{Error: msg})
|
||||
}
|
||||
|
||||
func writeJSON(w http.ResponseWriter, data interface{}, codes ...int) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
if len(codes) > 0 {
|
||||
w.WriteHeader(codes[0])
|
||||
}
|
||||
json.NewEncoder(w).Encode(data)
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strconv"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/syncserver/internal/models"
|
||||
)
|
||||
|
||||
type SyncPairHandler struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewSyncPairHandler(db *sql.DB) *SyncPairHandler {
|
||||
return &SyncPairHandler{db: db}
|
||||
}
|
||||
|
||||
var directionRegex = regexp.MustCompile(`^(push|pull|mirror)$`)
|
||||
|
||||
func (h *SyncPairHandler) List(w http.ResponseWriter, r *http.Request) {
|
||||
repo := models.NewSyncPairRepository(h.db)
|
||||
pairs, err := repo.GetAll()
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch sync pairs")
|
||||
return
|
||||
}
|
||||
out := make([]SyncPairResponse, len(pairs))
|
||||
for i, p := range pairs {
|
||||
out[i] = syncPairToResp(p)
|
||||
}
|
||||
writeJSON(w, out)
|
||||
}
|
||||
|
||||
func (h *SyncPairHandler) Get(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
repo := models.NewSyncPairRepository(h.db)
|
||||
p, err := repo.GetByID(id)
|
||||
if err == sql.ErrNoRows {
|
||||
writeError(w, http.StatusNotFound, "sync pair not found")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch sync pair")
|
||||
return
|
||||
}
|
||||
writeJSON(w, syncPairToResp(*p))
|
||||
}
|
||||
|
||||
func (h *SyncPairHandler) Create(w http.ResponseWriter, r *http.Request) {
|
||||
var req SyncPairRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid request body")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" || req.SourcePath == "" || req.DestPath == "" {
|
||||
writeError(w, http.StatusBadRequest, "name, source_path and dest_path are required")
|
||||
return
|
||||
}
|
||||
if req.Direction == "" {
|
||||
req.Direction = "push"
|
||||
}
|
||||
if !directionRegex.MatchString(req.Direction) {
|
||||
writeError(w, http.StatusBadRequest, "direction must be push, pull, or mirror")
|
||||
return
|
||||
}
|
||||
if req.RsyncFlags == "" {
|
||||
req.RsyncFlags = "-aP"
|
||||
}
|
||||
if !req.Enabled {
|
||||
req.Enabled = true
|
||||
}
|
||||
|
||||
sp := &models.SyncPair{
|
||||
Name: req.Name,
|
||||
SourceMachineID: req.SourceMachineID,
|
||||
SourcePath: req.SourcePath,
|
||||
DestMachineID: req.DestMachineID,
|
||||
DestPath: req.DestPath,
|
||||
Direction: req.Direction,
|
||||
RsyncFlags: req.RsyncFlags,
|
||||
ExcludePatterns: req.ExcludePatterns,
|
||||
Enabled: req.Enabled,
|
||||
}
|
||||
|
||||
repo := models.NewSyncPairRepository(h.db)
|
||||
id, err := repo.Create(sp)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to create sync pair")
|
||||
return
|
||||
}
|
||||
sp.ID = id
|
||||
w.Header().Set("Location", "/api/sync-pairs/"+strconv.FormatInt(id, 10))
|
||||
writeJSON(w, syncPairToResp(*sp), http.StatusCreated)
|
||||
}
|
||||
|
||||
func (h *SyncPairHandler) Update(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
var req SyncPairRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid request body")
|
||||
return
|
||||
}
|
||||
|
||||
if req.Name == "" || req.SourcePath == "" || req.DestPath == "" {
|
||||
writeError(w, http.StatusBadRequest, "name, source_path and dest_path are required")
|
||||
return
|
||||
}
|
||||
if !directionRegex.MatchString(req.Direction) {
|
||||
writeError(w, http.StatusBadRequest, "direction must be push, pull, or mirror")
|
||||
return
|
||||
}
|
||||
|
||||
repo := models.NewSyncPairRepository(h.db)
|
||||
existing, err := repo.GetByID(id)
|
||||
if err == sql.ErrNoRows {
|
||||
writeError(w, http.StatusNotFound, "sync pair not found")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to fetch sync pair")
|
||||
return
|
||||
}
|
||||
|
||||
existing.Name = req.Name
|
||||
existing.SourceMachineID = req.SourceMachineID
|
||||
existing.SourcePath = req.SourcePath
|
||||
existing.DestMachineID = req.DestMachineID
|
||||
existing.DestPath = req.DestPath
|
||||
existing.Direction = req.Direction
|
||||
existing.RsyncFlags = req.RsyncFlags
|
||||
existing.ExcludePatterns = req.ExcludePatterns
|
||||
existing.Enabled = req.Enabled
|
||||
|
||||
if err := repo.Update(existing); err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to update sync pair")
|
||||
return
|
||||
}
|
||||
writeJSON(w, syncPairToResp(*existing))
|
||||
}
|
||||
|
||||
func (h *SyncPairHandler) Delete(w http.ResponseWriter, r *http.Request) {
|
||||
id, err := strconv.ParseInt(chi.URLParam(r, "id"), 10, 64)
|
||||
if err != nil {
|
||||
writeError(w, http.StatusBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
repo := models.NewSyncPairRepository(h.db)
|
||||
if err := repo.Delete(id); err != nil {
|
||||
writeError(w, http.StatusInternalServerError, "failed to delete sync pair")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
}
|
||||
|
||||
func syncPairToResp(p models.SyncPair) SyncPairResponse {
|
||||
return SyncPairResponse{
|
||||
ID: p.ID,
|
||||
Name: p.Name,
|
||||
SourceMachineID: p.SourceMachineID,
|
||||
SourcePath: p.SourcePath,
|
||||
DestMachineID: p.DestMachineID,
|
||||
DestPath: p.DestPath,
|
||||
Direction: p.Direction,
|
||||
RsyncFlags: p.RsyncFlags,
|
||||
ExcludePatterns: p.ExcludePatterns,
|
||||
Enabled: p.Enabled,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,64 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/syncserver/internal/syncengine"
|
||||
)
|
||||
|
||||
type SSEHandler struct {
|
||||
engine *syncengine.Engine
|
||||
}
|
||||
|
||||
func NewSSEHandler(engine *syncengine.Engine) *SSEHandler {
|
||||
return &SSEHandler{engine: engine}
|
||||
}
|
||||
|
||||
func (h *SSEHandler) Stream(w http.ResponseWriter, r *http.Request) {
|
||||
jobIDStr := r.URL.Query().Get("job_id")
|
||||
var filterJobID int64
|
||||
if jobIDStr != "" {
|
||||
filterJobID, _ = strconv.ParseInt(jobIDStr, 10, 64)
|
||||
}
|
||||
|
||||
flusher, ok := w.(http.Flusher)
|
||||
if !ok {
|
||||
http.Error(w, "SSE not supported", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Connection", "keep-alive")
|
||||
w.Header().Set("X-Accel-Buffering", "no")
|
||||
|
||||
flusher.Flush()
|
||||
|
||||
if h.engine == nil {
|
||||
return
|
||||
}
|
||||
|
||||
events := h.engine.Events()
|
||||
for {
|
||||
select {
|
||||
case evt := <-events:
|
||||
if filterJobID != 0 && evt.JobID != filterJobID {
|
||||
continue
|
||||
}
|
||||
data, _ := json.Marshal(evt)
|
||||
fmt.Fprintf(w, "event: %s\ndata: %s\n\n", evt.Type, data)
|
||||
flusher.Flush()
|
||||
|
||||
case <-r.Context().Done():
|
||||
return
|
||||
|
||||
case <-time.After(30 * time.Second):
|
||||
fmt.Fprintf(w, ": keepalive\n\n")
|
||||
flusher.Flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"net/http"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"github.com/go-chi/chi/v5/middleware"
|
||||
"github.com/syncserver/internal/auth"
|
||||
"github.com/syncserver/internal/config"
|
||||
"github.com/syncserver/internal/sshmanager"
|
||||
"github.com/syncserver/internal/syncengine"
|
||||
"github.com/syncserver/internal/webui"
|
||||
)
|
||||
|
||||
type Server struct {
|
||||
router *chi.Mux
|
||||
cfg *config.Config
|
||||
engine *syncengine.Engine
|
||||
}
|
||||
|
||||
func NewServer(cfg *config.Config, db *sql.DB, engine *syncengine.Engine) *Server {
|
||||
auth.InitJWTManager(cfg.Auth.JWTSecret, cfg.Auth.JWTExpiryH)
|
||||
|
||||
r := chi.NewRouter()
|
||||
r.Use(middleware.RequestID)
|
||||
r.Use(middleware.RealIP)
|
||||
r.Use(middleware.Logger)
|
||||
r.Use(middleware.Recoverer)
|
||||
|
||||
s := &Server{router: r, cfg: cfg, engine: engine}
|
||||
|
||||
authHandler := NewAuthHandler(db)
|
||||
machineHandler := NewMachineHandler(db)
|
||||
syncPairHandler := NewSyncPairHandler(db)
|
||||
jobHandler := NewJobHandler(db, engine)
|
||||
sseHandler := NewSSEHandler(engine)
|
||||
|
||||
r.Route("/api", func(r chi.Router) {
|
||||
r.Route("/auth", func(r chi.Router) {
|
||||
r.Post("/login", authHandler.Login)
|
||||
r.Post("/logout", authHandler.Logout)
|
||||
r.With(auth.RequireAuth).Get("/me", authHandler.Me)
|
||||
})
|
||||
|
||||
r.With(auth.RequireAuth).Route("/machines", func(r chi.Router) {
|
||||
r.Get("/", machineHandler.List)
|
||||
r.Post("/", machineHandler.Create)
|
||||
r.Get("/{id}", machineHandler.Get)
|
||||
r.Put("/{id}", machineHandler.Update)
|
||||
r.Delete("/{id}", machineHandler.Delete)
|
||||
})
|
||||
|
||||
r.With(auth.RequireAuth).Route("/sync-pairs", func(r chi.Router) {
|
||||
r.Get("/", syncPairHandler.List)
|
||||
r.Post("/", syncPairHandler.Create)
|
||||
r.Get("/{id}", syncPairHandler.Get)
|
||||
r.Put("/{id}", syncPairHandler.Update)
|
||||
r.Delete("/{id}", syncPairHandler.Delete)
|
||||
r.Post("/{id}/run", jobHandler.TriggerRun)
|
||||
})
|
||||
|
||||
r.With(auth.RequireAuth).Route("/jobs", func(r chi.Router) {
|
||||
r.Get("/", jobHandler.List)
|
||||
r.Get("/{id}", jobHandler.Get)
|
||||
r.Post("/{id}/cancel", jobHandler.Cancel)
|
||||
r.Get("/{id}/log", jobHandler.StreamLog)
|
||||
})
|
||||
|
||||
r.With(auth.RequireAuth).Get("/jobs/stream", sseHandler.Stream)
|
||||
|
||||
r.With(auth.RequireAuth).Get("/settings/pubkey", func(w http.ResponseWriter, r *http.Request) {
|
||||
_, _, pubKey, _ := sshmanager.EnsureServerKey(cfg.SSHDir())
|
||||
w.Header().Set("Content-Type", "text/plain")
|
||||
w.Write([]byte(pubKey))
|
||||
})
|
||||
})
|
||||
|
||||
r.Get("/health", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Write([]byte("ok"))
|
||||
}))
|
||||
|
||||
r.NotFound(func(w http.ResponseWriter, r *http.Request) {
|
||||
if _, ok := webui.DistFS.Open("dist" + r.URL.Path); ok == nil {
|
||||
http.FileServer(http.FS(webui.DistFS)).ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
data, err := webui.DistFS.ReadFile("dist/index.html")
|
||||
if err != nil {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
w.Write(data)
|
||||
})
|
||||
|
||||
return s
|
||||
}
|
||||
|
||||
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
s.router.ServeHTTP(w, r)
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
const CookieName = "ss_token"
|
||||
|
||||
func SetAuthCookie(w http.ResponseWriter, token string, expiresAt time.Time) {
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: CookieName,
|
||||
Value: token,
|
||||
Path: "/",
|
||||
Expires: expiresAt,
|
||||
HttpOnly: true,
|
||||
Secure: false,
|
||||
SameSite: http.SameSiteStrictMode,
|
||||
})
|
||||
}
|
||||
|
||||
func ClearAuthCookie(w http.ResponseWriter) {
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: CookieName,
|
||||
Value: "",
|
||||
Path: "/",
|
||||
Expires: time.Unix(0, 0),
|
||||
HttpOnly: true,
|
||||
Secure: false,
|
||||
SameSite: http.SameSiteStrictMode,
|
||||
})
|
||||
}
|
||||
|
||||
func GetTokenFromRequest(r *http.Request) string {
|
||||
cookie, err := r.Cookie(CookieName)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return cookie.Value
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidToken = errors.New("invalid token")
|
||||
ErrExpiredToken = errors.New("token expired")
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
UserID int64 `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
Role string `json:"role"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
type JWTManager struct {
|
||||
secret []byte
|
||||
expiryH int
|
||||
}
|
||||
|
||||
func NewJWTManager(secret string, expiryH int) *JWTManager {
|
||||
return &JWTManager{
|
||||
secret: []byte(secret),
|
||||
expiryH: expiryH,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *JWTManager) Generate(userID int64, username, role string) (string, time.Time, error) {
|
||||
expiresAt := time.Now().Add(time.Duration(m.expiryH) * time.Hour)
|
||||
claims := &Claims{
|
||||
UserID: userID,
|
||||
Username: username,
|
||||
Role: role,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(expiresAt),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
NotBefore: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
signed, err := token.SignedString(m.secret)
|
||||
return signed, expiresAt, err
|
||||
}
|
||||
|
||||
func (m *JWTManager) Validate(tokenStr string) (*Claims, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenStr, &Claims{}, func(token *jwt.Token) (interface{}, error) {
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, ErrInvalidToken
|
||||
}
|
||||
return m.secret, nil
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, jwt.ErrTokenExpired) {
|
||||
return nil, ErrExpiredToken
|
||||
}
|
||||
return nil, ErrInvalidToken
|
||||
}
|
||||
claims, ok := token.Claims.(*Claims)
|
||||
if !ok || !token.Valid {
|
||||
return nil, ErrInvalidToken
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
type ctxKey string
|
||||
|
||||
const ClaimsCtxKey ctxKey = "claims"
|
||||
|
||||
type contextKey struct{}
|
||||
|
||||
func RequireAuth(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
token := GetTokenFromRequest(r)
|
||||
if token == "" {
|
||||
http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
claims, err := GlobalJWTManager.Validate(token)
|
||||
if err != nil {
|
||||
http.Error(w, `{"error":"unauthorized"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
ctx := r.Context()
|
||||
ctx = context.WithValue(ctx, ClaimsCtxKey, claims)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
|
||||
func RequireAdmin(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
claims := GetClaims(r.Context())
|
||||
if claims == nil || claims.Role != "admin" {
|
||||
http.Error(w, `{"error":"forbidden"}`, http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func GetClaims(ctx context.Context) *Claims {
|
||||
v := ctx.Value(ClaimsCtxKey)
|
||||
if v == nil {
|
||||
return nil
|
||||
}
|
||||
return v.(*Claims)
|
||||
}
|
||||
|
||||
var GlobalJWTManager *JWTManager
|
||||
|
||||
func InitJWTManager(secret string, expiryH int) {
|
||||
GlobalJWTManager = NewJWTManager(secret, expiryH)
|
||||
}
|
||||
|
||||
func GetJWTManager() *JWTManager {
|
||||
return GlobalJWTManager
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
var bcryptCost = bcrypt.DefaultCost
|
||||
|
||||
func HashPassword(plain string) ([]byte, error) {
|
||||
return bcrypt.GenerateFromPassword([]byte(plain), bcryptCost)
|
||||
}
|
||||
|
||||
func VerifyPassword(hash []byte, plain string) bool {
|
||||
err := bcrypt.CompareHashAndPassword(hash, []byte(plain))
|
||||
return err == nil
|
||||
}
|
||||
@@ -0,0 +1,37 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"log/slog"
|
||||
)
|
||||
|
||||
func SeedAdmin(db *sql.DB, username, password string) error {
|
||||
if username == "" || password == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
var exists bool
|
||||
err := db.QueryRow("SELECT EXISTS(SELECT 1 FROM users WHERE username = ?)", username).Scan(&exists)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if exists {
|
||||
return nil
|
||||
}
|
||||
|
||||
hash, err := HashPassword(password)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, err = db.Exec(
|
||||
"INSERT INTO users (username, password_hash, role) VALUES (?, ?, ?)",
|
||||
username, string(hash), "admin",
|
||||
)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
slog.Info("admin user created", "username", username)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Version string `yaml:"-" json:"-"`
|
||||
DataDir string `yaml:"data_dir" env:"SYNCSERVER_DATA_DIR" default:"./data"`
|
||||
ConfigDir string `yaml:"config_dir" env:"SYNCSERVER_CONFIG_DIR" default:"./etc/syncserver"`
|
||||
Addr string `yaml:"addr" env:"SYNCSERVER_ADDR" default:":8080"`
|
||||
|
||||
Auth AuthConfig `yaml:"auth"`
|
||||
|
||||
Scheduler SchedulerConfig `yaml:"scheduler"`
|
||||
}
|
||||
|
||||
type AuthConfig struct {
|
||||
JWTSecret string `yaml:"jwt_secret" env:"SYNCSERVER_JWT_SECRET"`
|
||||
JWTExpiryH int `yaml:"jwt_expiry_hours" env:"SYNCSERVER_JWT_EXPIRY_HOURS" default:"24"`
|
||||
AdminUser string `yaml:"-" env:"SYNCSERVER_ADMIN_USER"`
|
||||
AdminPass string `yaml:"-" env:"SYNCSERVER_ADMIN_PASSWORD"`
|
||||
}
|
||||
|
||||
type SchedulerConfig struct {
|
||||
Timezone string `yaml:"timezone" env:"SYNCSERVER_SCHEDULER_TZ" default:"UTC"`
|
||||
}
|
||||
|
||||
var globalCfg *Config
|
||||
|
||||
func Load(configPath, dataDir, addr string) (*Config, error) {
|
||||
cfg := &Config{
|
||||
DataDir: dataDir,
|
||||
ConfigDir: "./etc/syncserver",
|
||||
Addr: addr,
|
||||
Auth: AuthConfig{
|
||||
JWTExpiryH: 24,
|
||||
},
|
||||
Scheduler: SchedulerConfig{
|
||||
Timezone: "UTC",
|
||||
},
|
||||
}
|
||||
|
||||
if configPath != "" {
|
||||
data, err := os.ReadFile(configPath)
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return nil, fmt.Errorf("reading config: %w", err)
|
||||
}
|
||||
if err == nil {
|
||||
if err := yaml.Unmarshal(data, cfg); err != nil {
|
||||
return nil, fmt.Errorf("parsing config: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
prefix := "SYNCSERVER_"
|
||||
envs := []struct {
|
||||
key *string
|
||||
envName string
|
||||
required bool
|
||||
}{
|
||||
{&cfg.Auth.JWTSecret, prefix + "JWT_SECRET", false},
|
||||
{&cfg.Auth.AdminUser, prefix + "ADMIN_USER", false},
|
||||
{&cfg.Auth.AdminPass, prefix + "ADMIN_PASSWORD", false},
|
||||
{&cfg.DataDir, prefix + "DATA_DIR", false},
|
||||
{&cfg.Addr, prefix + "ADDR", false},
|
||||
{&cfg.Scheduler.Timezone, prefix + "SCHEDULER_TZ", false},
|
||||
}
|
||||
|
||||
for _, e := range envs {
|
||||
if v := os.Getenv(e.envName); v != "" {
|
||||
*e.key = v
|
||||
}
|
||||
}
|
||||
|
||||
if cfg.Auth.JWTSecret == "" {
|
||||
b := make([]byte, 32)
|
||||
f, err := os.Open("/dev/urandom")
|
||||
if err == nil {
|
||||
defer f.Close()
|
||||
n, _ := f.Read(b)
|
||||
if n == 32 {
|
||||
cfg.Auth.JWTSecret = fmt.Sprintf("%x", b)
|
||||
}
|
||||
}
|
||||
if cfg.Auth.JWTSecret == "" {
|
||||
cfg.Auth.JWTSecret = "insecure-dev-secret-change-in-production"
|
||||
}
|
||||
}
|
||||
|
||||
if dataDir := os.Getenv("SYNCSERVER_DATA_DIR"); dataDir != "" {
|
||||
cfg.DataDir = dataDir
|
||||
}
|
||||
if addr := os.Getenv("SYNCSERVER_ADDR"); addr != "" {
|
||||
cfg.Addr = addr
|
||||
}
|
||||
|
||||
absDataDir, err := filepath.Abs(cfg.DataDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
cfg.DataDir = absDataDir
|
||||
|
||||
globalCfg = cfg
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func Get() *Config {
|
||||
return globalCfg
|
||||
}
|
||||
|
||||
func (c *Config) DBPath() string {
|
||||
return filepath.Join(c.DataDir, "app.db")
|
||||
}
|
||||
|
||||
func (c *Config) SSHDir() string {
|
||||
return filepath.Join(c.DataDir, "ssh")
|
||||
}
|
||||
|
||||
func (c *Config) LogsDir() string {
|
||||
return filepath.Join(c.DataDir, "logs")
|
||||
}
|
||||
|
||||
func (c *Config) EnsureDirs() error {
|
||||
dirs := []string{c.DataDir, c.SSHDir(), c.LogsDir()}
|
||||
for _, d := range dirs {
|
||||
if err := os.MkdirAll(d, 0700); err != nil {
|
||||
return fmt.Errorf("creating dir %s: %w", d, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Config) LogPath() string {
|
||||
return filepath.Join(c.LogsDir(), "app.log")
|
||||
}
|
||||
|
||||
func NormalizeAddr(addr string) string {
|
||||
addr = strings.TrimSpace(addr)
|
||||
if !strings.Contains(addr, ":") {
|
||||
addr = ":" + addr
|
||||
}
|
||||
return addr
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
type DB struct {
|
||||
*sql.DB
|
||||
}
|
||||
|
||||
func (d *DB) SQLDB() *sql.DB {
|
||||
return d.DB
|
||||
}
|
||||
|
||||
func Open(dbPath string) (*DB, error) {
|
||||
if err := os.MkdirAll(filepath.Dir(dbPath), 0700); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
db, err := sql.Open("sqlite", dbPath+"?_journal_mode=WAL&_foreign_keys=ON&_busy_timeout=5000")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("opening db: %w", err)
|
||||
}
|
||||
|
||||
if _, err := db.Exec("PRAGMA foreign_keys = ON"); err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("enabling foreign_keys: %w", err)
|
||||
}
|
||||
|
||||
return &DB{db}, nil
|
||||
}
|
||||
|
||||
func (db *DB) Close() error {
|
||||
return db.DB.Close()
|
||||
}
|
||||
|
||||
func (db *DB) RunMigrations() error {
|
||||
return db.runMigrationsInternal(migrationsFS, "migrations")
|
||||
}
|
||||
|
||||
func (db *DB) runMigrationsInternal(mfs embedFS, migrationsRoot string) error {
|
||||
entries, err := mfs.ReadDir(migrationsRoot)
|
||||
if err != nil {
|
||||
return fmt.Errorf("reading migrations dir: %w", err)
|
||||
}
|
||||
|
||||
var names []string
|
||||
for _, e := range entries {
|
||||
if !e.IsDir() && strings.HasSuffix(e.Name(), ".sql") {
|
||||
names = append(names, e.Name())
|
||||
}
|
||||
}
|
||||
sort.Strings(names)
|
||||
|
||||
if _, err := db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version TEXT PRIMARY KEY,
|
||||
applied_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
`); err != nil {
|
||||
return fmt.Errorf("creating schema_migrations table: %w", err)
|
||||
}
|
||||
|
||||
for _, name := range names {
|
||||
var applied bool
|
||||
row := db.QueryRow("SELECT 1 FROM schema_migrations WHERE version = ?", name)
|
||||
if err := row.Scan(&applied); err == nil {
|
||||
applied = true
|
||||
}
|
||||
|
||||
if applied {
|
||||
continue
|
||||
}
|
||||
|
||||
data, err := mfs.ReadFile(filepath.Join(migrationsRoot, name))
|
||||
if err != nil {
|
||||
return fmt.Errorf("reading migration %s: %w", name, err)
|
||||
}
|
||||
|
||||
if _, err := db.Exec(string(data)); err != nil {
|
||||
return fmt.Errorf("applying migration %s: %w", name, err)
|
||||
}
|
||||
|
||||
if _, err := db.Exec("INSERT INTO schema_migrations (version) VALUES (?)", name); err != nil {
|
||||
return fmt.Errorf("recording migration %s: %w", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type embedFS interface {
|
||||
ReadDir(name string) ([]os.DirEntry, error)
|
||||
ReadFile(name string) ([]byte, error)
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
package db
|
||||
|
||||
import "embed"
|
||||
|
||||
// migrationsFS is the embedded filesystem containing SQL migration files.
|
||||
//go:embed migrations
|
||||
var migrationsFS embed.FS
|
||||
@@ -0,0 +1,86 @@
|
||||
-- 0001_init.sql
|
||||
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username TEXT UNIQUE NOT NULL,
|
||||
password_hash TEXT NOT NULL,
|
||||
role TEXT NOT NULL DEFAULT 'user',
|
||||
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS ssh_keys (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
label TEXT NOT NULL,
|
||||
private_key_path TEXT NOT NULL,
|
||||
public_key TEXT NOT NULL,
|
||||
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS machines (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL,
|
||||
host TEXT NOT NULL,
|
||||
port INTEGER NOT NULL DEFAULT 22,
|
||||
ssh_user TEXT NOT NULL DEFAULT 'root',
|
||||
ssh_key_id INTEGER REFERENCES ssh_keys(id),
|
||||
mac_address TEXT,
|
||||
wol_enabled INTEGER NOT NULL DEFAULT 0,
|
||||
broadcast_addr TEXT,
|
||||
wake_timeout_seconds INTEGER NOT NULL DEFAULT 120,
|
||||
wake_check_interval_seconds INTEGER NOT NULL DEFAULT 5,
|
||||
fingerprint_confirmed INTEGER NOT NULL DEFAULT 0,
|
||||
status TEXT NOT NULL DEFAULT 'unknown',
|
||||
last_seen_at DATETIME,
|
||||
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sync_pairs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL,
|
||||
source_machine_id INTEGER REFERENCES machines(id),
|
||||
source_path TEXT NOT NULL,
|
||||
dest_machine_id INTEGER REFERENCES machines(id),
|
||||
dest_path TEXT NOT NULL,
|
||||
direction TEXT NOT NULL DEFAULT 'push',
|
||||
rsync_flags TEXT NOT NULL DEFAULT '-aP',
|
||||
exclude_patterns TEXT NOT NULL DEFAULT '',
|
||||
enabled INTEGER NOT NULL DEFAULT 1,
|
||||
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS schedules (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
sync_pair_id INTEGER NOT NULL REFERENCES sync_pairs(id) ON DELETE CASCADE,
|
||||
cron_expr TEXT NOT NULL,
|
||||
next_run_at DATETIME,
|
||||
enabled INTEGER NOT NULL DEFAULT 1,
|
||||
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS jobs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
sync_pair_id INTEGER NOT NULL REFERENCES sync_pairs(id),
|
||||
trigger_type TEXT NOT NULL DEFAULT 'manual',
|
||||
status TEXT NOT NULL DEFAULT 'queued',
|
||||
started_at DATETIME,
|
||||
finished_at DATETIME,
|
||||
log_file TEXT,
|
||||
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS job_logs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
job_id INTEGER NOT NULL REFERENCES jobs(id) ON DELETE CASCADE,
|
||||
stream TEXT NOT NULL,
|
||||
content TEXT NOT NULL,
|
||||
timestamp DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
id TEXT PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
token_hash TEXT NOT NULL,
|
||||
expires_at DATETIME NOT NULL,
|
||||
revoked INTEGER NOT NULL DEFAULT 0,
|
||||
created_at DATETIME DEFAULT CURRENT_TIMESTAMP
|
||||
);
|
||||
@@ -0,0 +1,145 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Job struct {
|
||||
ID int64 `db:"id" json:"id"`
|
||||
SyncPairID int64 `db:"sync_pair_id" json:"sync_pair_id"`
|
||||
TriggerType string `db:"trigger_type" json:"trigger_type"`
|
||||
Status string `db:"status" json:"status"`
|
||||
StartedAt *time.Time `db:"started_at" json:"started_at"`
|
||||
FinishedAt *time.Time `db:"finished_at" json:"finished_at"`
|
||||
LogFile *string `db:"log_file" json:"log_file"`
|
||||
CreatedAt time.Time `db:"created_at" json:"created_at"`
|
||||
}
|
||||
|
||||
type JobRepository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewJobRepository(db *sql.DB) *JobRepository {
|
||||
return &JobRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *JobRepository) Create(syncPairID int64, triggerType, status string) (int64, error) {
|
||||
res, err := r.db.Exec(`
|
||||
INSERT INTO jobs (sync_pair_id, trigger_type, status) VALUES (?, ?, ?)`,
|
||||
syncPairID, triggerType, status,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
func (r *JobRepository) GetByID(id int64) (*Job, error) {
|
||||
var j Job
|
||||
var started, finished sql.NullTime
|
||||
var logFile sql.NullString
|
||||
err := r.db.QueryRow(`
|
||||
SELECT id, sync_pair_id, trigger_type, status, started_at, finished_at,
|
||||
log_file, created_at FROM jobs WHERE id = ?`, id).Scan(
|
||||
&j.ID, &j.SyncPairID, &j.TriggerType, &j.Status, &started, &finished,
|
||||
&logFile, &j.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if started.Valid {
|
||||
j.StartedAt = &started.Time
|
||||
}
|
||||
if finished.Valid {
|
||||
j.FinishedAt = &finished.Time
|
||||
}
|
||||
if logFile.Valid {
|
||||
j.LogFile = &logFile.String
|
||||
}
|
||||
return &j, nil
|
||||
}
|
||||
|
||||
func (r *JobRepository) GetAll(limit, offset int) ([]Job, error) {
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, sync_pair_id, trigger_type, status, started_at, finished_at,
|
||||
log_file, created_at FROM jobs ORDER BY created_at DESC LIMIT ? OFFSET ?`,
|
||||
limit, offset)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var jobs []Job
|
||||
for rows.Next() {
|
||||
var j Job
|
||||
var started, finished sql.NullTime
|
||||
var logFile sql.NullString
|
||||
if err := rows.Scan(&j.ID, &j.SyncPairID, &j.TriggerType, &j.Status,
|
||||
&started, &finished, &logFile, &j.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if started.Valid {
|
||||
j.StartedAt = &started.Time
|
||||
}
|
||||
if finished.Valid {
|
||||
j.FinishedAt = &finished.Time
|
||||
}
|
||||
if logFile.Valid {
|
||||
j.LogFile = &logFile.String
|
||||
}
|
||||
jobs = append(jobs, j)
|
||||
}
|
||||
return jobs, rows.Err()
|
||||
}
|
||||
|
||||
func (r *JobRepository) UpdateStatus(id int64, status string) error {
|
||||
var query string
|
||||
var args []interface{}
|
||||
switch status {
|
||||
case "running", "waking_up":
|
||||
query = "UPDATE jobs SET status = ?, started_at = COALESCE(started_at, CURRENT_TIMESTAMP) WHERE id = ?"
|
||||
args = []interface{}{status, id}
|
||||
case "success", "failed", "cancelled":
|
||||
query = "UPDATE jobs SET status = ?, finished_at = CURRENT_TIMESTAMP WHERE id = ?"
|
||||
args = []interface{}{status, id}
|
||||
default:
|
||||
query = "UPDATE jobs SET status = ? WHERE id = ?"
|
||||
args = []interface{}{status, id}
|
||||
}
|
||||
_, err := r.db.Exec(query, args...)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *JobRepository) SetLogFile(id int64, path string) error {
|
||||
_, err := r.db.Exec("UPDATE jobs SET log_file = ? WHERE id = ?", path, id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *JobRepository) GetRunningBySyncPair(syncPairID int64) (*Job, error) {
|
||||
var j Job
|
||||
var started sql.NullTime
|
||||
var logFile sql.NullString
|
||||
err := r.db.QueryRow(`
|
||||
SELECT id, sync_pair_id, trigger_type, status, started_at, finished_at,
|
||||
log_file, created_at FROM jobs
|
||||
WHERE sync_pair_id = ? AND status IN ('queued','waking_up','running')
|
||||
ORDER BY created_at DESC LIMIT 1`, syncPairID).Scan(
|
||||
&j.ID, &j.SyncPairID, &j.TriggerType, &j.Status, &started,
|
||||
&j.FinishedAt, &logFile, &j.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if started.Valid {
|
||||
j.StartedAt = &started.Time
|
||||
}
|
||||
if logFile.Valid {
|
||||
j.LogFile = &logFile.String
|
||||
}
|
||||
return &j, nil
|
||||
}
|
||||
|
||||
func (r *JobRepository) Count() (int64, error) {
|
||||
var n int64
|
||||
err := r.db.QueryRow("SELECT COUNT(*) FROM jobs").Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
@@ -0,0 +1,161 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Machine struct {
|
||||
ID int64 `db:"id" json:"id"`
|
||||
Name string `db:"name" json:"name"`
|
||||
Host string `db:"host" json:"host"`
|
||||
Port int `db:"port" json:"port"`
|
||||
SSHUser string `db:"ssh_user" json:"ssh_user"`
|
||||
SSHKeyID *int64 `db:"ssh_key_id" json:"ssh_key_id"`
|
||||
MACAddress *string `db:"mac_address" json:"mac_address"`
|
||||
WoLEnabled bool `db:"wol_enabled" json:"wol_enabled"`
|
||||
BroadcastAddr *string `db:"broadcast_addr" json:"broadcast_addr"`
|
||||
WakeTimeoutSeconds int `db:"wake_timeout_seconds" json:"wake_timeout_seconds"`
|
||||
WakeCheckIntervalSeconds int `db:"wake_check_interval_seconds" json:"wake_check_interval_seconds"`
|
||||
FingerprintConfirmed bool `db:"fingerprint_confirmed" json:"fingerprint_confirmed"`
|
||||
Status string `db:"status" json:"status"`
|
||||
LastSeenAt *time.Time `db:"last_seen_at" json:"last_seen_at"`
|
||||
CreatedAt time.Time `db:"created_at" json:"created_at"`
|
||||
}
|
||||
|
||||
type MachineRepository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewMachineRepository(db *sql.DB) *MachineRepository {
|
||||
return &MachineRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *MachineRepository) Create(m *Machine) (int64, error) {
|
||||
res, err := r.db.Exec(`
|
||||
INSERT INTO machines (name, host, port, ssh_user, ssh_key_id, mac_address,
|
||||
wol_enabled, broadcast_addr, wake_timeout_seconds, wake_check_interval_seconds,
|
||||
fingerprint_confirmed, status)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
m.Name, m.Host, m.Port, m.SSHUser, m.SSHKeyID, m.MACAddress,
|
||||
boolToInt(m.WoLEnabled), m.BroadcastAddr, m.WakeTimeoutSeconds,
|
||||
m.WakeCheckIntervalSeconds, boolToInt(m.FingerprintConfirmed), m.Status,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
func (r *MachineRepository) GetAll() ([]Machine, error) {
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, name, host, port, ssh_user, ssh_key_id, mac_address,
|
||||
wol_enabled, broadcast_addr, wake_timeout_seconds, wake_check_interval_seconds,
|
||||
fingerprint_confirmed, status, last_seen_at, created_at
|
||||
FROM machines ORDER BY name`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var ms []Machine
|
||||
for rows.Next() {
|
||||
var m Machine
|
||||
var mac, bcast sql.NullString
|
||||
var keyID sql.NullInt64
|
||||
var lastSeen sql.NullTime
|
||||
err := rows.Scan(&m.ID, &m.Name, &m.Host, &m.Port, &m.SSHUser, &keyID,
|
||||
&mac, &m.WoLEnabled, &bcast, &m.WakeTimeoutSeconds,
|
||||
&m.WakeCheckIntervalSeconds, &m.FingerprintConfirmed,
|
||||
&m.Status, &lastSeen, &m.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if keyID.Valid {
|
||||
v := keyID.Int64
|
||||
m.SSHKeyID = &v
|
||||
}
|
||||
if mac.Valid {
|
||||
m.MACAddress = &mac.String
|
||||
}
|
||||
if bcast.Valid {
|
||||
m.BroadcastAddr = &bcast.String
|
||||
}
|
||||
if lastSeen.Valid {
|
||||
m.LastSeenAt = &lastSeen.Time
|
||||
}
|
||||
ms = append(ms, m)
|
||||
}
|
||||
return ms, rows.Err()
|
||||
}
|
||||
|
||||
func (r *MachineRepository) GetByID(id int64) (*Machine, error) {
|
||||
var m Machine
|
||||
var mac, bcast sql.NullString
|
||||
var keyID sql.NullInt64
|
||||
var lastSeen sql.NullTime
|
||||
err := r.db.QueryRow(`
|
||||
SELECT id, name, host, port, ssh_user, ssh_key_id, mac_address,
|
||||
wol_enabled, broadcast_addr, wake_timeout_seconds, wake_check_interval_seconds,
|
||||
fingerprint_confirmed, status, last_seen_at, created_at
|
||||
FROM machines WHERE id = ?`, id).Scan(
|
||||
&m.ID, &m.Name, &m.Host, &m.Port, &m.SSHUser, &keyID,
|
||||
&mac, &m.WoLEnabled, &bcast, &m.WakeTimeoutSeconds,
|
||||
&m.WakeCheckIntervalSeconds, &m.FingerprintConfirmed,
|
||||
&m.Status, &lastSeen, &m.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if keyID.Valid {
|
||||
v := keyID.Int64
|
||||
m.SSHKeyID = &v
|
||||
}
|
||||
if mac.Valid {
|
||||
m.MACAddress = &mac.String
|
||||
}
|
||||
if bcast.Valid {
|
||||
m.BroadcastAddr = &bcast.String
|
||||
}
|
||||
if lastSeen.Valid {
|
||||
m.LastSeenAt = &lastSeen.Time
|
||||
}
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
func (r *MachineRepository) Update(m *Machine) error {
|
||||
_, err := r.db.Exec(`
|
||||
UPDATE machines SET name=?, host=?, port=?, ssh_user=?, ssh_key_id=?,
|
||||
mac_address=?, wol_enabled=?, broadcast_addr=?, wake_timeout_seconds=?,
|
||||
wake_check_interval_seconds=?, fingerprint_confirmed=?, status=?, last_seen_at=?
|
||||
WHERE id=?`,
|
||||
m.Name, m.Host, m.Port, m.SSHUser, m.SSHKeyID, m.MACAddress,
|
||||
boolToInt(m.WoLEnabled), m.BroadcastAddr, m.WakeTimeoutSeconds,
|
||||
m.WakeCheckIntervalSeconds, boolToInt(m.FingerprintConfirmed),
|
||||
m.Status, m.LastSeenAt, m.ID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *MachineRepository) Delete(id int64) error {
|
||||
_, err := r.db.Exec("DELETE FROM machines WHERE id = ?", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *MachineRepository) UpdateStatus(id int64, status string) error {
|
||||
_, err := r.db.Exec(
|
||||
"UPDATE machines SET status = ?, last_seen_at = CURRENT_TIMESTAMP WHERE id = ?",
|
||||
status, id,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func boolToInt(b bool) int {
|
||||
if b {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func intToBool(i int) bool {
|
||||
return i != 0
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Schedule struct {
|
||||
ID int64 `db:"id" json:"id"`
|
||||
SyncPairID int64 `db:"sync_pair_id" json:"sync_pair_id"`
|
||||
CronExpr string `db:"cron_expr" json:"cron_expr"`
|
||||
NextRunAt *time.Time `db:"next_run_at" json:"next_run_at"`
|
||||
Enabled bool `db:"enabled" json:"enabled"`
|
||||
CreatedAt time.Time `db:"created_at" json:"created_at"`
|
||||
}
|
||||
|
||||
type ScheduleRepository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewScheduleRepository(db *sql.DB) *ScheduleRepository {
|
||||
return &ScheduleRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *ScheduleRepository) Create(s *Schedule) (int64, error) {
|
||||
res, err := r.db.Exec(`
|
||||
INSERT INTO schedules (sync_pair_id, cron_expr, next_run_at, enabled)
|
||||
VALUES (?, ?, ?, ?)`,
|
||||
s.SyncPairID, s.CronExpr, s.NextRunAt, boolToInt(s.Enabled),
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
func (r *ScheduleRepository) GetAll() ([]Schedule, error) {
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, sync_pair_id, cron_expr, next_run_at, enabled, created_at
|
||||
FROM schedules ORDER BY id`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var schedules []Schedule
|
||||
for rows.Next() {
|
||||
var s Schedule
|
||||
var nextRun sql.NullTime
|
||||
if err := rows.Scan(&s.ID, &s.SyncPairID, &s.CronExpr, &nextRun,
|
||||
&s.Enabled, &s.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if nextRun.Valid {
|
||||
s.NextRunAt = &nextRun.Time
|
||||
}
|
||||
s.Enabled = intToBool(intToInt(s.Enabled))
|
||||
schedules = append(schedules, s)
|
||||
}
|
||||
return schedules, rows.Err()
|
||||
}
|
||||
|
||||
func (r *ScheduleRepository) GetByID(id int64) (*Schedule, error) {
|
||||
var s Schedule
|
||||
var nextRun sql.NullTime
|
||||
err := r.db.QueryRow(`
|
||||
SELECT id, sync_pair_id, cron_expr, next_run_at, enabled, created_at
|
||||
FROM schedules WHERE id = ?`, id).Scan(
|
||||
&s.ID, &s.SyncPairID, &s.CronExpr, &nextRun, &s.Enabled, &s.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if nextRun.Valid {
|
||||
s.NextRunAt = &nextRun.Time
|
||||
}
|
||||
s.Enabled = intToBool(intToInt(s.Enabled))
|
||||
return &s, nil
|
||||
}
|
||||
|
||||
func (r *ScheduleRepository) Update(s *Schedule) error {
|
||||
_, err := r.db.Exec(`
|
||||
UPDATE schedules SET sync_pair_id=?, cron_expr=?, next_run_at=?, enabled=?
|
||||
WHERE id=?`,
|
||||
s.SyncPairID, s.CronExpr, s.NextRunAt, boolToInt(s.Enabled), s.ID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *ScheduleRepository) Delete(id int64) error {
|
||||
_, err := r.db.Exec("DELETE FROM schedules WHERE id = ?", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *ScheduleRepository) GetEnabledDue(before time.Time) ([]Schedule, error) {
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, sync_pair_id, cron_expr, next_run_at, enabled, created_at
|
||||
FROM schedules WHERE enabled = 1 AND next_run_at IS NOT NULL AND next_run_at <= ?
|
||||
ORDER BY next_run_at`, before)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var schedules []Schedule
|
||||
for rows.Next() {
|
||||
var s Schedule
|
||||
var nextRun sql.NullTime
|
||||
if err := rows.Scan(&s.ID, &s.SyncPairID, &s.CronExpr, &nextRun,
|
||||
&s.Enabled, &s.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if nextRun.Valid {
|
||||
s.NextRunAt = &nextRun.Time
|
||||
}
|
||||
s.Enabled = true
|
||||
schedules = append(schedules, s)
|
||||
}
|
||||
return schedules, rows.Err()
|
||||
}
|
||||
|
||||
func (r *ScheduleRepository) UpdateNextRun(id int64, nextRun time.Time) error {
|
||||
_, err := r.db.Exec("UPDATE schedules SET next_run_at = ? WHERE id = ?", nextRun, id)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
type SSHKey struct {
|
||||
ID int64 `db:"id" json:"id"`
|
||||
Label string `db:"label" json:"label"`
|
||||
PrivateKeyPath string `db:"private_key_path" json:"-"`
|
||||
PublicKey string `db:"public_key" json:"public_key"`
|
||||
CreatedAt time.Time `db:"created_at" json:"created_at"`
|
||||
}
|
||||
|
||||
type SSHKeyRepository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewSSHKeyRepository(db *sql.DB) *SSHKeyRepository {
|
||||
return &SSHKeyRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *SSHKeyRepository) Create(label, privPath, pubKey string) (int64, error) {
|
||||
res, err := r.db.Exec(
|
||||
"INSERT INTO ssh_keys (label, private_key_path, public_key) VALUES (?, ?, ?)",
|
||||
label, privPath, pubKey,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
func (r *SSHKeyRepository) GetByID(id int64) (*SSHKey, error) {
|
||||
var k SSHKey
|
||||
err := r.db.QueryRow(
|
||||
"SELECT id, label, private_key_path, public_key, created_at FROM ssh_keys WHERE id = ?",
|
||||
id,
|
||||
).Scan(&k.ID, &k.Label, &k.PrivateKeyPath, &k.PublicKey, &k.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &k, nil
|
||||
}
|
||||
|
||||
func (r *SSHKeyRepository) GetAll() ([]SSHKey, error) {
|
||||
rows, err := r.db.Query(
|
||||
"SELECT id, label, private_key_path, public_key, created_at FROM ssh_keys ORDER BY label")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var keys []SSHKey
|
||||
for rows.Next() {
|
||||
var k SSHKey
|
||||
if err := rows.Scan(&k.ID, &k.Label, &k.PrivateKeyPath, &k.PublicKey, &k.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
keys = append(keys, k)
|
||||
}
|
||||
return keys, rows.Err()
|
||||
}
|
||||
|
||||
func (r *SSHKeyRepository) Delete(id int64) error {
|
||||
_, err := r.db.Exec("DELETE FROM ssh_keys WHERE id = ?", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *SSHKeyRepository) GetServerKey() (*SSHKey, error) {
|
||||
var k SSHKey
|
||||
err := r.db.QueryRow(
|
||||
"SELECT id, label, private_key_path, public_key, created_at FROM ssh_keys WHERE label = 'server' LIMIT 1",
|
||||
).Scan(&k.ID, &k.Label, &k.PrivateKeyPath, &k.PublicKey, &k.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &k, nil
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type SyncPair struct {
|
||||
ID int64 `db:"id" json:"id"`
|
||||
Name string `db:"name" json:"name"`
|
||||
SourceMachineID *int64 `db:"source_machine_id" json:"source_machine_id"`
|
||||
SourcePath string `db:"source_path" json:"source_path"`
|
||||
DestMachineID *int64 `db:"dest_machine_id" json:"dest_machine_id"`
|
||||
DestPath string `db:"dest_path" json:"dest_path"`
|
||||
Direction string `db:"direction" json:"direction"`
|
||||
RsyncFlags string `db:"rsync_flags" json:"rsync_flags"`
|
||||
ExcludePatterns string `db:"exclude_patterns" json:"exclude_patterns"`
|
||||
Enabled bool `db:"enabled" json:"enabled"`
|
||||
CreatedAt time.Time `db:"created_at" json:"created_at"`
|
||||
}
|
||||
|
||||
type SyncPairRepository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewSyncPairRepository(db *sql.DB) *SyncPairRepository {
|
||||
return &SyncPairRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *SyncPairRepository) Create(sp *SyncPair) (int64, error) {
|
||||
res, err := r.db.Exec(`
|
||||
INSERT INTO sync_pairs (name, source_machine_id, source_path, dest_machine_id,
|
||||
dest_path, direction, rsync_flags, exclude_patterns, enabled)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
sp.Name, sp.SourceMachineID, sp.SourcePath, sp.DestMachineID,
|
||||
sp.DestPath, sp.Direction, sp.RsyncFlags, sp.ExcludePatterns, boolToInt(sp.Enabled),
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
func (r *SyncPairRepository) GetAll() ([]SyncPair, error) {
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, name, source_machine_id, source_path, dest_machine_id, dest_path,
|
||||
direction, rsync_flags, exclude_patterns, enabled, created_at
|
||||
FROM sync_pairs ORDER BY name`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var pairs []SyncPair
|
||||
for rows.Next() {
|
||||
var p SyncPair
|
||||
var srcID, dstID sql.NullInt64
|
||||
if err := rows.Scan(&p.ID, &p.Name, &srcID, &p.SourcePath, &dstID,
|
||||
&p.DestPath, &p.Direction, &p.RsyncFlags, &p.ExcludePatterns,
|
||||
&p.Enabled, &p.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if srcID.Valid {
|
||||
v := srcID.Int64
|
||||
p.SourceMachineID = &v
|
||||
}
|
||||
if dstID.Valid {
|
||||
v := dstID.Int64
|
||||
p.DestMachineID = &v
|
||||
}
|
||||
p.Enabled = intToBool(intToInt(p.Enabled))
|
||||
pairs = append(pairs, p)
|
||||
}
|
||||
return pairs, rows.Err()
|
||||
}
|
||||
|
||||
func intToInt(v interface{}) int {
|
||||
switch x := v.(type) {
|
||||
case int:
|
||||
return x
|
||||
case int64:
|
||||
return int(x)
|
||||
case bool:
|
||||
if x {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
func (r *SyncPairRepository) GetByID(id int64) (*SyncPair, error) {
|
||||
var p SyncPair
|
||||
var srcID, dstID sql.NullInt64
|
||||
err := r.db.QueryRow(`
|
||||
SELECT id, name, source_machine_id, source_path, dest_machine_id, dest_path,
|
||||
direction, rsync_flags, exclude_patterns, enabled, created_at
|
||||
FROM sync_pairs WHERE id = ?`, id).Scan(
|
||||
&p.ID, &p.Name, &srcID, &p.SourcePath, &dstID,
|
||||
&p.DestPath, &p.Direction, &p.RsyncFlags, &p.ExcludePatterns,
|
||||
&p.Enabled, &p.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if srcID.Valid {
|
||||
v := srcID.Int64
|
||||
p.SourceMachineID = &v
|
||||
}
|
||||
if dstID.Valid {
|
||||
v := dstID.Int64
|
||||
p.DestMachineID = &v
|
||||
}
|
||||
p.Enabled = intToBool(intToInt(p.Enabled))
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
func (r *SyncPairRepository) Update(sp *SyncPair) error {
|
||||
_, err := r.db.Exec(`
|
||||
UPDATE sync_pairs SET name=?, source_machine_id=?, source_path=?, dest_machine_id=?,
|
||||
dest_path=?, direction=?, rsync_flags=?, exclude_patterns=?, enabled=?
|
||||
WHERE id=?`,
|
||||
sp.Name, sp.SourceMachineID, sp.SourcePath, sp.DestMachineID,
|
||||
sp.DestPath, sp.Direction, sp.RsyncFlags, sp.ExcludePatterns,
|
||||
boolToInt(sp.Enabled), sp.ID,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *SyncPairRepository) Delete(id int64) error {
|
||||
_, err := r.db.Exec("DELETE FROM sync_pairs WHERE id = ?", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *SyncPairRepository) GetEnabled() ([]SyncPair, error) {
|
||||
rows, err := r.db.Query(`
|
||||
SELECT id, name, source_machine_id, source_path, dest_machine_id, dest_path,
|
||||
direction, rsync_flags, exclude_patterns, enabled, created_at
|
||||
FROM sync_pairs WHERE enabled = 1`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var pairs []SyncPair
|
||||
for rows.Next() {
|
||||
var p SyncPair
|
||||
var srcID, dstID sql.NullInt64
|
||||
if err := rows.Scan(&p.ID, &p.Name, &srcID, &p.SourcePath, &dstID,
|
||||
&p.DestPath, &p.Direction, &p.RsyncFlags, &p.ExcludePatterns,
|
||||
&p.Enabled, &p.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if srcID.Valid {
|
||||
v := srcID.Int64
|
||||
p.SourceMachineID = &v
|
||||
}
|
||||
if dstID.Valid {
|
||||
v := dstID.Int64
|
||||
p.DestMachineID = &v
|
||||
}
|
||||
p.Enabled = true
|
||||
pairs = append(pairs, p)
|
||||
}
|
||||
return pairs, rows.Err()
|
||||
}
|
||||
|
||||
func (sp *SyncPair) ExcludePatternsList() []string {
|
||||
if sp.ExcludePatterns == "" {
|
||||
return nil
|
||||
}
|
||||
var patterns []string
|
||||
for _, p := range strings.Split(sp.ExcludePatterns, "\n") {
|
||||
p = strings.TrimSpace(p)
|
||||
if p != "" {
|
||||
patterns = append(patterns, p)
|
||||
}
|
||||
}
|
||||
return patterns
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
type User struct {
|
||||
ID int64 `db:"id" json:"id"`
|
||||
Username string `db:"username" json:"username"`
|
||||
PasswordHash string `db:"password_hash" json:"-"`
|
||||
Role string `db:"role" json:"role"`
|
||||
CreatedAt time.Time `db:"created_at" json:"created_at"`
|
||||
}
|
||||
|
||||
type UserRepository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewUserRepository(db *sql.DB) *UserRepository {
|
||||
return &UserRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *UserRepository) Create(username, passwordHash, role string) (int64, error) {
|
||||
res, err := r.db.Exec(
|
||||
"INSERT INTO users (username, password_hash, role) VALUES (?, ?, ?)",
|
||||
username, passwordHash, role,
|
||||
)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return res.LastInsertId()
|
||||
}
|
||||
|
||||
func (r *UserRepository) GetByUsername(username string) (*User, error) {
|
||||
var u User
|
||||
err := r.db.QueryRow(
|
||||
"SELECT id, username, password_hash, role, created_at FROM users WHERE username = ?",
|
||||
username,
|
||||
).Scan(&u.ID, &u.Username, &u.PasswordHash, &u.Role, &u.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
func (r *UserRepository) Exists() (bool, error) {
|
||||
var n int
|
||||
err := r.db.QueryRow("SELECT COUNT(*) FROM users").Scan(&n)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return n > 0, nil
|
||||
}
|
||||
|
||||
func (r *UserRepository) GetByID(id int64) (*User, error) {
|
||||
var u User
|
||||
err := r.db.QueryRow(
|
||||
"SELECT id, username, password_hash, role, created_at FROM users WHERE id = ?",
|
||||
id,
|
||||
).Scan(&u.ID, &u.Username, &u.PasswordHash, &u.Role, &u.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &u, nil
|
||||
}
|
||||
@@ -0,0 +1,153 @@
|
||||
package scheduler
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type CronExpr struct {
|
||||
Minute []int
|
||||
Hour []int
|
||||
DayOfMonth []int
|
||||
Month []int
|
||||
DayOfWeek []int
|
||||
}
|
||||
|
||||
func ParseCron(expr string) (*CronExpr, error) {
|
||||
parts := strings.Fields(expr)
|
||||
if len(parts) != 5 {
|
||||
return nil, fmt.Errorf("expected 5 fields, got %d", len(parts))
|
||||
}
|
||||
|
||||
minute, err := parseField(parts[0], 0, 59)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("minute: %w", err)
|
||||
}
|
||||
hour, err := parseField(parts[1], 0, 23)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("hour: %w", err)
|
||||
}
|
||||
dom, err := parseField(parts[2], 1, 31)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("day of month: %w", err)
|
||||
}
|
||||
month, err := parseField(parts[3], 1, 12)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("month: %w", err)
|
||||
}
|
||||
dow, err := parseField(parts[4], 0, 6)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("day of week: %w", err)
|
||||
}
|
||||
|
||||
return &CronExpr{
|
||||
Minute: minute,
|
||||
Hour: hour,
|
||||
DayOfMonth: dom,
|
||||
Month: month,
|
||||
DayOfWeek: dow,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func parseField(field string, min, max int) ([]int, error) {
|
||||
if field == "*" {
|
||||
var vals []int
|
||||
for i := min; i <= max; i++ {
|
||||
vals = append(vals, i)
|
||||
}
|
||||
return vals, nil
|
||||
}
|
||||
|
||||
var result []int
|
||||
parts := strings.Split(field, ",")
|
||||
for _, part := range parts {
|
||||
if strings.Contains(part, "/") {
|
||||
stepParts := strings.Split(part, "/")
|
||||
if len(stepParts) != 2 {
|
||||
return nil, fmt.Errorf("invalid step: %s", part)
|
||||
}
|
||||
rangePart := stepParts[0]
|
||||
step, err := strconv.Atoi(stepParts[1])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid step value: %s", stepParts[1])
|
||||
}
|
||||
var start, end int
|
||||
if rangePart == "*" {
|
||||
start, end = min, max
|
||||
} else if strings.Contains(rangePart, "-") {
|
||||
rp := strings.Split(rangePart, "-")
|
||||
if len(rp) != 2 {
|
||||
return nil, fmt.Errorf("invalid range: %s", rangePart)
|
||||
}
|
||||
start, _ = strconv.Atoi(rp[0])
|
||||
end, _ = strconv.Atoi(rp[1])
|
||||
} else {
|
||||
v, _ := strconv.Atoi(rangePart)
|
||||
start, end = v, v
|
||||
}
|
||||
for i := start; i <= end; i += step {
|
||||
result = append(result, i)
|
||||
}
|
||||
} else if strings.Contains(part, "-") {
|
||||
rp := strings.Split(part, "-")
|
||||
if len(rp) != 2 {
|
||||
return nil, fmt.Errorf("invalid range: %s", part)
|
||||
}
|
||||
start, _ := strconv.Atoi(rp[0])
|
||||
end, _ := strconv.Atoi(rp[1])
|
||||
for i := start; i <= end; i++ {
|
||||
result = append(result, i)
|
||||
}
|
||||
} else {
|
||||
v, err := strconv.Atoi(part)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid value: %s", part)
|
||||
}
|
||||
if v < min || v > max {
|
||||
return nil, fmt.Errorf("value %d out of range [%d,%d]", v, min, max)
|
||||
}
|
||||
result = append(result, v)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (c *CronExpr) Matches(t time.Time) bool {
|
||||
if !contains(c.Minute, t.Minute()) {
|
||||
return false
|
||||
}
|
||||
if !contains(c.Hour, t.Hour()) {
|
||||
return false
|
||||
}
|
||||
if !contains(c.DayOfMonth, t.Day()) {
|
||||
return false
|
||||
}
|
||||
if !contains(c.Month, int(t.Month())) {
|
||||
return false
|
||||
}
|
||||
if !contains(c.DayOfWeek, int(t.Weekday())) {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func contains(slice []int, val int) bool {
|
||||
for _, v := range slice {
|
||||
if v == val {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func NextRun(expr *CronExpr, from time.Time) time.Time {
|
||||
for i := 1; i <= 525600; i++ {
|
||||
t := from.Add(time.Duration(i) * time.Minute)
|
||||
if expr.Matches(t) {
|
||||
return t
|
||||
}
|
||||
}
|
||||
return from.AddDate(0, 0, 1)
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package scheduler
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestParseCron(t *testing.T) {
|
||||
tests := []struct {
|
||||
expr string
|
||||
wantErr bool
|
||||
}{
|
||||
{"* * * * *", false},
|
||||
{"0 * * * *", false},
|
||||
{"*/5 * * * *", false},
|
||||
{"0,30 * * * *", false},
|
||||
{"0-30 * * * *", false},
|
||||
{"*/15 9-17 * * *", false},
|
||||
{"0 0 1 * *", false},
|
||||
{"0 0 * * 0", false},
|
||||
{"0 0 1,15 * *", false},
|
||||
{"invalid", true},
|
||||
{"* * * *", true},
|
||||
{"60 * * * *", true},
|
||||
{"* 24 * * *", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
_, err := ParseCron(tt.expr)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("ParseCron(%q) error = %v, wantErr %v", tt.expr, err, tt.wantErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCronMatches(t *testing.T) {
|
||||
expr, err := ParseCron("*/5 * * * *")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
minute int
|
||||
want bool
|
||||
}{
|
||||
{0, true},
|
||||
{5, true},
|
||||
{10, true},
|
||||
{15, true},
|
||||
{1, false},
|
||||
{2, false},
|
||||
{7, false},
|
||||
}
|
||||
|
||||
now := time.Date(2024, 1, 1, 12, 0, 0, 0, time.UTC)
|
||||
for _, tt := range tests {
|
||||
tm := time.Date(now.Year(), now.Month(), now.Day(), now.Hour(), tt.minute, 0, 0, time.UTC)
|
||||
if got := expr.Matches(tm); got != tt.want {
|
||||
t.Errorf("Matches(minute=%d) = %v, want %v", tt.minute, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCronMatchesSpecific(t *testing.T) {
|
||||
expr, err := ParseCron("30 9 15 * *")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
matches := time.Date(2024, 6, 15, 9, 30, 0, 0, time.UTC)
|
||||
if !expr.Matches(matches) {
|
||||
t.Error("should match 9:30 on 15th of month")
|
||||
}
|
||||
|
||||
notMatch := time.Date(2024, 6, 16, 9, 30, 0, 0, time.UTC)
|
||||
if expr.Matches(notMatch) {
|
||||
t.Error("should not match 9:30 on 16th of month")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNextRun(t *testing.T) {
|
||||
expr, err := ParseCron("*/5 * * * *")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
from := time.Date(2024, 1, 1, 12, 0, 0, 0, time.UTC)
|
||||
next := NextRun(expr, from)
|
||||
|
||||
if next.Minute() != 5 || next.Hour() != 12 {
|
||||
t.Errorf("NextRun = %v, want 12:05", next)
|
||||
}
|
||||
|
||||
if !next.After(from) {
|
||||
t.Error("NextRun should be after from time")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package scheduler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/syncserver/internal/config"
|
||||
"github.com/syncserver/internal/models"
|
||||
"github.com/syncserver/internal/syncengine"
|
||||
)
|
||||
|
||||
type Scheduler struct {
|
||||
db *sql.DB
|
||||
engine *syncengine.Engine
|
||||
cfg *config.Config
|
||||
stopCh chan struct{}
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func New(database interface{ SQLDB() *sql.DB }, engine *syncengine.Engine, cfg *config.Config) *Scheduler {
|
||||
return &Scheduler{
|
||||
db: database.SQLDB(),
|
||||
engine: engine,
|
||||
cfg: cfg,
|
||||
stopCh: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Scheduler) Start() {
|
||||
s.wg.Add(1)
|
||||
go s.run()
|
||||
slog.Info("scheduler started")
|
||||
}
|
||||
|
||||
func (s *Scheduler) Stop() {
|
||||
close(s.stopCh)
|
||||
s.wg.Wait()
|
||||
slog.Info("scheduler stopped")
|
||||
}
|
||||
|
||||
func (s *Scheduler) run() {
|
||||
defer s.wg.Done()
|
||||
ticker := time.NewTicker(1 * time.Minute)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-s.stopCh:
|
||||
return
|
||||
case <-ticker.C:
|
||||
s.tick()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Scheduler) tick() {
|
||||
scheduleRepo := models.NewScheduleRepository(s.db)
|
||||
now := time.Now().UTC()
|
||||
|
||||
schedules, err := scheduleRepo.GetEnabledDue(now)
|
||||
if err != nil {
|
||||
slog.Error("scheduler: failed to get due schedules", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
for _, sch := range schedules {
|
||||
pairRepo := models.NewSyncPairRepository(s.db)
|
||||
pair, err := pairRepo.GetByID(sch.SyncPairID)
|
||||
if err != nil || !pair.Enabled {
|
||||
continue
|
||||
}
|
||||
|
||||
jobID, err := s.engine.CreateJob(sch.SyncPairID, "scheduled")
|
||||
if err != nil {
|
||||
slog.Error("scheduler: failed to create job", "schedule_id", sch.ID, "error", err)
|
||||
continue
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
go func(jobID int64, pairID int64, schID int64) {
|
||||
if err := s.engine.Run(ctx, jobID, pairID); err != nil {
|
||||
slog.Warn("scheduler: job failed", "job_id", jobID, "error", err)
|
||||
}
|
||||
|
||||
expr, _ := ParseCron(sch.CronExpr)
|
||||
if expr != nil {
|
||||
next := NextRun(expr, time.Now().UTC())
|
||||
scheduleRepo.UpdateNextRun(schID, next)
|
||||
}
|
||||
}(jobID, sch.SyncPairID, sch.ID)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package sshmanager
|
||||
|
||||
import (
|
||||
"crypto/ed25519"
|
||||
"crypto/rand"
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const ServerKeyLabel = "server"
|
||||
|
||||
func EnsureServerKey(sshDir string) (privPath, pubPath string, pubKey string, err error) {
|
||||
if err := os.MkdirAll(sshDir, 0700); err != nil {
|
||||
return "", "", "", fmt.Errorf("creating ssh dir: %w", err)
|
||||
}
|
||||
|
||||
privPath = filepath.Join(sshDir, "id_ed25519")
|
||||
pubPath = filepath.Join(sshDir, "id_ed25519.pub")
|
||||
|
||||
if _, err := os.Stat(privPath); os.IsNotExist(err) {
|
||||
pub, priv, err := ed25519.GenerateKey(rand.Reader)
|
||||
if err != nil {
|
||||
return "", "", "", fmt.Errorf("generating ed25519 key: %w", err)
|
||||
}
|
||||
|
||||
privFile, err := os.OpenFile(privPath, os.O_CREATE|os.O_WRONLY, 0600)
|
||||
if err != nil {
|
||||
return "", "", "", fmt.Errorf("creating private key file: %w", err)
|
||||
}
|
||||
defer privFile.Close()
|
||||
|
||||
privBytes, err := x509.MarshalPKCS8PrivateKey(priv)
|
||||
if err != nil {
|
||||
return "", "", "", fmt.Errorf("marshaling private key: %w", err)
|
||||
}
|
||||
pem.Encode(privFile, &pem.Block{Type: "PRIVATE KEY", Bytes: privBytes})
|
||||
|
||||
pubKey = fmt.Sprintf("%s %s", strings.TrimSpace(string(pub)), "syncserver")
|
||||
if err := os.WriteFile(pubPath, []byte(pubKey), 0644); err != nil {
|
||||
return "", "", "", fmt.Errorf("writing public key: %w", err)
|
||||
}
|
||||
return privPath, pubPath, pubKey, nil
|
||||
} else if err != nil {
|
||||
return "", "", "", fmt.Errorf("checking private key: %w", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(pubPath)
|
||||
if err != nil {
|
||||
return "", "", "", fmt.Errorf("reading public key: %w", err)
|
||||
}
|
||||
return privPath, pubPath, strings.TrimSpace(string(data)), nil
|
||||
}
|
||||
|
||||
func ReadPrivateKey(path string) ([]byte, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
block, _ := pem.Decode(data)
|
||||
if block == nil {
|
||||
return nil, fmt.Errorf("no PEM block found")
|
||||
}
|
||||
return block.Bytes, nil
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package sshmanager
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type KnownHost struct {
|
||||
Host string
|
||||
Port int
|
||||
KeyType string
|
||||
Fingerprint string
|
||||
}
|
||||
|
||||
func EnsureKnownHosts(sshDir string) (string, error) {
|
||||
path := filepath.Join(sshDir, "known_hosts")
|
||||
f, err := os.OpenFile(path, os.O_CREATE|os.O_RDONLY, 0644)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
f.Close()
|
||||
return path, nil
|
||||
}
|
||||
|
||||
func AddKnownHost(sshDir, host string, port int, keyData []byte) error {
|
||||
path := filepath.Join(sshDir, "known_hosts")
|
||||
f, err := os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
addr := host
|
||||
if port != 22 {
|
||||
addr = fmt.Sprintf("[%s]:%d", host, port)
|
||||
}
|
||||
|
||||
line := fmt.Sprintf("%s %s\n", addr, strings.TrimSpace(string(keyData)))
|
||||
if _, err := f.WriteString(line); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseHostPort(entry string) (string, int) {
|
||||
if strings.HasPrefix(entry, "[") {
|
||||
var h string
|
||||
var p int
|
||||
if n, _ := fmt.Sscanf(entry, "[%[^]]]:%d", &h, &p); n == 2 {
|
||||
return h, p
|
||||
}
|
||||
}
|
||||
parts := strings.Split(entry, ":")
|
||||
if len(parts) == 2 {
|
||||
return parts[0], 22
|
||||
}
|
||||
return entry, 22
|
||||
}
|
||||
|
||||
func GetKnownHost(sshDir, host string, port int) (*KnownHost, error) {
|
||||
path := filepath.Join(sshDir, "known_hosts")
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
var targetAddr string
|
||||
if port != 22 {
|
||||
targetAddr = fmt.Sprintf("[%s]:%d", host, port)
|
||||
} else {
|
||||
targetAddr = host
|
||||
}
|
||||
|
||||
scanner := bufio.NewScanner(f)
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
parts := strings.Fields(line)
|
||||
if len(parts) < 2 {
|
||||
continue
|
||||
}
|
||||
h, p := parseHostPort(parts[0])
|
||||
if (h == host || parts[0] == targetAddr) && p == port {
|
||||
return &KnownHost{
|
||||
Host: h,
|
||||
Port: p,
|
||||
KeyType: parts[1],
|
||||
Fingerprint: parts[1] + " " + parts[2],
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func HasKnownHost(sshDir, host string, port int) (bool, error) {
|
||||
kh, err := GetKnownHost(sshDir, host, port)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return kh != nil, nil
|
||||
}
|
||||
|
||||
func RemoveKnownHost(sshDir, host string, port int) error {
|
||||
path := filepath.Join(sshDir, "known_hosts")
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
var lines []string
|
||||
targetAddr := host
|
||||
if port != 22 {
|
||||
targetAddr = fmt.Sprintf("[%s]:%d", host, port)
|
||||
}
|
||||
|
||||
scanner := bufio.NewScanner(f)
|
||||
for scanner.Scan() {
|
||||
line := scanner.Text()
|
||||
h, p := parseHostPort(line)
|
||||
if h == host && p == port {
|
||||
continue
|
||||
}
|
||||
if line == targetAddr {
|
||||
continue
|
||||
}
|
||||
lines = append(lines, line)
|
||||
}
|
||||
|
||||
tmp := path + ".tmp"
|
||||
wf, err := os.Create(tmp)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, l := range lines {
|
||||
wf.WriteString(l + "\n")
|
||||
}
|
||||
wf.Close()
|
||||
return os.Rename(tmp, path)
|
||||
}
|
||||
@@ -0,0 +1,107 @@
|
||||
package sshmanager
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/ssh"
|
||||
)
|
||||
|
||||
type ConnResult struct {
|
||||
Success bool
|
||||
Output string
|
||||
Error string
|
||||
Fingerprint string
|
||||
}
|
||||
|
||||
func TestSSHConnection(ctx context.Context, host string, port int, user, privKeyPath, knownHostsPath string, strictHostKeyChecking bool) (*ConnResult, error) {
|
||||
addr := fmt.Sprintf("%s:%d", host, port)
|
||||
|
||||
auths := []ssh.AuthMethod{}
|
||||
if privKeyPath != "" {
|
||||
key, err := os.ReadFile(privKeyPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reading private key: %w", err)
|
||||
}
|
||||
signer, err := ssh.ParsePrivateKey(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parsing private key: %w", err)
|
||||
}
|
||||
auths = append(auths, ssh.PublicKeys(signer))
|
||||
}
|
||||
|
||||
hostKeyPolicy := ssh.InsecureIgnoreHostKey()
|
||||
if strictHostKeyChecking && knownHostsPath != "" {
|
||||
hostKeyCallback, err := getHostKeyCallback(knownHostsPath, host, port)
|
||||
if err != nil {
|
||||
return &ConnResult{Success: false, Error: fmt.Sprintf("known_hosts: %v", err)}, nil
|
||||
}
|
||||
hostKeyPolicy = hostKeyCallback
|
||||
}
|
||||
|
||||
cfg := &ssh.ClientConfig{
|
||||
User: user,
|
||||
Auth: auths,
|
||||
HostKeyCallback: hostKeyPolicy,
|
||||
Timeout: 10 * time.Second,
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(ctx, 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
conn, err := ssh.Dial("tcp", addr, cfg)
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "known_hosts") || strings.Contains(err.Error(), "host key") {
|
||||
return &ConnResult{
|
||||
Success: false,
|
||||
Error: fmt.Sprintf("host key verification failed: %v", err),
|
||||
}, nil
|
||||
}
|
||||
return &ConnResult{
|
||||
Success: false,
|
||||
Error: fmt.Sprintf("connection failed: %v", err),
|
||||
}, nil
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
session, err := conn.NewSession()
|
||||
if err != nil {
|
||||
return &ConnResult{Success: false, Error: fmt.Sprintf("session: %v", err)}, nil
|
||||
}
|
||||
defer session.Close()
|
||||
|
||||
var stdout, stderr bytes.Buffer
|
||||
session.Stdout = &stdout
|
||||
session.Stderr = &stderr
|
||||
|
||||
if err := session.Run("echo ok && uname -a"); err != nil {
|
||||
return &ConnResult{
|
||||
Success: false,
|
||||
Error: fmt.Sprintf("exec: %v, stderr: %s", err, stderr.String()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
return &ConnResult{
|
||||
Success: true,
|
||||
Output: stdout.String(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func getHostKeyCallback(knownHostsPath, host string, port int) (ssh.HostKeyCallback, error) {
|
||||
return ssh.HostKeyCallback(func(hostname string, remote net.Addr, key ssh.PublicKey) error {
|
||||
kh, err := GetKnownHost(filepath.Dir(knownHostsPath), host, port)
|
||||
if err != nil {
|
||||
return fmt.Errorf("checking known_hosts: %w", err)
|
||||
}
|
||||
if kh == nil {
|
||||
return fmt.Errorf("host key not found in known_hosts: %s", hostname)
|
||||
}
|
||||
return nil
|
||||
}), nil
|
||||
}
|
||||
@@ -0,0 +1,223 @@
|
||||
package syncengine
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/syncserver/internal/config"
|
||||
"github.com/syncserver/internal/models"
|
||||
"github.com/syncserver/internal/wol"
|
||||
)
|
||||
|
||||
type Engine struct {
|
||||
db *sql.DB
|
||||
cfg *config.Config
|
||||
queue *Queue
|
||||
eventBus chan Event
|
||||
mu sync.RWMutex
|
||||
stopped bool
|
||||
}
|
||||
|
||||
type Event struct {
|
||||
Type string
|
||||
JobID int64
|
||||
Status string
|
||||
Line string
|
||||
Stream string
|
||||
}
|
||||
|
||||
func New(database interface{ SQLDB() *sql.DB }, cfg *config.Config) *Engine {
|
||||
e := &Engine{
|
||||
db: database.SQLDB(),
|
||||
cfg: cfg,
|
||||
queue: NewQueue(),
|
||||
eventBus: make(chan Event, 100),
|
||||
}
|
||||
return e
|
||||
}
|
||||
|
||||
func (e *Engine) Start() {}
|
||||
func (e *Engine) Stop() {}
|
||||
|
||||
func (e *Engine) Events() <-chan Event {
|
||||
return e.eventBus
|
||||
}
|
||||
|
||||
func (e *Engine) Run(ctx context.Context, jobID int64, pairID int64) error {
|
||||
if e.queue.IsRunning(pairID) {
|
||||
existingJobID, _ := e.queue.GetJobID(pairID)
|
||||
return fmt.Errorf("%w: job %d is already running", ErrAlreadyRunning, existingJobID)
|
||||
}
|
||||
|
||||
jobCtx, cancel := context.WithCancel(ctx)
|
||||
enqueueErr := e.queue.Enqueue(pairID, jobID, cancel)
|
||||
if enqueueErr != nil {
|
||||
return enqueueErr
|
||||
}
|
||||
defer e.queue.Dequeue(pairID)
|
||||
|
||||
pairRepo := models.NewSyncPairRepository(e.db)
|
||||
pair, err := pairRepo.GetByID(pairID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("fetching sync pair: %w", err)
|
||||
}
|
||||
|
||||
machineRepo := models.NewMachineRepository(e.db)
|
||||
var srcMachine, dstMachine *models.Machine
|
||||
|
||||
if pair.SourceMachineID != nil {
|
||||
m, _ := machineRepo.GetByID(*pair.SourceMachineID)
|
||||
srcMachine = m
|
||||
}
|
||||
if pair.DestMachineID != nil {
|
||||
m, _ := machineRepo.GetByID(*pair.DestMachineID)
|
||||
dstMachine = m
|
||||
}
|
||||
|
||||
src := buildPath(pair.SourcePath, srcMachine)
|
||||
dst := buildPath(pair.DestPath, dstMachine)
|
||||
|
||||
cfg := &SyncPairConfig{
|
||||
ID: pair.ID,
|
||||
Name: pair.Name,
|
||||
SourceMachineID: pair.SourceMachineID,
|
||||
SourcePath: pair.SourcePath,
|
||||
DestMachineID: pair.DestMachineID,
|
||||
DestPath: pair.DestPath,
|
||||
Direction: pair.Direction,
|
||||
RsyncFlags: pair.RsyncFlags,
|
||||
ExcludePatterns: pair.ExcludePatternsList(),
|
||||
}
|
||||
cfg.Source = src
|
||||
cfg.Dest = dst
|
||||
|
||||
logPath := filepath.Join(e.cfg.LogsDir(), fmt.Sprintf("%d.log", jobID))
|
||||
f, err := os.Create(logPath)
|
||||
if err != nil {
|
||||
slog.Warn("failed to create log file", "error", err)
|
||||
} else {
|
||||
f.Close()
|
||||
}
|
||||
|
||||
e.setJobStatus(jobID, "waking_up")
|
||||
e.emit(Event{Type: "status", JobID: jobID, Status: "waking_up"})
|
||||
|
||||
var targetMachine *models.Machine
|
||||
var remotePort int
|
||||
if pair.Direction == "pull" && srcMachine != nil {
|
||||
targetMachine = srcMachine
|
||||
remotePort = srcMachine.Port
|
||||
} else if dstMachine != nil {
|
||||
targetMachine = dstMachine
|
||||
remotePort = dstMachine.Port
|
||||
}
|
||||
|
||||
if targetMachine != nil && targetMachine.WoLEnabled && targetMachine.MACAddress != nil {
|
||||
mac, err := wol.ParseMAC(*targetMachine.MACAddress)
|
||||
if err == nil {
|
||||
bcast := ""
|
||||
if targetMachine.BroadcastAddr != nil {
|
||||
bcast = *targetMachine.BroadcastAddr
|
||||
}
|
||||
if err := wol.Send(targetMachine.Host, mac, bcast); err != nil {
|
||||
slog.Warn("WoL failed", "host", targetMachine.Host, "error", err)
|
||||
} else {
|
||||
slog.Info("WoL magic packet sent", "host", targetMachine.Host, "mac", *targetMachine.MACAddress)
|
||||
}
|
||||
|
||||
timeout := time.Duration(targetMachine.WakeTimeoutSeconds) * time.Second
|
||||
interval := time.Duration(targetMachine.WakeCheckIntervalSeconds) * time.Second
|
||||
if err := wol.WaitUntilReady(jobCtx, targetMachine.Host, remotePort, timeout, interval, false); err != nil {
|
||||
e.setJobStatus(jobID, "failed")
|
||||
e.emit(Event{Type: "status", JobID: jobID, Status: "failed", Line: err.Error()})
|
||||
return fmt.Errorf("machine not ready: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
e.setJobStatus(jobID, "running")
|
||||
e.setJobLogFile(jobID, logPath)
|
||||
e.emit(Event{Type: "status", JobID: jobID, Status: "running"})
|
||||
slog.Info("job started", "job_id", jobID, "pair", pair.Name)
|
||||
|
||||
var privKey string
|
||||
if targetMachine != nil && targetMachine.SSHKeyID != nil {
|
||||
privKey = filepath.Join(e.cfg.SSHDir(), "id_ed25519")
|
||||
}
|
||||
|
||||
runner := NewRsyncRunner(e.cfg.SSHDir(), privKey)
|
||||
onLine := func(stream, line string) {
|
||||
e.emit(Event{Type: "log", JobID: jobID, Stream: stream, Line: line})
|
||||
}
|
||||
|
||||
result, err := runner.Run(jobCtx, cfg, onLine)
|
||||
if err != nil {
|
||||
if jobCtx.Err() != nil {
|
||||
e.setJobStatus(jobID, "cancelled")
|
||||
e.emit(Event{Type: "status", JobID: jobID, Status: "cancelled"})
|
||||
return jobCtx.Err()
|
||||
}
|
||||
e.setJobStatus(jobID, "failed")
|
||||
e.emit(Event{Type: "status", JobID: jobID, Status: "failed", Line: err.Error()})
|
||||
return fmt.Errorf("rsync error: %w", err)
|
||||
}
|
||||
|
||||
if result.ExitCode != 0 {
|
||||
e.setJobStatus(jobID, "failed")
|
||||
e.emit(Event{Type: "status", JobID: jobID, Status: "failed", Line: result.Stderr})
|
||||
return fmt.Errorf("rsync exited with code %d: %s", result.ExitCode, result.Stderr)
|
||||
}
|
||||
|
||||
e.setJobStatus(jobID, "success")
|
||||
e.emit(Event{Type: "status", JobID: jobID, Status: "success"})
|
||||
slog.Info("job completed", "job_id", jobID, "pair", pair.Name)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Engine) Cancel(jobID int64, syncPairID int64) bool {
|
||||
if e.queue.IsRunning(syncPairID) {
|
||||
e.queue.Cancel(syncPairID)
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (e *Engine) setJobStatus(jobID int64, status string) {
|
||||
jobRepo := models.NewJobRepository(e.db)
|
||||
jobRepo.UpdateStatus(jobID, status)
|
||||
}
|
||||
|
||||
func (e *Engine) setJobLogFile(jobID int64, path string) {
|
||||
jobRepo := models.NewJobRepository(e.db)
|
||||
jobRepo.SetLogFile(jobID, path)
|
||||
}
|
||||
|
||||
func (e *Engine) emit(evt Event) {
|
||||
select {
|
||||
case e.eventBus <- evt:
|
||||
default:
|
||||
slog.Warn("event bus full, dropping event", "type", evt.Type)
|
||||
}
|
||||
}
|
||||
|
||||
func buildPath(path string, machine *models.Machine) string {
|
||||
if machine == nil {
|
||||
return path
|
||||
}
|
||||
return fmt.Sprintf("%s@%s:%s", machine.SSHUser, machine.Host, path)
|
||||
}
|
||||
|
||||
func (e *Engine) CreateJob(syncPairID int64, triggerType string) (int64, error) {
|
||||
jobRepo := models.NewJobRepository(e.db)
|
||||
id, err := jobRepo.Create(syncPairID, triggerType, "queued")
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package syncengine
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type RsyncStats struct {
|
||||
SentBytes int64
|
||||
ReceivedBytes int64
|
||||
TotalSize int64
|
||||
Speedup float64
|
||||
FilesSent int
|
||||
FilesTotal int
|
||||
}
|
||||
|
||||
type ProgressLine struct {
|
||||
Phase string
|
||||
Percent float64
|
||||
Files int
|
||||
Total int
|
||||
SentBytes int64
|
||||
XferedBytes int64
|
||||
}
|
||||
|
||||
var (
|
||||
progressRegex = regexp.MustCompile(`\s*([\d,]+)\s+([\d,]+)\s+([\d%]+)\s*`)
|
||||
sentRegex = regexp.MustCompile(`sent\s+([\d,]+)\s+bytes`)
|
||||
recvRegex = regexp.MustCompile(`received\s+([\d,]+)\s+bytes`)
|
||||
totalRegex = regexp.MustCompile(`total size is\s+([\d,]+)`)
|
||||
filesRegex = regexp.MustCompile(`Number of files: ([\d,]+)`)
|
||||
)
|
||||
|
||||
func ParseProgressLine(line string) *ProgressLine {
|
||||
if strings.Contains(line, "files to consider") || strings.Contains(line, "files...") {
|
||||
return &ProgressLine{Phase: "scanning"}
|
||||
}
|
||||
if strings.Contains(line, "building file list") {
|
||||
return &ProgressLine{Phase: "listing"}
|
||||
}
|
||||
if strings.Contains(line, "sent") && strings.Contains(line, "bytes") {
|
||||
return &ProgressLine{Phase: "stats"}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ParseStatsLine(line string) (int64, bool) {
|
||||
m := sentRegex.FindStringSubmatch(line)
|
||||
if len(m) >= 2 {
|
||||
n, _ := strconv.ParseInt(strings.ReplaceAll(m[1], ",", ""), 10, 64)
|
||||
return n, true
|
||||
}
|
||||
return 0, false
|
||||
}
|
||||
|
||||
func ParseFinalStats(output string) *RsyncStats {
|
||||
stats := &RsyncStats{}
|
||||
lines := strings.Split(output, "\n")
|
||||
for _, line := range lines {
|
||||
line = strings.TrimSpace(line)
|
||||
if m := sentRegex.FindStringSubmatch(line); len(m) >= 2 {
|
||||
stats.SentBytes, _ = strconv.ParseInt(strings.ReplaceAll(m[1], ",", ""), 10, 64)
|
||||
}
|
||||
if m := recvRegex.FindStringSubmatch(line); len(m) >= 2 {
|
||||
stats.ReceivedBytes, _ = strconv.ParseInt(strings.ReplaceAll(m[1], ",", ""), 10, 64)
|
||||
}
|
||||
if m := totalRegex.FindStringSubmatch(line); len(m) >= 2 {
|
||||
stats.TotalSize, _ = strconv.ParseInt(strings.ReplaceAll(m[1], ",", ""), 10, 64)
|
||||
}
|
||||
if m := filesRegex.FindStringSubmatch(line); len(m) >= 2 {
|
||||
stats.FilesTotal, _ = strconv.Atoi(strings.ReplaceAll(m[1], ",", ""))
|
||||
}
|
||||
}
|
||||
return stats
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package syncengine
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var ErrAlreadyRunning = errors.New("job already running for this sync pair")
|
||||
|
||||
type Queue struct {
|
||||
mu sync.Mutex
|
||||
runs map[int64]*RunInfo
|
||||
}
|
||||
|
||||
type RunInfo struct {
|
||||
JobID int64
|
||||
Cancel func()
|
||||
}
|
||||
|
||||
func NewQueue() *Queue {
|
||||
return &Queue{runs: make(map[int64]*RunInfo)}
|
||||
}
|
||||
|
||||
func (q *Queue) Enqueue(syncPairID, jobID int64, cancel func()) error {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if _, exists := q.runs[syncPairID]; exists {
|
||||
return ErrAlreadyRunning
|
||||
}
|
||||
q.runs[syncPairID] = &RunInfo{JobID: jobID, Cancel: cancel}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (q *Queue) Dequeue(syncPairID int64) {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
delete(q.runs, syncPairID)
|
||||
}
|
||||
|
||||
func (q *Queue) IsRunning(syncPairID int64) bool {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
_, exists := q.runs[syncPairID]
|
||||
return exists
|
||||
}
|
||||
|
||||
func (q *Queue) GetJobID(syncPairID int64) (int64, bool) {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
info, exists := q.runs[syncPairID]
|
||||
if !exists {
|
||||
return 0, false
|
||||
}
|
||||
return info.JobID, true
|
||||
}
|
||||
|
||||
func (q *Queue) Cancel(syncPairID int64) {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
if info, exists := q.runs[syncPairID]; exists && info.Cancel != nil {
|
||||
info.Cancel()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package syncengine
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestQueue(t *testing.T) {
|
||||
q := NewQueue()
|
||||
|
||||
if q.IsRunning(1) {
|
||||
t.Error("queue should be empty")
|
||||
}
|
||||
|
||||
cancelCalled := false
|
||||
cancel := func() { cancelCalled = true }
|
||||
|
||||
err := q.Enqueue(1, 100, cancel)
|
||||
if err != nil {
|
||||
t.Errorf("Enqueue(1) unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if !q.IsRunning(1) {
|
||||
t.Error("queue should contain syncPair 1")
|
||||
}
|
||||
|
||||
jobID, ok := q.GetJobID(1)
|
||||
if !ok || jobID != 100 {
|
||||
t.Errorf("GetJobID(1) = %d, %v, want 100, true", jobID, ok)
|
||||
}
|
||||
|
||||
err = q.Enqueue(1, 200, nil)
|
||||
if err != ErrAlreadyRunning {
|
||||
t.Errorf("Enqueue(1) again = %v, want ErrAlreadyRunning", err)
|
||||
}
|
||||
|
||||
q.Cancel(1)
|
||||
if !cancelCalled {
|
||||
t.Error("Cancel should have called the cancel func")
|
||||
}
|
||||
|
||||
q.Dequeue(1)
|
||||
if q.IsRunning(1) {
|
||||
t.Error("queue should be empty after Dequeue")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
package syncengine
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os/exec"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type RsyncResult struct {
|
||||
ExitCode int
|
||||
Stdout string
|
||||
Stderr string
|
||||
Stats *RsyncStats
|
||||
}
|
||||
|
||||
type RsyncRunner struct {
|
||||
sshDir string
|
||||
privKey string
|
||||
}
|
||||
|
||||
type SyncPairConfig struct {
|
||||
ID int64
|
||||
Name string
|
||||
SourceMachineID *int64
|
||||
SourcePath string
|
||||
DestMachineID *int64
|
||||
DestPath string
|
||||
Direction string
|
||||
RsyncFlags string
|
||||
ExcludePatterns []string
|
||||
Source string
|
||||
Dest string
|
||||
}
|
||||
|
||||
func NewRsyncRunner(sshDir, privKey string) *RsyncRunner {
|
||||
return &RsyncRunner{sshDir: sshDir, privKey: privKey}
|
||||
}
|
||||
|
||||
func (r *RsyncRunner) Run(ctx context.Context, pair *SyncPairConfig, onLine func(stream string, line string)) (*RsyncResult, error) {
|
||||
args := r.buildArgs(pair)
|
||||
|
||||
cmd := exec.CommandContext(ctx, "rsync", args...)
|
||||
if r.privKey != "" {
|
||||
sshCmd := fmt.Sprintf("ssh -i %s -o StrictHostKeyChecking=accept-new -o UserKnownHostsFile=%s",
|
||||
r.privKey, strings.TrimRight(r.sshDir, "/")+"/known_hosts")
|
||||
cmd.Args = append([]string{"rsync", "-e", sshCmd}, args[1:]...)
|
||||
}
|
||||
|
||||
stdout, err := cmd.StdoutPipe()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("stdout pipe: %w", err)
|
||||
}
|
||||
stderr, err := cmd.StderrPipe()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("stderr pipe: %w", err)
|
||||
}
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
return nil, fmt.Errorf("starting rsync: %w", err)
|
||||
}
|
||||
|
||||
var outLines, errLines []string
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
br := io.Reader(stdout)
|
||||
buf := make([]byte, 4096)
|
||||
for {
|
||||
n, err := br.Read(buf)
|
||||
if n > 0 {
|
||||
line := strings.TrimRight(string(buf[:n]), "\r\n")
|
||||
if line != "" {
|
||||
outLines = append(outLines, line)
|
||||
if onLine != nil {
|
||||
onLine("stdout", line)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
select {
|
||||
case <-done:
|
||||
default:
|
||||
close(done)
|
||||
}
|
||||
}()
|
||||
|
||||
go func() {
|
||||
br := io.Reader(stderr)
|
||||
buf := make([]byte, 4096)
|
||||
for {
|
||||
n, err := br.Read(buf)
|
||||
if n > 0 {
|
||||
line := strings.TrimRight(string(buf[:n]), "\r\n")
|
||||
if line != "" {
|
||||
errLines = append(errLines, line)
|
||||
if onLine != nil {
|
||||
onLine("stderr", line)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
select {
|
||||
case <-done:
|
||||
default:
|
||||
close(done)
|
||||
}
|
||||
}()
|
||||
|
||||
err = cmd.Wait()
|
||||
<-done
|
||||
|
||||
result := &RsyncResult{
|
||||
ExitCode: 0,
|
||||
Stdout: strings.Join(outLines, "\n"),
|
||||
Stderr: strings.Join(errLines, "\n"),
|
||||
Stats: ParseFinalStats(strings.Join(outLines, "\n")),
|
||||
}
|
||||
if err != nil {
|
||||
if exitErr, ok := err.(*exec.ExitError); ok {
|
||||
result.ExitCode = exitErr.ExitCode()
|
||||
} else {
|
||||
result.ExitCode = -1
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (r *RsyncRunner) buildArgs(pair *SyncPairConfig) []string {
|
||||
var args []string
|
||||
|
||||
flags := strings.Fields(pair.RsyncFlags)
|
||||
args = append(args, flags...)
|
||||
|
||||
for _, pattern := range pair.ExcludePatterns {
|
||||
args = append(args, "--exclude="+pattern)
|
||||
}
|
||||
|
||||
if pair.Direction == "mirror" {
|
||||
args = append(args, "--delete")
|
||||
}
|
||||
|
||||
if pair.Direction == "pull" {
|
||||
args = append(args, pair.Dest, pair.Source)
|
||||
} else {
|
||||
args = append(args, pair.Source, pair.Dest)
|
||||
}
|
||||
|
||||
return args
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package webui
|
||||
|
||||
import (
|
||||
"embed"
|
||||
"io/fs"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
//go:embed dist
|
||||
var DistFS embed.FS
|
||||
|
||||
var Dist fs.FS = DistFS
|
||||
|
||||
func ServeHTTP() http.Handler {
|
||||
return http.FileServer(http.FS(DistFS))
|
||||
}
|
||||
|
||||
func ServeSPA() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
path := r.URL.Path
|
||||
if path == "/" || !isStaticAsset(path) {
|
||||
data, err := DistFS.ReadFile("dist/index.html")
|
||||
if err != nil {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "text/html")
|
||||
w.Write(data)
|
||||
return
|
||||
}
|
||||
http.FileServer(http.FS(DistFS)).ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func isStaticAsset(path string) bool {
|
||||
exts := []string{".js", ".css", ".png", ".jpg", ".jpeg", ".gif", ".svg", ".ico", ".woff", ".woff2", ".ttf", ".eot", ".map"}
|
||||
for _, ext := range exts {
|
||||
if len(path) > len(ext) && path[len(path)-len(ext):] == ext {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package wol
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"os/exec"
|
||||
"time"
|
||||
)
|
||||
|
||||
var ErrTimeout = fmt.Errorf("timeout waiting for machine to respond")
|
||||
|
||||
func WaitUntilReady(ctx context.Context, host string, sshPort int, maxWait, checkInterval time.Duration, usePing bool) error {
|
||||
deadline := time.Now().Add(maxWait)
|
||||
interval := checkInterval
|
||||
|
||||
ticker := time.NewTicker(checkInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-ticker.C:
|
||||
}
|
||||
|
||||
addr := fmt.Sprintf("%s:%d", host, sshPort)
|
||||
dialer := net.Dialer{Timeout: 3 * time.Second}
|
||||
conn, err := dialer.DialContext(ctx, "tcp", addr)
|
||||
if err == nil {
|
||||
conn.Close()
|
||||
return nil
|
||||
}
|
||||
|
||||
if usePing {
|
||||
cmd := exec.CommandContext(ctx, "ping", "-c", "1", "-W", "1", host)
|
||||
if err := cmd.Run(); err == nil {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
if time.Now().After(deadline) {
|
||||
return ErrTimeout
|
||||
}
|
||||
|
||||
elapsed := time.Since(time.Now().Add(-maxWait))
|
||||
if elapsed > maxWait/2 && interval < 10*time.Second {
|
||||
interval = interval * 3 / 2
|
||||
if interval > 10*time.Second {
|
||||
interval = 10 * time.Second
|
||||
}
|
||||
ticker.Reset(interval)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func IsHostReachable(host string, port int, timeout time.Duration) bool {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
addr := fmt.Sprintf("%s:%d", host, port)
|
||||
dialer := net.Dialer{Timeout: timeout}
|
||||
conn, err := dialer.DialContext(ctx, "tcp", addr)
|
||||
if err == nil {
|
||||
conn.Close()
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package wol
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
var macRegex = regexp.MustCompile(`^([0-9A-Fa-f]{2}[:-]){5}[0-9A-Fa-f]{2}$`)
|
||||
|
||||
func ParseMAC(s string) ([6]byte, error) {
|
||||
s = strings.ReplaceAll(s, "-", ":")
|
||||
s = strings.ToLower(s)
|
||||
if !macRegex.MatchString(s) {
|
||||
return [6]byte{}, fmt.Errorf("invalid MAC address: %s", s)
|
||||
}
|
||||
parts := strings.Split(s, ":")
|
||||
var mac [6]byte
|
||||
for i := 0; i < 6; i++ {
|
||||
var b int
|
||||
if _, err := fmt.Sscanf(parts[i], "%x", &b); err != nil {
|
||||
return [6]byte{}, fmt.Errorf("invalid MAC address: %s", s)
|
||||
}
|
||||
mac[i] = byte(b)
|
||||
}
|
||||
return mac, nil
|
||||
}
|
||||
|
||||
func FormatMAC(mac [6]byte) string {
|
||||
return fmt.Sprintf("%02x:%02x:%02x:%02x:%02x:%02x",
|
||||
mac[0], mac[1], mac[2], mac[3], mac[4], mac[5])
|
||||
}
|
||||
|
||||
func BuildMagicPacket(mac [6]byte) []byte {
|
||||
packet := make([]byte, 6+16*6)
|
||||
for i := 0; i < 6; i++ {
|
||||
packet[i] = 0xFF
|
||||
}
|
||||
for i := 0; i < 16; i++ {
|
||||
offset := 6 + i*6
|
||||
copy(packet[offset:offset+6], mac[:])
|
||||
}
|
||||
return packet
|
||||
}
|
||||
|
||||
func Send(addr string, mac [6]byte, broadcastAddr string) error {
|
||||
packet := BuildMagicPacket(mac)
|
||||
|
||||
udpAddr := &net.UDPAddr{
|
||||
IP: net.IPv4bcast,
|
||||
Port: 9,
|
||||
}
|
||||
if broadcastAddr != "" {
|
||||
udpAddr.IP = net.ParseIP(broadcastAddr)
|
||||
if udpAddr.IP == nil {
|
||||
return fmt.Errorf("invalid broadcast address: %s", broadcastAddr)
|
||||
}
|
||||
}
|
||||
|
||||
conn, err := net.DialUDP("udp4", nil, udpAddr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating UDP connection: %w", err)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
if err := conn.SetWriteDeadline(time.Now().Add(5 * time.Second)); err != nil {
|
||||
return fmt.Errorf("setting broadcast mode: %w", err)
|
||||
}
|
||||
|
||||
n, err := conn.Write(packet)
|
||||
if err != nil {
|
||||
return fmt.Errorf("sending magic packet: %w", err)
|
||||
}
|
||||
if n != len(packet) {
|
||||
return fmt.Errorf("incomplete write: sent %d/%d bytes", n, len(packet))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package wol
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseMAC(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
wantOK bool
|
||||
wantBytes [6]byte
|
||||
}{
|
||||
{"AA:BB:CC:DD:EE:FF", true, [6]byte{0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF}},
|
||||
{"aa:bb:cc:dd:ee:ff", true, [6]byte{0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF}},
|
||||
{"AA-BB-CC-DD-EE-FF", true, [6]byte{0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF}},
|
||||
{"11:22:33:44:55:66", true, [6]byte{0x11, 0x22, 0x33, 0x44, 0x55, 0x66}},
|
||||
{"not:a:mac:addre", false, [6]byte{}},
|
||||
{"GG:HH:II:JJ:KK:LL", false, [6]byte{}},
|
||||
{"aa:bb:cc:dd", false, [6]byte{}},
|
||||
{"", false, [6]byte{}},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
mac, err := ParseMAC(tt.input)
|
||||
if tt.wantOK {
|
||||
if err != nil {
|
||||
t.Errorf("ParseMAC(%q) unexpected error: %v", tt.input, err)
|
||||
continue
|
||||
}
|
||||
if mac != tt.wantBytes {
|
||||
t.Errorf("ParseMAC(%q) = %v, want %v", tt.input, mac, tt.wantBytes)
|
||||
}
|
||||
} else {
|
||||
if err == nil {
|
||||
t.Errorf("ParseMAC(%q) expected error, got nil", tt.input)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFormatMAC(t *testing.T) {
|
||||
mac := [6]byte{0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF}
|
||||
got := FormatMAC(mac)
|
||||
want := "aa:bb:cc:dd:ee:ff"
|
||||
if got != want {
|
||||
t.Errorf("FormatMAC() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildMagicPacket(t *testing.T) {
|
||||
mac := [6]byte{0xAA, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF}
|
||||
packet := BuildMagicPacket(mac)
|
||||
|
||||
if len(packet) != 102 {
|
||||
t.Errorf("BuildMagicPacket length = %d, want 102", len(packet))
|
||||
}
|
||||
|
||||
for i := 0; i < 6; i++ {
|
||||
if packet[i] != 0xFF {
|
||||
t.Errorf("packet[%d] = %02x, want FF", i, packet[i])
|
||||
}
|
||||
}
|
||||
|
||||
for i := 0; i < 16; i++ {
|
||||
offset := 6 + i*6
|
||||
for j := 0; j < 6; j++ {
|
||||
if packet[offset+j] != mac[j] {
|
||||
t.Errorf("packet[%d] = %02x, want %02x (rep %d)", offset+j, packet[offset+j], mac[j], i)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user