@@ -0,0 +1,99 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Entry struct {
|
||||
ID int64 `json:"id"`
|
||||
UserID *int64 `json:"user_id,omitempty"`
|
||||
Actor string `json:"actor"`
|
||||
Action string `json:"action"`
|
||||
Resource string `json:"resource"`
|
||||
Detail map[string]any `json:"detail,omitempty"`
|
||||
IP string `json:"ip"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Status int `json:"status"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
}
|
||||
|
||||
type Service struct{ db *sql.DB }
|
||||
|
||||
func New(db *sql.DB) *Service { return &Service{db: db} }
|
||||
|
||||
func (s *Service) Log(ctx context.Context, e Entry) error {
|
||||
if e.CreatedAt == 0 {
|
||||
e.CreatedAt = time.Now().Unix()
|
||||
}
|
||||
b, _ := json.Marshal(e.Detail)
|
||||
_, err := s.db.ExecContext(ctx, `INSERT INTO audit_log(user_id,actor,action,resource,detail_json,ip,user_agent,status,created_at) VALUES(?,?,?,?,?,?,?,?,?)`, e.UserID, e.Actor, e.Action, e.Resource, string(b), e.IP, e.UserAgent, e.Status, e.CreatedAt)
|
||||
return err
|
||||
}
|
||||
func (s *Service) List(ctx context.Context, limit, offset int, action string) ([]Entry, error) {
|
||||
if limit < 1 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > 500 {
|
||||
limit = 500
|
||||
}
|
||||
if offset < 0 {
|
||||
offset = 0
|
||||
}
|
||||
q := `SELECT id,user_id,actor,action,resource,detail_json,ip,user_agent,status,created_at FROM audit_log`
|
||||
args := []any{}
|
||||
if strings.TrimSpace(action) != "" {
|
||||
q += ` WHERE action LIKE ?`
|
||||
args = append(args, "%"+strings.TrimSpace(action)+"%")
|
||||
}
|
||||
q += ` ORDER BY id DESC LIMIT ? OFFSET ?`
|
||||
args = append(args, limit, offset)
|
||||
rows, err := s.db.QueryContext(ctx, q, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Entry{}
|
||||
for rows.Next() {
|
||||
var e Entry
|
||||
var uid sql.NullInt64
|
||||
var raw string
|
||||
if err := rows.Scan(&e.ID, &uid, &e.Actor, &e.Action, &e.Resource, &raw, &e.IP, &e.UserAgent, &e.Status, &e.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if uid.Valid {
|
||||
v := uid.Int64
|
||||
e.UserID = &v
|
||||
}
|
||||
_ = json.Unmarshal([]byte(raw), &e.Detail)
|
||||
out = append(out, e)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// Run periodically prunes old audit records. A retention of 0 keeps the audit
|
||||
// trail indefinitely.
|
||||
func (s *Service) Run(ctx context.Context, retentionDays int) {
|
||||
if retentionDays <= 0 {
|
||||
<-ctx.Done()
|
||||
return
|
||||
}
|
||||
cleanup := func() {
|
||||
cut := time.Now().Add(-time.Duration(retentionDays) * 24 * time.Hour).Unix()
|
||||
_, _ = s.db.ExecContext(ctx, `DELETE FROM audit_log WHERE created_at<?`, cut)
|
||||
}
|
||||
cleanup()
|
||||
t := time.NewTicker(24 * time.Hour)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
cleanup()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,212 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/config"
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"golang.org/x/oauth2"
|
||||
"net/http"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type contextKey string
|
||||
|
||||
const userKey contextKey = "user"
|
||||
|
||||
type User struct {
|
||||
ID int64 `json:"id"`
|
||||
Sub string `json:"sub"`
|
||||
Email string `json:"email"`
|
||||
Name string `json:"name"`
|
||||
Role string `json:"role"`
|
||||
}
|
||||
type Service struct {
|
||||
cfg config.Config
|
||||
db *sql.DB
|
||||
verifier *oidc.IDTokenVerifier
|
||||
oauth oauth2.Config
|
||||
dev User
|
||||
}
|
||||
|
||||
func New(ctx context.Context, c config.Config, db *sql.DB) (*Service, error) {
|
||||
s := &Service{cfg: c, db: db}
|
||||
if c.Mode == config.ModeAgent {
|
||||
return s, nil
|
||||
}
|
||||
if c.AuthDisabled {
|
||||
now := time.Now().Unix()
|
||||
_, e := db.ExecContext(ctx, `INSERT INTO users(oidc_sub,email,name,role,last_login_at,created_at) VALUES(?,?,?,?,?,?) ON CONFLICT(oidc_sub) DO UPDATE SET last_login_at=excluded.last_login_at`, "dev", "dev@local", "Development Admin", "admin", now, now)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
e = db.QueryRowContext(ctx, `SELECT id,oidc_sub,email,name,role FROM users WHERE oidc_sub='dev'`).Scan(&s.dev.ID, &s.dev.Sub, &s.dev.Email, &s.dev.Name, &s.dev.Role)
|
||||
return s, e
|
||||
}
|
||||
dctx, cancel := context.WithTimeout(ctx, c.HTTPTimeout)
|
||||
defer cancel()
|
||||
p, e := oidc.NewProvider(dctx, c.OIDCIssuer)
|
||||
if e != nil {
|
||||
return nil, fmt.Errorf("oidc discovery: %w", e)
|
||||
}
|
||||
s.verifier = p.Verifier(&oidc.Config{ClientID: c.OIDCClientID})
|
||||
s.oauth = oauth2.Config{ClientID: c.OIDCClientID, ClientSecret: c.OIDCClientSecret, Endpoint: p.Endpoint(), RedirectURL: c.OIDCRedirectURL, Scopes: []string{oidc.ScopeOpenID, "profile", "email", "groups"}}
|
||||
return s, nil
|
||||
}
|
||||
func (s *Service) Login(w http.ResponseWriter, r *http.Request) {
|
||||
if s.cfg.AuthDisabled {
|
||||
http.Redirect(w, r, "/", 302)
|
||||
return
|
||||
}
|
||||
state, err := token(24)
|
||||
if err != nil {
|
||||
http.Error(w, "could not initialize login", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
nonce, err := token(24)
|
||||
if err != nil {
|
||||
http.Error(w, "could not initialize login", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
s.temp(w, "dw_state", state)
|
||||
s.temp(w, "dw_nonce", nonce)
|
||||
http.Redirect(w, r, s.oauth.AuthCodeURL(state, oidc.Nonce(nonce)), 302)
|
||||
}
|
||||
func (s *Service) Callback(w http.ResponseWriter, r *http.Request) error {
|
||||
ctx, cancel := context.WithTimeout(r.Context(), s.cfg.HTTPTimeout)
|
||||
defer cancel()
|
||||
sc, e := r.Cookie("dw_state")
|
||||
if e != nil || sc.Value != r.URL.Query().Get("state") {
|
||||
return errors.New("invalid oidc state")
|
||||
}
|
||||
nc, e := r.Cookie("dw_nonce")
|
||||
if e != nil {
|
||||
return errors.New("missing nonce")
|
||||
}
|
||||
tok, e := s.oauth.Exchange(ctx, r.URL.Query().Get("code"))
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
raw, ok := tok.Extra("id_token").(string)
|
||||
if !ok {
|
||||
return errors.New("missing id_token")
|
||||
}
|
||||
id, e := s.verifier.Verify(ctx, raw)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
var c struct {
|
||||
Sub string `json:"sub"`
|
||||
Email string `json:"email"`
|
||||
Name string `json:"name"`
|
||||
Preferred string `json:"preferred_username"`
|
||||
Nonce string `json:"nonce"`
|
||||
Groups []string `json:"groups"`
|
||||
}
|
||||
if e = id.Claims(&c); e != nil {
|
||||
return e
|
||||
}
|
||||
if c.Nonce != nc.Value {
|
||||
return errors.New("invalid nonce")
|
||||
}
|
||||
s.clearTemp(w, "dw_state")
|
||||
s.clearTemp(w, "dw_nonce")
|
||||
if c.Name == "" {
|
||||
c.Name = c.Preferred
|
||||
}
|
||||
role := "viewer"
|
||||
if slices.Contains(c.Groups, s.cfg.OIDCOperatorGroup) {
|
||||
role = "operator"
|
||||
}
|
||||
if slices.Contains(c.Groups, s.cfg.OIDCAdminGroup) {
|
||||
role = "admin"
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
_, e = s.db.ExecContext(ctx, `INSERT INTO users(oidc_sub,email,name,role,last_login_at,created_at) VALUES(?,?,?,?,?,?) ON CONFLICT(oidc_sub) DO UPDATE SET email=excluded.email,name=excluded.name,role=excluded.role,last_login_at=excluded.last_login_at`, c.Sub, c.Email, c.Name, role, now, now)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
var uid int64
|
||||
if e = s.db.QueryRowContext(ctx, `SELECT id FROM users WHERE oidc_sub=?`, c.Sub).Scan(&uid); e != nil {
|
||||
return e
|
||||
}
|
||||
v, e := token(32)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
h := sha256.Sum256([]byte(v))
|
||||
exp := time.Now().Add(12 * time.Hour)
|
||||
_, e = s.db.ExecContext(ctx, `INSERT INTO sessions(token_hash,user_id,expires_at,created_at) VALUES(?,?,?,?)`, h[:], uid, exp.Unix(), now)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
http.SetCookie(w, &http.Cookie{Name: "dockwatch_session", Value: v, Path: "/", HttpOnly: true, Secure: s.cfg.SecureCookies, SameSite: http.SameSiteLaxMode, Expires: exp, MaxAge: int(time.Until(exp).Seconds())})
|
||||
return nil
|
||||
}
|
||||
func (s *Service) Logout(w http.ResponseWriter, r *http.Request) {
|
||||
if c, e := r.Cookie("dockwatch_session"); e == nil {
|
||||
h := sha256.Sum256([]byte(c.Value))
|
||||
_, _ = s.db.ExecContext(r.Context(), `DELETE FROM sessions WHERE token_hash=?`, h[:])
|
||||
}
|
||||
http.SetCookie(w, &http.Cookie{Name: "dockwatch_session", Path: "/", HttpOnly: true, Secure: s.cfg.SecureCookies, SameSite: http.SameSiteLaxMode, MaxAge: -1})
|
||||
}
|
||||
func (s *Service) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if s.cfg.AuthDisabled {
|
||||
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), userKey, s.dev)))
|
||||
return
|
||||
}
|
||||
c, e := r.Cookie("dockwatch_session")
|
||||
if e != nil {
|
||||
http.Error(w, "unauthorized", 401)
|
||||
return
|
||||
}
|
||||
h := sha256.Sum256([]byte(c.Value))
|
||||
var u User
|
||||
e = s.db.QueryRowContext(r.Context(), `SELECT u.id,u.oidc_sub,u.email,u.name,u.role FROM sessions s JOIN users u ON u.id=s.user_id WHERE s.token_hash=? AND s.expires_at>?`, h[:], time.Now().Unix()).Scan(&u.ID, &u.Sub, &u.Email, &u.Name, &u.Role)
|
||||
if e != nil {
|
||||
http.Error(w, "unauthorized", 401)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), userKey, u)))
|
||||
})
|
||||
}
|
||||
func UserFrom(ctx context.Context) (User, bool) { u, ok := ctx.Value(userKey).(User); return u, ok }
|
||||
func RequireRole(min string, next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
u, ok := UserFrom(r.Context())
|
||||
if !ok || rank(u.Role) < rank(min) {
|
||||
http.Error(w, "forbidden", 403)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
func rank(r string) int {
|
||||
switch strings.ToLower(r) {
|
||||
case "admin":
|
||||
return 3
|
||||
case "operator":
|
||||
return 2
|
||||
default:
|
||||
return 1
|
||||
}
|
||||
}
|
||||
func token(n int) (string, error) {
|
||||
b := make([]byte, n)
|
||||
_, e := rand.Read(b)
|
||||
return base64.RawURLEncoding.EncodeToString(b), e
|
||||
}
|
||||
func (s *Service) temp(w http.ResponseWriter, n, v string) {
|
||||
http.SetCookie(w, &http.Cookie{Name: n, Value: v, Path: "/", HttpOnly: true, Secure: s.cfg.SecureCookies, SameSite: http.SameSiteLaxMode, MaxAge: 600})
|
||||
}
|
||||
func (s *Service) clearTemp(w http.ResponseWriter, n string) {
|
||||
http.SetCookie(w, &http.Cookie{Name: n, Path: "/", HttpOnly: true, Secure: s.cfg.SecureCookies, SameSite: http.SameSiteLaxMode, MaxAge: -1})
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package buildinfo
|
||||
|
||||
import "runtime"
|
||||
|
||||
// Values may be overridden at build time with -ldflags -X.
|
||||
var (
|
||||
Version = "0.9.0"
|
||||
Commit = "dev"
|
||||
Date = "unknown"
|
||||
)
|
||||
|
||||
type Info struct {
|
||||
Version string `json:"version"`
|
||||
Commit string `json:"commit"`
|
||||
BuildDate string `json:"build_date"`
|
||||
GoVersion string `json:"go_version"`
|
||||
}
|
||||
|
||||
func Current() Info {
|
||||
return Info{Version: Version, Commit: Commit, BuildDate: Date, GoVersion: runtime.Version()}
|
||||
}
|
||||
@@ -0,0 +1,215 @@
|
||||
package composeedit
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
type ParseResult struct {
|
||||
Value any `json:"value"`
|
||||
}
|
||||
|
||||
type Patch struct {
|
||||
Path []string `json:"path"`
|
||||
Value any `json:"value"`
|
||||
Delete bool `json:"delete"`
|
||||
}
|
||||
|
||||
func Parse(src string) (ParseResult, error) {
|
||||
var doc yaml.Node
|
||||
if err := yaml.Unmarshal([]byte(src), &doc); err != nil {
|
||||
return ParseResult{}, fmt.Errorf("yaml: %w", err)
|
||||
}
|
||||
if len(doc.Content) == 0 {
|
||||
return ParseResult{Value: map[string]any{}}, nil
|
||||
}
|
||||
var v any
|
||||
if err := doc.Content[0].Decode(&v); err != nil {
|
||||
return ParseResult{}, err
|
||||
}
|
||||
v = normalize(v)
|
||||
return ParseResult{Value: v}, nil
|
||||
}
|
||||
|
||||
func Apply(src string, p Patch) (string, error) {
|
||||
if len(p.Path) == 0 {
|
||||
return "", errors.New("path is required")
|
||||
}
|
||||
var doc yaml.Node
|
||||
if err := yaml.Unmarshal([]byte(src), &doc); err != nil {
|
||||
return "", fmt.Errorf("yaml: %w", err)
|
||||
}
|
||||
if len(doc.Content) == 0 {
|
||||
return "", errors.New("empty yaml document")
|
||||
}
|
||||
root := doc.Content[0]
|
||||
parent, last, err := walkParent(root, p.Path, !p.Delete)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if p.Delete {
|
||||
if err := deleteChild(parent, last); err != nil {
|
||||
return "", err
|
||||
}
|
||||
} else {
|
||||
n, err := valueNode(p.Value)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := setChild(parent, last, n); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
var b bytes.Buffer
|
||||
enc := yaml.NewEncoder(&b)
|
||||
enc.SetIndent(2)
|
||||
if err := enc.Encode(&doc); err != nil {
|
||||
return "", err
|
||||
}
|
||||
_ = enc.Close()
|
||||
return strings.TrimSuffix(b.String(), "\n") + "\n", nil
|
||||
}
|
||||
|
||||
func normalize(v any) any {
|
||||
switch x := v.(type) {
|
||||
case map[string]any:
|
||||
out := make(map[string]any, len(x))
|
||||
for k, v := range x {
|
||||
out[k] = normalize(v)
|
||||
}
|
||||
return out
|
||||
case map[any]any:
|
||||
out := map[string]any{}
|
||||
for k, v := range x {
|
||||
out[fmt.Sprint(k)] = normalize(v)
|
||||
}
|
||||
return out
|
||||
case []any:
|
||||
out := make([]any, len(x))
|
||||
for i, v := range x {
|
||||
out[i] = normalize(v)
|
||||
}
|
||||
return out
|
||||
default:
|
||||
return x
|
||||
}
|
||||
}
|
||||
|
||||
func walkParent(root *yaml.Node, path []string, create bool) (*yaml.Node, string, error) {
|
||||
cur := root
|
||||
for _, seg := range path[:len(path)-1] {
|
||||
if cur.Kind == yaml.DocumentNode && len(cur.Content) > 0 {
|
||||
cur = cur.Content[0]
|
||||
}
|
||||
switch cur.Kind {
|
||||
case yaml.MappingNode:
|
||||
n := mapGet(cur, seg)
|
||||
if n == nil {
|
||||
if !create {
|
||||
return nil, "", fmt.Errorf("path %q not found", seg)
|
||||
}
|
||||
n = &yaml.Node{Kind: yaml.MappingNode, Tag: "!!map"}
|
||||
mapSet(cur, seg, n)
|
||||
}
|
||||
cur = n
|
||||
case yaml.SequenceNode:
|
||||
i, err := strconv.Atoi(seg)
|
||||
if err != nil || i < 0 || i >= len(cur.Content) {
|
||||
return nil, "", fmt.Errorf("invalid array index %q", seg)
|
||||
}
|
||||
cur = cur.Content[i]
|
||||
default:
|
||||
return nil, "", fmt.Errorf("cannot descend through scalar at %q", seg)
|
||||
}
|
||||
}
|
||||
return cur, path[len(path)-1], nil
|
||||
}
|
||||
func mapGet(m *yaml.Node, key string) *yaml.Node {
|
||||
for i := 0; i+1 < len(m.Content); i += 2 {
|
||||
if m.Content[i].Value == key {
|
||||
return m.Content[i+1]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func mapSet(m *yaml.Node, key string, v *yaml.Node) {
|
||||
for i := 0; i+1 < len(m.Content); i += 2 {
|
||||
if m.Content[i].Value == key {
|
||||
old := m.Content[i+1]
|
||||
v.HeadComment = old.HeadComment
|
||||
v.LineComment = old.LineComment
|
||||
v.FootComment = old.FootComment
|
||||
m.Content[i+1] = v
|
||||
return
|
||||
}
|
||||
}
|
||||
m.Content = append(m.Content, &yaml.Node{Kind: yaml.ScalarNode, Tag: "!!str", Value: key}, v)
|
||||
}
|
||||
func setChild(p *yaml.Node, key string, v *yaml.Node) error {
|
||||
switch p.Kind {
|
||||
case yaml.MappingNode:
|
||||
mapSet(p, key, v)
|
||||
return nil
|
||||
case yaml.SequenceNode:
|
||||
if key == "-" {
|
||||
p.Content = append(p.Content, v)
|
||||
return nil
|
||||
}
|
||||
i, err := strconv.Atoi(key)
|
||||
if err != nil || i < 0 || i > len(p.Content) {
|
||||
return fmt.Errorf("invalid array index %q", key)
|
||||
}
|
||||
if i == len(p.Content) {
|
||||
p.Content = append(p.Content, v)
|
||||
} else {
|
||||
old := p.Content[i]
|
||||
v.HeadComment, v.LineComment, v.FootComment = old.HeadComment, old.LineComment, old.FootComment
|
||||
p.Content[i] = v
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return errors.New("parent is not a map or array")
|
||||
}
|
||||
}
|
||||
func deleteChild(p *yaml.Node, key string) error {
|
||||
switch p.Kind {
|
||||
case yaml.MappingNode:
|
||||
for i := 0; i+1 < len(p.Content); i += 2 {
|
||||
if p.Content[i].Value == key {
|
||||
p.Content = append(p.Content[:i], p.Content[i+2:]...)
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil
|
||||
case yaml.SequenceNode:
|
||||
i, err := strconv.Atoi(key)
|
||||
if err != nil || i < 0 || i >= len(p.Content) {
|
||||
return fmt.Errorf("invalid array index %q", key)
|
||||
}
|
||||
p.Content = append(p.Content[:i], p.Content[i+1:]...)
|
||||
return nil
|
||||
default:
|
||||
return errors.New("parent is not a map or array")
|
||||
}
|
||||
}
|
||||
func valueNode(v any) (*yaml.Node, error) {
|
||||
raw, err := json.Marshal(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var x any
|
||||
if err := json.Unmarshal(raw, &x); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var n yaml.Node
|
||||
if err := n.Encode(x); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &n, nil
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package composeedit
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestPatchPreservesUnknownAndComments(t *testing.T) {
|
||||
src := "# top\nservices:\n app:\n image: nginx:old # keep\n x-future:\n magic: true\n deploy:\n replicas: 2\n"
|
||||
out, e := Apply(src, Patch{Path: []string{"services", "app", "image"}, Value: "nginx:new"})
|
||||
if e != nil {
|
||||
t.Fatal(e)
|
||||
}
|
||||
if !strings.Contains(out, "x-future:") || !strings.Contains(out, "magic: true") || !strings.Contains(out, "replicas: 2") || !strings.Contains(out, "# top") {
|
||||
t.Fatalf("preservation failed:\n%s", out)
|
||||
}
|
||||
r, e := Parse(out)
|
||||
if e != nil || r.Value == nil {
|
||||
t.Fatalf("parse %v", e)
|
||||
}
|
||||
}
|
||||
func TestArrayPatch(t *testing.T) {
|
||||
src := "services:\n app:\n ports:\n - 8080:80\n"
|
||||
out, e := Apply(src, Patch{Path: []string{"services", "app", "ports", "0"}, Value: "9090:80"})
|
||||
if e != nil {
|
||||
t.Fatal(e)
|
||||
}
|
||||
if !strings.Contains(out, "9090:80") {
|
||||
t.Fatal(out)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Mode string
|
||||
|
||||
const (
|
||||
ModeStandalone Mode = "standalone"
|
||||
ModeMaster Mode = "master"
|
||||
ModeAgent Mode = "agent"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
Mode Mode
|
||||
ListenAddr, BaseURL, DataDir, StacksDir, AppSecret string
|
||||
SecureCookies, AuthDisabled bool
|
||||
OIDCIssuer, OIDCClientID, OIDCClientSecret, OIDCRedirectURL, OIDCAdminGroup, OIDCOperatorGroup string
|
||||
AgentToken string
|
||||
CheckConcurrency, RetentionDays, AuditRetentionDays int
|
||||
HTTPTimeout time.Duration
|
||||
}
|
||||
|
||||
func Load() (Config, error) {
|
||||
checkConcurrency, err := envIntStrict("CHECK_CONCURRENCY", 8)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
retentionDays, err := envIntStrict("CHECK_RETENTION_DAYS", 30)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
auditRetentionDays, err := envIntStrict("AUDIT_RETENTION_DAYS", 180)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
httpTimeoutSeconds, err := envIntStrict("HTTP_TIMEOUT_SECONDS", 10)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
authDisabled, err := envBoolStrict("AUTH_DISABLED", false)
|
||||
if err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
c := Config{
|
||||
Mode: Mode(env("APP_MODE", "standalone")),
|
||||
ListenAddr: env("LISTEN_ADDR", ":8080"),
|
||||
BaseURL: strings.TrimRight(env("BASE_URL", "http://localhost:8080"), "/"),
|
||||
DataDir: env("DATA_DIR", "/data"),
|
||||
StacksDir: env("STACKS_DIR", "/stacks"),
|
||||
AppSecret: os.Getenv("APP_SECRET"),
|
||||
AuthDisabled: authDisabled,
|
||||
OIDCIssuer: strings.TrimRight(os.Getenv("OIDC_ISSUER"), "/"),
|
||||
OIDCClientID: os.Getenv("OIDC_CLIENT_ID"),
|
||||
OIDCClientSecret: os.Getenv("OIDC_CLIENT_SECRET"),
|
||||
OIDCRedirectURL: os.Getenv("OIDC_REDIRECT_URL"),
|
||||
OIDCAdminGroup: env("OIDC_ADMIN_GROUP", "dockwatch-admins"),
|
||||
OIDCOperatorGroup: env("OIDC_OPERATOR_GROUP", "dockwatch-operators"),
|
||||
AgentToken: os.Getenv("AGENT_TOKEN"),
|
||||
CheckConcurrency: checkConcurrency,
|
||||
RetentionDays: retentionDays,
|
||||
AuditRetentionDays: auditRetentionDays,
|
||||
HTTPTimeout: time.Duration(httpTimeoutSeconds) * time.Second,
|
||||
}
|
||||
c.SecureCookies = strings.HasPrefix(c.BaseURL, "https://")
|
||||
if c.OIDCRedirectURL == "" {
|
||||
c.OIDCRedirectURL = c.BaseURL + "/auth/callback"
|
||||
}
|
||||
switch c.Mode {
|
||||
case ModeStandalone, ModeMaster, ModeAgent:
|
||||
default:
|
||||
return c, fmt.Errorf("APP_MODE must be standalone, master, or agent")
|
||||
}
|
||||
if c.CheckConcurrency < 1 || c.CheckConcurrency > 128 {
|
||||
return c, fmt.Errorf("CHECK_CONCURRENCY must be between 1 and 128")
|
||||
}
|
||||
if c.RetentionDays < 0 || c.RetentionDays > 3650 {
|
||||
return c, fmt.Errorf("CHECK_RETENTION_DAYS must be between 0 and 3650")
|
||||
}
|
||||
if c.AuditRetentionDays < 0 || c.AuditRetentionDays > 3650 {
|
||||
return c, fmt.Errorf("AUDIT_RETENTION_DAYS must be between 0 and 3650")
|
||||
}
|
||||
if c.HTTPTimeout < time.Second || c.HTTPTimeout > 5*time.Minute {
|
||||
return c, fmt.Errorf("HTTP_TIMEOUT_SECONDS must be between 1 and 300")
|
||||
}
|
||||
if c.Mode != ModeAgent {
|
||||
u, err := url.Parse(c.BaseURL)
|
||||
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" {
|
||||
return c, errors.New("BASE_URL must be an absolute http(s) URL without credentials, query or fragment")
|
||||
}
|
||||
}
|
||||
if c.Mode == ModeAgent {
|
||||
if len(c.AgentToken) < 24 {
|
||||
return c, errors.New("AGENT_TOKEN must be at least 24 characters in agent mode")
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
if len(c.AppSecret) < 32 {
|
||||
return c, errors.New("APP_SECRET must be at least 32 characters")
|
||||
}
|
||||
if !c.AuthDisabled && (c.OIDCIssuer == "" || c.OIDCClientID == "" || c.OIDCClientSecret == "") {
|
||||
return c, errors.New("OIDC_ISSUER, OIDC_CLIENT_ID and OIDC_CLIENT_SECRET are required unless AUTH_DISABLED=true")
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
func (c Config) DBPath() string { return c.DataDir + "/dockwatch.db" }
|
||||
func (c Config) EncryptionKey() []byte { s := sha256.Sum256([]byte(c.AppSecret)); return s[:] }
|
||||
func (c Config) SecretFingerprint() string {
|
||||
h := sha256.Sum256([]byte(c.AppSecret))
|
||||
return base64.RawURLEncoding.EncodeToString(h[:6])
|
||||
}
|
||||
func env(k, f string) string {
|
||||
if v := os.Getenv(k); v != "" {
|
||||
return v
|
||||
}
|
||||
return f
|
||||
}
|
||||
func envIntStrict(k string, f int) (int, error) {
|
||||
v := strings.TrimSpace(os.Getenv(k))
|
||||
if v == "" {
|
||||
return f, nil
|
||||
}
|
||||
n, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("%s must be an integer: %w", k, err)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
func envBoolStrict(k string, f bool) (bool, error) {
|
||||
v := strings.TrimSpace(os.Getenv(k))
|
||||
if v == "" {
|
||||
return f, nil
|
||||
}
|
||||
b, err := strconv.ParseBool(v)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("%s must be a boolean: %w", k, err)
|
||||
}
|
||||
return b, nil
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package config
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestStrictEnvironmentParsing(t *testing.T) {
|
||||
t.Setenv("CHECK_CONCURRENCY", "not-a-number")
|
||||
if _, err := Load(); err == nil {
|
||||
t.Fatal("expected invalid CHECK_CONCURRENCY to fail instead of silently using a default")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStrictBooleanParsing(t *testing.T) {
|
||||
t.Setenv("CHECK_CONCURRENCY", "8")
|
||||
t.Setenv("AUTH_DISABLED", "sometimes")
|
||||
if _, err := Load(); err == nil {
|
||||
t.Fatal("expected invalid AUTH_DISABLED to fail")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
_ "modernc.org/sqlite"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
)
|
||||
|
||||
func Open(path string) (*sql.DB, error) {
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0750); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dsn := "file:" + path + "?_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)&_pragma=foreign_keys(ON)"
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// WAL allows concurrent readers while SQLite still serializes writes. A small
|
||||
// pool keeps UI reads responsive while monitor checks are being persisted.
|
||||
db.SetMaxOpenConns(4)
|
||||
db.SetMaxIdleConns(4)
|
||||
db.SetConnMaxLifetime(30 * time.Minute)
|
||||
ctx, c := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer c()
|
||||
if err := db.PingContext(ctx); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
if err := migrate(ctx, db); err != nil {
|
||||
_ = db.Close()
|
||||
return nil, err
|
||||
}
|
||||
return db, nil
|
||||
}
|
||||
func migrate(ctx context.Context, db *sql.DB) error {
|
||||
stmts := []string{
|
||||
`CREATE TABLE IF NOT EXISTS users(id INTEGER PRIMARY KEY AUTOINCREMENT,oidc_sub TEXT NOT NULL UNIQUE,email TEXT NOT NULL DEFAULT '',name TEXT NOT NULL DEFAULT '',role TEXT NOT NULL DEFAULT 'viewer',last_login_at INTEGER NOT NULL,created_at INTEGER NOT NULL)`,
|
||||
`CREATE TABLE IF NOT EXISTS sessions(token_hash BLOB PRIMARY KEY,user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,expires_at INTEGER NOT NULL,created_at INTEGER NOT NULL)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_sessions_expires ON sessions(expires_at)`,
|
||||
`CREATE TABLE IF NOT EXISTS nodes(id INTEGER PRIMARY KEY AUTOINCREMENT,name TEXT NOT NULL UNIQUE,base_url TEXT NOT NULL,token_enc BLOB NOT NULL,enabled INTEGER NOT NULL DEFAULT 1,created_at INTEGER NOT NULL,updated_at INTEGER NOT NULL)`,
|
||||
`CREATE TABLE IF NOT EXISTS monitor_services(id INTEGER PRIMARY KEY AUTOINCREMENT,name TEXT NOT NULL UNIQUE,description TEXT NOT NULL DEFAULT '',created_at INTEGER NOT NULL,updated_at INTEGER NOT NULL)`,
|
||||
`CREATE TABLE IF NOT EXISTS monitors(id INTEGER PRIMARY KEY AUTOINCREMENT,name TEXT NOT NULL,type TEXT NOT NULL,target TEXT NOT NULL,node_id INTEGER NULL REFERENCES nodes(id) ON DELETE SET NULL,service_id INTEGER NULL REFERENCES monitor_services(id) ON DELETE SET NULL,interval_seconds INTEGER NOT NULL DEFAULT 60,timeout_ms INTEGER NOT NULL DEFAULT 5000,expected_min INTEGER NOT NULL DEFAULT 200,expected_max INTEGER NOT NULL DEFAULT 399,method TEXT NOT NULL DEFAULT 'GET',headers_json TEXT NOT NULL DEFAULT '{}',body TEXT NOT NULL DEFAULT '',keyword TEXT NOT NULL DEFAULT '',invert_keyword INTEGER NOT NULL DEFAULT 0,ignore_tls INTEGER NOT NULL DEFAULT 0,require_healthy INTEGER NOT NULL DEFAULT 0,enabled INTEGER NOT NULL DEFAULT 1,status TEXT NOT NULL DEFAULT 'pending',maintenance_until INTEGER NULL,maintenance_note TEXT NOT NULL DEFAULT '',last_checked_at INTEGER NULL,created_by INTEGER NULL REFERENCES users(id) ON DELETE SET NULL,created_at INTEGER NOT NULL,updated_at INTEGER NOT NULL)`,
|
||||
`CREATE TABLE IF NOT EXISTS monitor_checks(id INTEGER PRIMARY KEY AUTOINCREMENT,monitor_id INTEGER NOT NULL REFERENCES monitors(id) ON DELETE CASCADE,ok INTEGER NOT NULL,status_code INTEGER NOT NULL DEFAULT 0,latency_ms INTEGER NOT NULL DEFAULT 0,message TEXT NOT NULL DEFAULT '',checked_at INTEGER NOT NULL)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_monitor_checks_mon_time ON monitor_checks(monitor_id,checked_at DESC)`,
|
||||
`CREATE TABLE IF NOT EXISTS audit_log(id INTEGER PRIMARY KEY AUTOINCREMENT,user_id INTEGER NULL REFERENCES users(id) ON DELETE SET NULL,actor TEXT NOT NULL DEFAULT '',action TEXT NOT NULL,resource TEXT NOT NULL DEFAULT '',detail_json TEXT NOT NULL DEFAULT '{}',ip TEXT NOT NULL DEFAULT '',user_agent TEXT NOT NULL DEFAULT '',status INTEGER NOT NULL DEFAULT 0,created_at INTEGER NOT NULL)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_audit_created ON audit_log(created_at DESC)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_audit_action ON audit_log(action,created_at DESC)`,
|
||||
`CREATE TABLE IF NOT EXISTS notification_channels(id INTEGER PRIMARY KEY AUTOINCREMENT,name TEXT NOT NULL UNIQUE,type TEXT NOT NULL,config_json TEXT NOT NULL DEFAULT '{}',enabled INTEGER NOT NULL DEFAULT 1,created_at INTEGER NOT NULL,updated_at INTEGER NOT NULL)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_notifications_enabled ON notification_channels(enabled,type)`,
|
||||
`CREATE TABLE IF NOT EXISTS status_pages(id INTEGER PRIMARY KEY AUTOINCREMENT,name TEXT NOT NULL,slug TEXT NOT NULL UNIQUE,description TEXT NOT NULL DEFAULT '',enabled INTEGER NOT NULL DEFAULT 1,created_at INTEGER NOT NULL,updated_at INTEGER NOT NULL)`,
|
||||
`CREATE TABLE IF NOT EXISTS status_page_services(page_id INTEGER NOT NULL REFERENCES status_pages(id) ON DELETE CASCADE,service_id INTEGER NOT NULL REFERENCES monitor_services(id) ON DELETE CASCADE,sort_order INTEGER NOT NULL DEFAULT 0,PRIMARY KEY(page_id,service_id))`,
|
||||
`CREATE TABLE IF NOT EXISTS git_sources(id INTEGER PRIMARY KEY AUTOINCREMENT,stack_name TEXT NOT NULL UNIQUE,repo_url TEXT NOT NULL,branch TEXT NOT NULL DEFAULT 'main',workdir TEXT NOT NULL DEFAULT '.',compose_file TEXT NOT NULL DEFAULT 'compose.yaml',auto_deploy INTEGER NOT NULL DEFAULT 0,webhook_secret_enc BLOB NOT NULL,last_commit TEXT NOT NULL DEFAULT '',last_sync_at INTEGER NULL,last_error TEXT NOT NULL DEFAULT '',created_at INTEGER NOT NULL,updated_at INTEGER NOT NULL)`}
|
||||
for i, s := range stmts {
|
||||
if _, err := db.ExecContext(ctx, s); err != nil {
|
||||
return fmt.Errorf("migration %d: %w", i+1, err)
|
||||
}
|
||||
}
|
||||
// Additive migrations keep existing installations compatible.
|
||||
cols := map[string]string{
|
||||
"method": "TEXT NOT NULL DEFAULT 'GET'",
|
||||
"headers_json": "TEXT NOT NULL DEFAULT '{}'",
|
||||
"body": "TEXT NOT NULL DEFAULT ''",
|
||||
"keyword": "TEXT NOT NULL DEFAULT ''",
|
||||
"invert_keyword": "INTEGER NOT NULL DEFAULT 0",
|
||||
"ignore_tls": "INTEGER NOT NULL DEFAULT 0",
|
||||
"maintenance_until": "INTEGER NULL",
|
||||
"maintenance_note": "TEXT NOT NULL DEFAULT ''",
|
||||
"service_id": "INTEGER NULL REFERENCES monitor_services(id) ON DELETE SET NULL",
|
||||
"require_healthy": "INTEGER NOT NULL DEFAULT 0",
|
||||
}
|
||||
for name, def := range cols {
|
||||
if err := ensureColumn(ctx, db, "monitors", name, def); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := ensureColumn(ctx, db, "git_sources", "node_id", "INTEGER NULL REFERENCES nodes(id) ON DELETE SET NULL"); err != nil {
|
||||
return err
|
||||
}
|
||||
indexes := []string{
|
||||
`CREATE INDEX IF NOT EXISTS idx_monitors_service ON monitors(service_id,name)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_monitors_schedule ON monitors(enabled,last_checked_at,interval_seconds)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_monitors_status ON monitors(status)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_status_page_services_page ON status_page_services(page_id,sort_order)`,
|
||||
`CREATE INDEX IF NOT EXISTS idx_git_sources_node ON git_sources(node_id)`,
|
||||
}
|
||||
for _, stmt := range indexes {
|
||||
if _, err := db.ExecContext(ctx, stmt); err != nil {
|
||||
return fmt.Errorf("create index: %w", err)
|
||||
}
|
||||
}
|
||||
_, _ = db.ExecContext(ctx, `PRAGMA optimize`)
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureColumn(ctx context.Context, db *sql.DB, table, column, definition string) error {
|
||||
rows, err := db.QueryContext(ctx, "PRAGMA table_info("+table+")")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
found := false
|
||||
for rows.Next() {
|
||||
var cid int
|
||||
var name, typ string
|
||||
var notnull, pk int
|
||||
var dflt any
|
||||
if err := rows.Scan(&cid, &name, &typ, ¬null, &dflt, &pk); err != nil {
|
||||
_ = rows.Close()
|
||||
return err
|
||||
}
|
||||
if name == column {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
_ = rows.Close()
|
||||
return err
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
if found {
|
||||
return nil
|
||||
}
|
||||
if _, err := db.ExecContext(ctx, "ALTER TABLE "+table+" ADD COLUMN "+column+" "+definition); err != nil {
|
||||
return fmt.Errorf("add %s.%s: %w", table, column, err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,57 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestFreshDatabaseHasCurrentMonitorSchemaAndIndexes(t *testing.T) {
|
||||
db, err := Open(filepath.Join(t.TempDir(), "dockwatch.db"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
ctx := context.Background()
|
||||
|
||||
rows, err := db.QueryContext(ctx, `PRAGMA table_info(monitors)`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cols := map[string]bool{}
|
||||
for rows.Next() {
|
||||
var cid, notnull, pk int
|
||||
var name, typ string
|
||||
var dflt any
|
||||
if err := rows.Scan(&cid, &name, &typ, ¬null, &dflt, &pk); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cols[name] = true
|
||||
}
|
||||
_ = rows.Close()
|
||||
for _, name := range []string{"service_id", "method", "headers_json", "maintenance_until", "require_healthy"} {
|
||||
if !cols[name] {
|
||||
t.Fatalf("fresh monitors schema is missing %s", name)
|
||||
}
|
||||
}
|
||||
|
||||
idx, err := db.QueryContext(ctx, `PRAGMA index_list(monitors)`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
indexes := map[string]bool{}
|
||||
for idx.Next() {
|
||||
var seq, unique, partial int
|
||||
var name, origin string
|
||||
if err := idx.Scan(&seq, &name, &unique, &origin, &partial); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
indexes[name] = true
|
||||
}
|
||||
_ = idx.Close()
|
||||
for _, name := range []string{"idx_monitors_service", "idx_monitors_schedule", "idx_monitors_status"} {
|
||||
if !indexes[name] {
|
||||
t.Fatalf("fresh database is missing index %s", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,765 @@
|
||||
package gitops
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/nodes"
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/stacks"
|
||||
)
|
||||
|
||||
var stackNameRx = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._-]{0,63}$`)
|
||||
|
||||
type Source struct {
|
||||
ID int64 `json:"id"`
|
||||
NodeID *int64 `json:"node_id,omitempty"`
|
||||
StackName string `json:"stack_name"`
|
||||
RepoURL string `json:"repo_url"`
|
||||
Branch string `json:"branch"`
|
||||
Workdir string `json:"workdir"`
|
||||
ComposeFile string `json:"compose_file"`
|
||||
AutoDeploy bool `json:"auto_deploy"`
|
||||
LastCommit string `json:"last_commit"`
|
||||
LastSyncAt *int64 `json:"last_sync_at,omitempty"`
|
||||
LastError string `json:"last_error"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
UpdatedAt int64 `json:"updated_at"`
|
||||
}
|
||||
type Input struct {
|
||||
NodeID *int64 `json:"node_id"`
|
||||
StackName string `json:"stack_name"`
|
||||
RepoURL string `json:"repo_url"`
|
||||
Branch string `json:"branch"`
|
||||
Workdir string `json:"workdir"`
|
||||
ComposeFile string `json:"compose_file"`
|
||||
AutoDeploy bool `json:"auto_deploy"`
|
||||
}
|
||||
type Service struct {
|
||||
db *sql.DB
|
||||
key []byte
|
||||
stacks *stacks.Service
|
||||
nodes *nodes.Manager
|
||||
locks sync.Map
|
||||
}
|
||||
|
||||
func New(db *sql.DB, key []byte, ss *stacks.Service, nm *nodes.Manager) *Service {
|
||||
return &Service{db: db, key: key, stacks: ss, nodes: nm}
|
||||
}
|
||||
func normalize(in *Input) error {
|
||||
in.StackName = strings.TrimSpace(in.StackName)
|
||||
in.RepoURL = strings.TrimSpace(in.RepoURL)
|
||||
in.Branch = strings.TrimSpace(in.Branch)
|
||||
in.Workdir = filepath.Clean(strings.TrimSpace(in.Workdir))
|
||||
in.ComposeFile = filepath.Clean(strings.TrimSpace(in.ComposeFile))
|
||||
if in.StackName == "" || in.RepoURL == "" {
|
||||
return errors.New("stack_name and repo_url required")
|
||||
}
|
||||
if !stackNameRx.MatchString(in.StackName) || strings.ContainsAny(in.RepoURL, "\r\n") || strings.HasPrefix(in.RepoURL, "-") {
|
||||
return errors.New("invalid stack name or repository URL")
|
||||
}
|
||||
if len(in.RepoURL) > 4096 || len(in.Branch) > 255 || strings.ContainsAny(in.Branch, "\r\n") {
|
||||
return errors.New("git source fields too long")
|
||||
}
|
||||
if in.Branch == "" {
|
||||
in.Branch = "main"
|
||||
}
|
||||
if in.Workdir == "." || in.Workdir == "" {
|
||||
in.Workdir = "."
|
||||
}
|
||||
if strings.HasPrefix(in.Workdir, "..") || filepath.IsAbs(in.Workdir) {
|
||||
return errors.New("invalid workdir")
|
||||
}
|
||||
if in.ComposeFile == "." || in.ComposeFile == "" {
|
||||
in.ComposeFile = "compose.yaml"
|
||||
}
|
||||
if strings.HasPrefix(in.ComposeFile, "..") || filepath.IsAbs(in.ComposeFile) {
|
||||
return errors.New("invalid compose_file")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) List(ctx context.Context) ([]Source, error) {
|
||||
rows, e := s.db.QueryContext(ctx, `SELECT id,node_id,stack_name,repo_url,branch,workdir,compose_file,auto_deploy,last_commit,last_sync_at,last_error,created_at,updated_at FROM git_sources ORDER BY stack_name`)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Source{}
|
||||
for rows.Next() {
|
||||
var x Source
|
||||
var sync, node sql.NullInt64
|
||||
if e := rows.Scan(&x.ID, &node, &x.StackName, &x.RepoURL, &x.Branch, &x.Workdir, &x.ComposeFile, &x.AutoDeploy, &x.LastCommit, &sync, &x.LastError, &x.CreatedAt, &x.UpdatedAt); e != nil {
|
||||
return nil, e
|
||||
}
|
||||
if node.Valid {
|
||||
v := node.Int64
|
||||
x.NodeID = &v
|
||||
}
|
||||
if sync.Valid {
|
||||
v := sync.Int64
|
||||
x.LastSyncAt = &v
|
||||
}
|
||||
out = append(out, x)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func (s *Service) Get(ctx context.Context, id int64) (Source, error) {
|
||||
var x Source
|
||||
var sync, node sql.NullInt64
|
||||
e := s.db.QueryRowContext(ctx, `SELECT id,node_id,stack_name,repo_url,branch,workdir,compose_file,auto_deploy,last_commit,last_sync_at,last_error,created_at,updated_at FROM git_sources WHERE id=?`, id).Scan(&x.ID, &node, &x.StackName, &x.RepoURL, &x.Branch, &x.Workdir, &x.ComposeFile, &x.AutoDeploy, &x.LastCommit, &sync, &x.LastError, &x.CreatedAt, &x.UpdatedAt)
|
||||
if node.Valid {
|
||||
v := node.Int64
|
||||
x.NodeID = &v
|
||||
}
|
||||
if sync.Valid {
|
||||
v := sync.Int64
|
||||
x.LastSyncAt = &v
|
||||
}
|
||||
return x, e
|
||||
}
|
||||
func (s *Service) Create(ctx context.Context, in Input) (Source, string, error) {
|
||||
if e := normalize(&in); e != nil {
|
||||
return Source{}, "", e
|
||||
}
|
||||
secret := make([]byte, 32)
|
||||
if _, e := rand.Read(secret); e != nil {
|
||||
return Source{}, "", e
|
||||
}
|
||||
sec := hex.EncodeToString(secret)
|
||||
enc, e := s.encrypt([]byte(sec))
|
||||
if e != nil {
|
||||
return Source{}, "", e
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
r, e := s.db.ExecContext(ctx, `INSERT INTO git_sources(node_id,stack_name,repo_url,branch,workdir,compose_file,auto_deploy,webhook_secret_enc,last_commit,last_error,created_at,updated_at) VALUES(?,?,?,?,?,?,?,?,'','',?,?)`, in.NodeID, in.StackName, in.RepoURL, in.Branch, in.Workdir, in.ComposeFile, in.AutoDeploy, enc, now, now)
|
||||
if e != nil {
|
||||
return Source{}, "", e
|
||||
}
|
||||
id, _ := r.LastInsertId()
|
||||
x, e := s.Get(ctx, id)
|
||||
return x, sec, e
|
||||
}
|
||||
func (s *Service) Update(ctx context.Context, id int64, in Input) (Source, error) {
|
||||
if e := normalize(&in); e != nil {
|
||||
return Source{}, e
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
r, e := s.db.ExecContext(ctx, `UPDATE git_sources SET node_id=?,stack_name=?,repo_url=?,branch=?,workdir=?,compose_file=?,auto_deploy=?,updated_at=? WHERE id=?`, in.NodeID, in.StackName, in.RepoURL, in.Branch, in.Workdir, in.ComposeFile, in.AutoDeploy, now, id)
|
||||
if e != nil {
|
||||
return Source{}, e
|
||||
}
|
||||
n, _ := r.RowsAffected()
|
||||
if n == 0 {
|
||||
return Source{}, sql.ErrNoRows
|
||||
}
|
||||
return s.Get(ctx, id)
|
||||
}
|
||||
func (s *Service) Delete(ctx context.Context, id int64) error {
|
||||
_, e := s.db.ExecContext(ctx, `DELETE FROM git_sources WHERE id=?`, id)
|
||||
return e
|
||||
}
|
||||
func (s *Service) RotateSecret(ctx context.Context, id int64) (string, error) {
|
||||
secret := make([]byte, 32)
|
||||
if _, e := rand.Read(secret); e != nil {
|
||||
return "", e
|
||||
}
|
||||
sec := hex.EncodeToString(secret)
|
||||
enc, e := s.encrypt([]byte(sec))
|
||||
if e != nil {
|
||||
return "", e
|
||||
}
|
||||
_, e = s.db.ExecContext(ctx, `UPDATE git_sources SET webhook_secret_enc=?,updated_at=? WHERE id=?`, enc, time.Now().Unix(), id)
|
||||
return sec, e
|
||||
}
|
||||
func (s *Service) VerifyWebhook(ctx context.Context, id int64, body []byte, signature, token string) error {
|
||||
var enc []byte
|
||||
if e := s.db.QueryRowContext(ctx, `SELECT webhook_secret_enc FROM git_sources WHERE id=?`, id).Scan(&enc); e != nil {
|
||||
return e
|
||||
}
|
||||
plain, e := s.decrypt(enc)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
secret := string(plain)
|
||||
if token != "" && hmac.Equal([]byte(token), []byte(secret)) {
|
||||
return nil
|
||||
}
|
||||
signature = strings.TrimPrefix(signature, "sha256=")
|
||||
if signature != "" {
|
||||
mac := hmac.New(sha256.New, []byte(secret))
|
||||
_, _ = mac.Write(body)
|
||||
want := hex.EncodeToString(mac.Sum(nil))
|
||||
if hmac.Equal([]byte(strings.ToLower(signature)), []byte(want)) {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return errors.New("invalid webhook signature")
|
||||
}
|
||||
func (s *Service) lockFor(id int64) *sync.Mutex {
|
||||
v, _ := s.locks.LoadOrStore(id, &sync.Mutex{})
|
||||
return v.(*sync.Mutex)
|
||||
}
|
||||
|
||||
func (s *Service) Sync(ctx context.Context, id int64) (Source, error) {
|
||||
mu := s.lockFor(id)
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
x, e := s.Get(ctx, id)
|
||||
if e != nil {
|
||||
return x, e
|
||||
}
|
||||
in := Input{NodeID: x.NodeID, StackName: x.StackName, RepoURL: x.RepoURL, Branch: x.Branch, Workdir: x.Workdir, ComposeFile: x.ComposeFile, AutoDeploy: x.AutoDeploy}
|
||||
var commit string
|
||||
if x.NodeID != nil {
|
||||
if s.nodes == nil {
|
||||
return s.syncFailed(ctx, x, errors.New("node manager unavailable"))
|
||||
}
|
||||
b, _, err := s.nodes.Do(ctx, *x.NodeID, "POST", "/agent/v1/git/sync", in)
|
||||
if err != nil {
|
||||
return s.syncFailed(ctx, x, err)
|
||||
}
|
||||
var resp struct {
|
||||
Commit string `json:"commit"`
|
||||
}
|
||||
if err = json.Unmarshal(b, &resp); err != nil {
|
||||
return s.syncFailed(ctx, x, err)
|
||||
}
|
||||
commit = resp.Commit
|
||||
} else {
|
||||
commit, e = s.SyncTransient(ctx, in)
|
||||
if e != nil {
|
||||
return s.syncFailed(ctx, x, e)
|
||||
}
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
_, _ = s.db.ExecContext(ctx, `UPDATE git_sources SET last_commit=?,last_sync_at=?,last_error='',updated_at=? WHERE id=?`, commit, now, now, id)
|
||||
return s.Get(ctx, id)
|
||||
}
|
||||
|
||||
// SyncTransient performs a Git-backed stack synchronization on the current
|
||||
// Docker environment. Agents expose this operation to the master without
|
||||
// persisting Git source metadata locally.
|
||||
func (s *Service) SyncTransient(ctx context.Context, in Input) (string, error) {
|
||||
if e := normalize(&in); e != nil {
|
||||
return "", e
|
||||
}
|
||||
tmp, e := os.MkdirTemp("", "dockwatch-git-*")
|
||||
if e != nil {
|
||||
return "", e
|
||||
}
|
||||
defer os.RemoveAll(tmp)
|
||||
cloneDir := filepath.Join(tmp, "repo")
|
||||
cctx, cancel := context.WithTimeout(ctx, 3*time.Minute)
|
||||
defer cancel()
|
||||
cmd := exec.CommandContext(cctx, "git", "clone", "--depth", "1", "--branch", in.Branch, "--single-branch", in.RepoURL, cloneDir)
|
||||
out, e := cmd.CombinedOutput()
|
||||
if e != nil {
|
||||
return "", fmt.Errorf("git clone: %w: %s", e, strings.TrimSpace(string(out)))
|
||||
}
|
||||
commitRaw, e := exec.CommandContext(cctx, "git", "-C", cloneDir, "rev-parse", "HEAD").Output()
|
||||
if e != nil {
|
||||
return "", e
|
||||
}
|
||||
src := filepath.Join(cloneDir, in.Workdir)
|
||||
if fi, e := os.Stat(src); e != nil || !fi.IsDir() {
|
||||
return "", errors.New("git workdir not found")
|
||||
}
|
||||
if _, e := os.Stat(filepath.Join(src, in.ComposeFile)); e != nil {
|
||||
return "", fmt.Errorf("compose file not found: %s", in.ComposeFile)
|
||||
}
|
||||
dst := filepath.Join(s.stacks.Root(), in.StackName)
|
||||
stage := filepath.Join(tmp, "stage")
|
||||
if e := copyDir(src, stage); e != nil {
|
||||
return "", e
|
||||
}
|
||||
if filepath.Clean(in.ComposeFile) != "compose.yaml" {
|
||||
b, e := os.ReadFile(filepath.Join(stage, in.ComposeFile))
|
||||
if e != nil {
|
||||
return "", e
|
||||
}
|
||||
if e = os.WriteFile(filepath.Join(stage, "compose.yaml"), b, 0640); e != nil {
|
||||
return "", e
|
||||
}
|
||||
}
|
||||
if e := s.stacks.ValidateProject(ctx, in.StackName, stage, "compose.yaml"); e != nil {
|
||||
return "", e
|
||||
}
|
||||
if e := syncManagedTree(stage, dst); e != nil {
|
||||
return "", e
|
||||
}
|
||||
if in.AutoDeploy {
|
||||
if _, e = s.stacks.Action(ctx, in.StackName, "up"); e != nil {
|
||||
return "", e
|
||||
}
|
||||
}
|
||||
return strings.TrimSpace(string(commitRaw)), nil
|
||||
}
|
||||
|
||||
func (s *Service) syncFailed(ctx context.Context, x Source, e error) (Source, error) {
|
||||
_, _ = s.db.ExecContext(ctx, `UPDATE git_sources SET last_error=?,updated_at=? WHERE id=?`, e.Error(), time.Now().Unix(), x.ID)
|
||||
x.LastError = e.Error()
|
||||
return x, e
|
||||
}
|
||||
|
||||
const gitManifestPath = ".dockwatch/git-manifest.json"
|
||||
|
||||
type gitManifest struct {
|
||||
Files []string `json:"files"`
|
||||
}
|
||||
|
||||
// syncManagedTree applies a Git checkout without treating the stack directory as
|
||||
// disposable storage. Only files previously managed by Git and files present in
|
||||
// the new checkout are changed. Unrelated files (for example bind-mount data)
|
||||
// survive a sync. The touched files are snapshotted so a failed apply can be
|
||||
// rolled back without copying or deleting the whole stack directory.
|
||||
func syncManagedTree(stage, dst string) error {
|
||||
files, err := collectManagedFiles(stage)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureSafeRoot(dst); err != nil {
|
||||
return err
|
||||
}
|
||||
old, err := readGitManifest(dst)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
backup, err := os.MkdirTemp("", "dockwatch-git-rollback-*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer os.RemoveAll(backup)
|
||||
|
||||
affected := map[string]struct{}{}
|
||||
for _, rel := range old {
|
||||
affected[rel] = struct{}{}
|
||||
}
|
||||
for _, rel := range files {
|
||||
affected[rel] = struct{}{}
|
||||
}
|
||||
manifestAbs := filepath.Join(dst, filepath.FromSlash(gitManifestPath))
|
||||
manifestBackup := filepath.Join(backup, "manifest.json")
|
||||
manifestExisted := false
|
||||
if info, err := os.Lstat(manifestAbs); err == nil {
|
||||
if !info.Mode().IsRegular() {
|
||||
return errors.New("Git manifest path is not a regular file")
|
||||
}
|
||||
if err := copyRegularFile(manifestAbs, manifestBackup, info.Mode().Perm()); err != nil {
|
||||
return err
|
||||
}
|
||||
manifestExisted = true
|
||||
} else if !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
|
||||
backedUp := map[string]bool{}
|
||||
for rel := range affected {
|
||||
target, err := safeManagedPath(dst, rel)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureSafeParent(dst, filepath.Dir(target)); err != nil {
|
||||
return err
|
||||
}
|
||||
info, err := os.Lstat(target)
|
||||
if os.IsNotExist(err) {
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
return fmt.Errorf("refusing to replace symlink in Git-managed path %q", rel)
|
||||
}
|
||||
if info.Mode().IsRegular() {
|
||||
bp := filepath.Join(backup, filepath.FromSlash(rel))
|
||||
if err := copyRegularFile(target, bp, info.Mode().Perm()); err != nil {
|
||||
return err
|
||||
}
|
||||
backedUp[rel] = true
|
||||
}
|
||||
}
|
||||
|
||||
rollback := func() {
|
||||
for rel := range affected {
|
||||
target, err := safeManagedPath(dst, rel)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if info, err := os.Lstat(target); err == nil && info.Mode().IsRegular() {
|
||||
_ = os.Remove(target)
|
||||
}
|
||||
if backedUp[rel] {
|
||||
bp := filepath.Join(backup, filepath.FromSlash(rel))
|
||||
if info, err := os.Stat(bp); err == nil {
|
||||
_ = copyRegularFile(bp, target, info.Mode().Perm())
|
||||
}
|
||||
}
|
||||
}
|
||||
_ = os.Remove(manifestAbs)
|
||||
if manifestExisted {
|
||||
if info, err := os.Stat(manifestBackup); err == nil {
|
||||
_ = copyRegularFile(manifestBackup, manifestAbs, info.Mode().Perm())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
newSet := make(map[string]struct{}, len(files))
|
||||
for _, rel := range files {
|
||||
newSet[rel] = struct{}{}
|
||||
}
|
||||
// Remove files that disappeared from Git, but never recursively remove a
|
||||
// directory. This deliberately leaves unrelated data untouched.
|
||||
for _, rel := range old {
|
||||
if _, ok := newSet[rel]; ok {
|
||||
continue
|
||||
}
|
||||
target, err := safeManagedPath(dst, rel)
|
||||
if err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
if info, err := os.Lstat(target); err == nil {
|
||||
if info.Mode()&os.ModeSymlink != 0 || (!info.Mode().IsRegular() && !info.IsDir()) {
|
||||
rollback()
|
||||
return fmt.Errorf("refusing to remove non-regular Git-managed path %q", rel)
|
||||
}
|
||||
if info.Mode().IsRegular() {
|
||||
if err := os.Remove(target); err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
}
|
||||
} else if !os.IsNotExist(err) {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
for _, rel := range files {
|
||||
source, err := safeManagedPath(stage, rel)
|
||||
if err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
target, err := safeManagedPath(dst, rel)
|
||||
if err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
if err := ensureSafeParent(dst, filepath.Dir(target)); err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
if info, err := os.Lstat(target); err == nil && info.IsDir() {
|
||||
// A file replacing a directory is safe only when that directory is
|
||||
// empty. os.Remove intentionally refuses non-empty directories.
|
||||
if err := os.Remove(target); err != nil {
|
||||
rollback()
|
||||
return fmt.Errorf("Git file %q conflicts with existing directory containing unmanaged data: %w", rel, err)
|
||||
}
|
||||
} else if err == nil && info.Mode()&os.ModeSymlink != 0 {
|
||||
rollback()
|
||||
return fmt.Errorf("refusing to replace symlink in Git-managed path %q", rel)
|
||||
} else if err != nil && !os.IsNotExist(err) {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
info, err := os.Stat(source)
|
||||
if err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
if err := copyRegularFile(source, target, info.Mode().Perm()); err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if err := writeGitManifest(dst, files); err != nil {
|
||||
rollback()
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func collectManagedFiles(root string) ([]string, error) {
|
||||
out := []string{}
|
||||
err := filepath.Walk(root, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if path == root {
|
||||
return nil
|
||||
}
|
||||
rel, err := filepath.Rel(root, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rel = filepath.ToSlash(rel)
|
||||
if rel == gitManifestPath {
|
||||
return nil
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
return fmt.Errorf("Git stack contains unsupported symlink %q", rel)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return nil
|
||||
}
|
||||
if info.Mode().IsRegular() {
|
||||
out = append(out, rel)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func readGitManifest(dst string) ([]string, error) {
|
||||
path := filepath.Join(dst, filepath.FromSlash(gitManifestPath))
|
||||
if err := ensureSafeParent(dst, filepath.Dir(path)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
b, err := os.ReadFile(path)
|
||||
if os.IsNotExist(err) {
|
||||
return []string{}, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var m gitManifest
|
||||
if err := json.Unmarshal(b, &m); err != nil {
|
||||
return nil, fmt.Errorf("invalid Git-managed file manifest: %w", err)
|
||||
}
|
||||
out := make([]string, 0, len(m.Files))
|
||||
seen := map[string]bool{}
|
||||
for _, rel := range m.Files {
|
||||
rel = filepath.ToSlash(filepath.Clean(filepath.FromSlash(rel)))
|
||||
if rel == "." || rel == gitManifestPath || strings.HasPrefix(rel, "../") || filepath.IsAbs(filepath.FromSlash(rel)) || seen[rel] {
|
||||
continue
|
||||
}
|
||||
seen[rel] = true
|
||||
out = append(out, rel)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func writeGitManifest(dst string, files []string) error {
|
||||
path := filepath.Join(dst, filepath.FromSlash(gitManifestPath))
|
||||
if err := ensureSafeParent(dst, filepath.Dir(path)); err != nil {
|
||||
return err
|
||||
}
|
||||
b, err := json.MarshalIndent(gitManifest{Files: files}, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
f, err := os.CreateTemp(filepath.Dir(path), ".dockwatch-git-manifest-*")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmp := f.Name()
|
||||
defer os.Remove(tmp)
|
||||
if err := f.Chmod(0640); err != nil {
|
||||
_ = f.Close()
|
||||
return err
|
||||
}
|
||||
if _, err := f.Write(append(b, '\n')); err != nil {
|
||||
_ = f.Close()
|
||||
return err
|
||||
}
|
||||
if err := f.Sync(); err != nil {
|
||||
_ = f.Close()
|
||||
return err
|
||||
}
|
||||
if err := f.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tmp, path)
|
||||
}
|
||||
|
||||
func safeManagedPath(root, rel string) (string, error) {
|
||||
rel = filepath.Clean(filepath.FromSlash(rel))
|
||||
if rel == "." || filepath.IsAbs(rel) || rel == ".." || strings.HasPrefix(rel, ".."+string(os.PathSeparator)) {
|
||||
return "", errors.New("invalid Git-managed path")
|
||||
}
|
||||
return filepath.Join(root, rel), nil
|
||||
}
|
||||
|
||||
func ensureSafeRoot(root string) error {
|
||||
info, err := os.Lstat(root)
|
||||
if os.IsNotExist(err) {
|
||||
return os.MkdirAll(root, 0750)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
|
||||
return errors.New("Git stack destination must be a real directory")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ensureSafeParent(root, parent string) error {
|
||||
rel, err := filepath.Rel(root, parent)
|
||||
if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(os.PathSeparator)) {
|
||||
return errors.New("path escapes stack root")
|
||||
}
|
||||
cur := root
|
||||
if err := ensureSafeRoot(root); err != nil {
|
||||
return err
|
||||
}
|
||||
if rel == "." {
|
||||
return nil
|
||||
}
|
||||
for _, part := range strings.Split(rel, string(os.PathSeparator)) {
|
||||
cur = filepath.Join(cur, part)
|
||||
info, err := os.Lstat(cur)
|
||||
if os.IsNotExist(err) {
|
||||
if err := os.Mkdir(cur, 0750); err != nil && !os.IsExist(err) {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
|
||||
return fmt.Errorf("unsafe parent path %q", cur)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func copyRegularFile(src, dst string, mode os.FileMode) error {
|
||||
if err := os.MkdirAll(filepath.Dir(dst), 0750); err != nil {
|
||||
return err
|
||||
}
|
||||
in, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
tmp, err := os.CreateTemp(filepath.Dir(dst), ".dockwatch-git-write-*")
|
||||
if err != nil {
|
||||
_ = in.Close()
|
||||
return err
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
defer os.Remove(tmpName)
|
||||
if err := tmp.Chmod(mode); err != nil {
|
||||
_ = in.Close()
|
||||
_ = tmp.Close()
|
||||
return err
|
||||
}
|
||||
_, copyErr := io.Copy(tmp, in)
|
||||
inErr := in.Close()
|
||||
if copyErr == nil {
|
||||
copyErr = tmp.Sync()
|
||||
}
|
||||
outErr := tmp.Close()
|
||||
if copyErr != nil {
|
||||
return copyErr
|
||||
}
|
||||
if inErr != nil {
|
||||
return inErr
|
||||
}
|
||||
if outErr != nil {
|
||||
return outErr
|
||||
}
|
||||
return os.Rename(tmpName, dst)
|
||||
}
|
||||
|
||||
func copyDir(src, dst string) error {
|
||||
return filepath.Walk(src, func(path string, info os.FileInfo, e error) error {
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
rel, e := filepath.Rel(src, path)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if rel == ".git" || strings.HasPrefix(rel, ".git"+string(os.PathSeparator)) {
|
||||
if info.IsDir() {
|
||||
return filepath.SkipDir
|
||||
}
|
||||
return nil
|
||||
}
|
||||
target := filepath.Join(dst, rel)
|
||||
if info.IsDir() {
|
||||
return os.MkdirAll(target, info.Mode().Perm())
|
||||
}
|
||||
if !info.Mode().IsRegular() {
|
||||
return nil
|
||||
}
|
||||
if e := os.MkdirAll(filepath.Dir(target), 0750); e != nil {
|
||||
return e
|
||||
}
|
||||
in, e := os.Open(path)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
out, e := os.OpenFile(target, os.O_CREATE|os.O_TRUNC|os.O_WRONLY, info.Mode().Perm())
|
||||
if e != nil {
|
||||
_ = in.Close()
|
||||
return e
|
||||
}
|
||||
_, copyErr := io.Copy(out, in)
|
||||
inErr := in.Close()
|
||||
outErr := out.Close()
|
||||
if copyErr != nil {
|
||||
return copyErr
|
||||
}
|
||||
if inErr != nil {
|
||||
return inErr
|
||||
}
|
||||
return outErr
|
||||
})
|
||||
}
|
||||
func (s *Service) encrypt(p []byte) ([]byte, error) {
|
||||
b, e := aes.NewCipher(s.key)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
g, e := cipher.NewGCM(b)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
nonce := make([]byte, g.NonceSize())
|
||||
if _, e = rand.Read(nonce); e != nil {
|
||||
return nil, e
|
||||
}
|
||||
return g.Seal(nonce, nonce, p, nil), nil
|
||||
}
|
||||
func (s *Service) decrypt(v []byte) ([]byte, error) {
|
||||
b, e := aes.NewCipher(s.key)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
g, e := cipher.NewGCM(b)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
if len(v) < g.NonceSize() {
|
||||
return nil, errors.New("invalid encrypted secret")
|
||||
}
|
||||
return g.Open(nil, v[:g.NonceSize()], v[g.NonceSize():], nil)
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package gitops
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestNormalizeRejectsUnsafeStackAndRepositoryArguments(t *testing.T) {
|
||||
cases := []Input{
|
||||
{StackName: "../escape", RepoURL: "https://example.invalid/repo.git"},
|
||||
{StackName: "demo", RepoURL: "--upload-pack=evil"},
|
||||
{StackName: "demo", RepoURL: "https://example.invalid/repo.git\n--option"},
|
||||
}
|
||||
for _, in := range cases {
|
||||
if err := normalize(&in); err == nil {
|
||||
t.Fatalf("expected unsafe input to be rejected: %#v", in)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncManagedTreePreservesUnmanagedDataAndRemovesStaleGitFiles(t *testing.T) {
|
||||
dst := t.TempDir()
|
||||
stage1 := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(stage1, "compose.yaml"), []byte("services: {}\n"), 0640); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Join(stage1, "config"), 0750); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(stage1, "config", "old.txt"), []byte("old"), 0640); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := syncManagedTree(stage1, dst); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Join(dst, "data"), 0750); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dst, "data", "runtime.db"), []byte("keep"), 0600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
stage2 := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(stage2, "compose.yaml"), []byte("services:\n web:\n image: nginx\n"), 0640); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Join(stage2, "config"), 0750); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(stage2, "config", "new.txt"), []byte("new"), 0640); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := syncManagedTree(stage2, dst); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if b, err := os.ReadFile(filepath.Join(dst, "data", "runtime.db")); err != nil || string(b) != "keep" {
|
||||
t.Fatalf("unmanaged data was not preserved: %q %v", b, err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(dst, "config", "old.txt")); !os.IsNotExist(err) {
|
||||
t.Fatalf("stale Git-managed file still exists: %v", err)
|
||||
}
|
||||
if b, err := os.ReadFile(filepath.Join(dst, "config", "new.txt")); err != nil || string(b) != "new" {
|
||||
t.Fatalf("new Git file missing: %q %v", b, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncManagedTreeRejectsSymlinkDestination(t *testing.T) {
|
||||
parent := t.TempDir()
|
||||
outside := t.TempDir()
|
||||
dst := filepath.Join(parent, "demo")
|
||||
if err := os.Symlink(outside, dst); err != nil {
|
||||
t.Skipf("symlinks unavailable: %v", err)
|
||||
}
|
||||
stage := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(stage, "compose.yaml"), []byte("services: {}\n"), 0640); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := syncManagedTree(stage, dst); err == nil {
|
||||
t.Fatal("expected symlink Git destination to be rejected")
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,18 @@
|
||||
package httpapi
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/audit"
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/auth"
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/config"
|
||||
)
|
||||
|
||||
func TestRouterPatternsDoNotConflict(t *testing.T) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
t.Fatalf("ServeMux route conflict: %v", r)
|
||||
}
|
||||
}()
|
||||
_ = New(config.Config{Mode: config.ModeStandalone}, &auth.Service{}, nil, nil, nil, (*audit.Service)(nil), nil, nil)
|
||||
}
|
||||
@@ -0,0 +1,481 @@
|
||||
package monitor
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type ProbeGroup struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
Monitors []Monitor `json:"monitors,omitempty"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
UpdatedAt int64 `json:"updated_at"`
|
||||
}
|
||||
|
||||
type ProbeGroupInput struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
MonitorIDs []int64 `json:"monitor_ids"`
|
||||
}
|
||||
|
||||
type StatusPage struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Slug string `json:"slug"`
|
||||
Description string `json:"description"`
|
||||
Status string `json:"status"`
|
||||
Enabled bool `json:"enabled"`
|
||||
ServiceIDs []int64 `json:"service_ids"`
|
||||
Services []ProbeGroup `json:"services,omitempty"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
UpdatedAt int64 `json:"updated_at"`
|
||||
}
|
||||
|
||||
type StatusPageInput struct {
|
||||
Name string `json:"name"`
|
||||
Slug string `json:"slug"`
|
||||
Description string `json:"description"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
ServiceIDs []int64 `json:"service_ids"`
|
||||
}
|
||||
|
||||
var slugRx = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,62}$`)
|
||||
|
||||
func aggregateStatus(ms []Monitor) string {
|
||||
if len(ms) == 0 {
|
||||
return "unknown"
|
||||
}
|
||||
hasMaint, hasPending, hasPaused := false, false, false
|
||||
for _, m := range ms {
|
||||
switch m.Status {
|
||||
case "down":
|
||||
return "down"
|
||||
case "maintenance":
|
||||
hasMaint = true
|
||||
case "pending":
|
||||
hasPending = true
|
||||
case "paused":
|
||||
hasPaused = true
|
||||
}
|
||||
}
|
||||
if hasMaint {
|
||||
return "maintenance"
|
||||
}
|
||||
if hasPending {
|
||||
return "pending"
|
||||
}
|
||||
if hasPaused {
|
||||
return "paused"
|
||||
}
|
||||
return "up"
|
||||
}
|
||||
|
||||
func (s *Service) ListGroups(ctx context.Context) ([]ProbeGroup, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `SELECT id,name,description,created_at,updated_at FROM monitor_services ORDER BY name`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out := []ProbeGroup{}
|
||||
for rows.Next() {
|
||||
var g ProbeGroup
|
||||
if err := rows.Scan(&g.ID, &g.Name, &g.Description, &g.CreatedAt, &g.UpdatedAt); err != nil {
|
||||
_ = rows.Close()
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, g)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
_ = rows.Close()
|
||||
return nil, err
|
||||
}
|
||||
if err := rows.Close(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Load monitors once and group in memory. This avoids an N+1 query pattern
|
||||
// on the dashboard and keeps refresh cost predictable with many services.
|
||||
all, err := s.List(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
byService := make(map[int64][]Monitor, len(out))
|
||||
for _, m := range all {
|
||||
if m.ServiceID != nil {
|
||||
byService[*m.ServiceID] = append(byService[*m.ServiceID], m)
|
||||
}
|
||||
}
|
||||
for i := range out {
|
||||
out[i].Monitors = byService[out[i].ID]
|
||||
if out[i].Monitors == nil {
|
||||
out[i].Monitors = []Monitor{}
|
||||
}
|
||||
out[i].Status = aggregateStatus(out[i].Monitors)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (s *Service) GetGroup(ctx context.Context, id int64) (ProbeGroup, error) {
|
||||
var g ProbeGroup
|
||||
if err := s.db.QueryRowContext(ctx, `SELECT id,name,description,created_at,updated_at FROM monitor_services WHERE id=?`, id).Scan(&g.ID, &g.Name, &g.Description, &g.CreatedAt, &g.UpdatedAt); err != nil {
|
||||
return g, err
|
||||
}
|
||||
ms, err := s.listGroupMonitors(ctx, id)
|
||||
if err != nil {
|
||||
return g, err
|
||||
}
|
||||
g.Monitors = ms
|
||||
g.Status = aggregateStatus(ms)
|
||||
return g, nil
|
||||
}
|
||||
|
||||
func (s *Service) listGroupMonitors(ctx context.Context, id int64) ([]Monitor, error) {
|
||||
rows, err := s.db.QueryContext(ctx, selectMonitor+` WHERE m.service_id=? ORDER BY m.name`, time.Now().Add(-24*time.Hour).Unix(), id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Monitor{}
|
||||
for rows.Next() {
|
||||
m, e := scanMonitor(rows)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
out = append(out, m)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func normalizeGroupInput(in *ProbeGroupInput) error {
|
||||
in.Name = strings.TrimSpace(in.Name)
|
||||
in.Description = strings.TrimSpace(in.Description)
|
||||
if in.Name == "" || len(in.Name) > 120 || strings.ContainsAny(in.Name, "\r\n") {
|
||||
return errors.New("valid service name required")
|
||||
}
|
||||
if len(in.Description) > 2000 {
|
||||
return errors.New("service description too long")
|
||||
}
|
||||
if in.MonitorIDs != nil {
|
||||
seen := map[int64]bool{}
|
||||
ids := make([]int64, 0, len(in.MonitorIDs))
|
||||
for _, id := range in.MonitorIDs {
|
||||
if id < 1 {
|
||||
return errors.New("monitor_ids must contain positive IDs")
|
||||
}
|
||||
if !seen[id] {
|
||||
seen[id] = true
|
||||
ids = append(ids, id)
|
||||
}
|
||||
}
|
||||
in.MonitorIDs = ids
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func assignGroupMonitors(ctx context.Context, tx *sql.Tx, groupID int64, monitorIDs []int64, clearExisting bool) error {
|
||||
if clearExisting {
|
||||
if _, err := tx.ExecContext(ctx, `UPDATE monitors SET service_id=NULL,updated_at=? WHERE service_id=?`, time.Now().Unix(), groupID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, monitorID := range monitorIDs {
|
||||
r, err := tx.ExecContext(ctx, `UPDATE monitors SET service_id=?,updated_at=? WHERE id=?`, groupID, time.Now().Unix(), monitorID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, _ := r.RowsAffected()
|
||||
if n == 0 {
|
||||
return fmt.Errorf("monitor %d not found", monitorID)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) CreateGroup(ctx context.Context, in ProbeGroupInput) (ProbeGroup, error) {
|
||||
if err := normalizeGroupInput(&in); err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
now := time.Now().Unix()
|
||||
r, err := tx.ExecContext(ctx, `INSERT INTO monitor_services(name,description,created_at,updated_at) VALUES(?,?,?,?)`, in.Name, in.Description, now, now)
|
||||
if err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
id, _ := r.LastInsertId()
|
||||
if in.MonitorIDs != nil {
|
||||
if err := assignGroupMonitors(ctx, tx, id, in.MonitorIDs, false); err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
return s.GetGroup(ctx, id)
|
||||
}
|
||||
func (s *Service) UpdateGroup(ctx context.Context, id int64, in ProbeGroupInput) (ProbeGroup, error) {
|
||||
if err := normalizeGroupInput(&in); err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
r, err := tx.ExecContext(ctx, `UPDATE monitor_services SET name=?,description=?,updated_at=? WHERE id=?`, in.Name, in.Description, time.Now().Unix(), id)
|
||||
if err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
n, _ := r.RowsAffected()
|
||||
if n == 0 {
|
||||
return ProbeGroup{}, sql.ErrNoRows
|
||||
}
|
||||
if in.MonitorIDs != nil {
|
||||
if err := assignGroupMonitors(ctx, tx, id, in.MonitorIDs, true); err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return ProbeGroup{}, err
|
||||
}
|
||||
return s.GetGroup(ctx, id)
|
||||
}
|
||||
func (s *Service) DeleteGroup(ctx context.Context, id int64) error {
|
||||
_, e := s.db.ExecContext(ctx, `DELETE FROM monitor_services WHERE id=?`, id)
|
||||
return e
|
||||
}
|
||||
|
||||
func normalizeStatusPage(in *StatusPageInput) error {
|
||||
in.Name = strings.TrimSpace(in.Name)
|
||||
in.Slug = strings.ToLower(strings.TrimSpace(in.Slug))
|
||||
in.Description = strings.TrimSpace(in.Description)
|
||||
if in.Name == "" || len(in.Name) > 120 || strings.ContainsAny(in.Name, "\r\n") || !slugRx.MatchString(in.Slug) {
|
||||
return errors.New("name and valid slug required (lowercase letters, numbers, hyphens)")
|
||||
}
|
||||
if len(in.Description) > 4000 {
|
||||
return errors.New("status page description too long")
|
||||
}
|
||||
seen := map[int64]bool{}
|
||||
var ids []int64
|
||||
for _, id := range in.ServiceIDs {
|
||||
if id > 0 && !seen[id] {
|
||||
seen[id] = true
|
||||
ids = append(ids, id)
|
||||
}
|
||||
}
|
||||
in.ServiceIDs = ids
|
||||
return nil
|
||||
}
|
||||
func (s *Service) ListStatusPages(ctx context.Context) ([]StatusPage, error) {
|
||||
rows, e := s.db.QueryContext(ctx, `SELECT id,name,slug,description,enabled,created_at,updated_at FROM status_pages ORDER BY name`)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
out := []StatusPage{}
|
||||
for rows.Next() {
|
||||
var p StatusPage
|
||||
if e = rows.Scan(&p.ID, &p.Name, &p.Slug, &p.Description, &p.Enabled, &p.CreatedAt, &p.UpdatedAt); e != nil {
|
||||
_ = rows.Close()
|
||||
return nil, e
|
||||
}
|
||||
p.ServiceIDs = []int64{}
|
||||
out = append(out, p)
|
||||
}
|
||||
if e = rows.Err(); e != nil {
|
||||
_ = rows.Close()
|
||||
return nil, e
|
||||
}
|
||||
if e = rows.Close(); e != nil {
|
||||
return nil, e
|
||||
}
|
||||
|
||||
// Resolve page memberships only after releasing the one SQLite connection.
|
||||
for i := range out {
|
||||
out[i].ServiceIDs, e = s.statusPageServiceIDs(ctx, out[i].ID)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
func (s *Service) GetStatusPage(ctx context.Context, id int64) (StatusPage, error) {
|
||||
var p StatusPage
|
||||
e := s.db.QueryRowContext(ctx, `SELECT id,name,slug,description,enabled,created_at,updated_at FROM status_pages WHERE id=?`, id).Scan(&p.ID, &p.Name, &p.Slug, &p.Description, &p.Enabled, &p.CreatedAt, &p.UpdatedAt)
|
||||
if e != nil {
|
||||
return p, e
|
||||
}
|
||||
p.ServiceIDs, e = s.statusPageServiceIDs(ctx, id)
|
||||
return p, e
|
||||
}
|
||||
func (s *Service) statusPageServiceIDs(ctx context.Context, id int64) ([]int64, error) {
|
||||
rows, e := s.db.QueryContext(ctx, `SELECT service_id FROM status_page_services WHERE page_id=? ORDER BY sort_order,service_id`, id)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []int64{}
|
||||
for rows.Next() {
|
||||
var x int64
|
||||
if e = rows.Scan(&x); e != nil {
|
||||
return nil, e
|
||||
}
|
||||
out = append(out, x)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func (s *Service) savePageServices(ctx context.Context, tx *sql.Tx, id int64, ids []int64) error {
|
||||
if _, e := tx.ExecContext(ctx, `DELETE FROM status_page_services WHERE page_id=?`, id); e != nil {
|
||||
return e
|
||||
}
|
||||
for i, sid := range ids {
|
||||
if _, e := tx.ExecContext(ctx, `INSERT INTO status_page_services(page_id,service_id,sort_order) VALUES(?,?,?)`, id, sid, i); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) CreateStatusPage(ctx context.Context, in StatusPageInput) (StatusPage, error) {
|
||||
if e := normalizeStatusPage(&in); e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
en := true
|
||||
if in.Enabled != nil {
|
||||
en = *in.Enabled
|
||||
}
|
||||
tx, e := s.db.BeginTx(ctx, nil)
|
||||
if e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
defer tx.Rollback()
|
||||
now := time.Now().Unix()
|
||||
r, e := tx.ExecContext(ctx, `INSERT INTO status_pages(name,slug,description,enabled,created_at,updated_at) VALUES(?,?,?,?,?,?)`, in.Name, in.Slug, in.Description, en, now, now)
|
||||
if e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
id, _ := r.LastInsertId()
|
||||
if e = s.savePageServices(ctx, tx, id, in.ServiceIDs); e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
if e = tx.Commit(); e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
return s.GetStatusPage(ctx, id)
|
||||
}
|
||||
func (s *Service) UpdateStatusPage(ctx context.Context, id int64, in StatusPageInput) (StatusPage, error) {
|
||||
if e := normalizeStatusPage(&in); e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
old, e := s.GetStatusPage(ctx, id)
|
||||
if e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
en := old.Enabled
|
||||
if in.Enabled != nil {
|
||||
en = *in.Enabled
|
||||
}
|
||||
tx, e := s.db.BeginTx(ctx, nil)
|
||||
if e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
defer tx.Rollback()
|
||||
if _, e = tx.ExecContext(ctx, `UPDATE status_pages SET name=?,slug=?,description=?,enabled=?,updated_at=? WHERE id=?`, in.Name, in.Slug, in.Description, en, time.Now().Unix(), id); e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
if e = s.savePageServices(ctx, tx, id, in.ServiceIDs); e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
if e = tx.Commit(); e != nil {
|
||||
return StatusPage{}, e
|
||||
}
|
||||
return s.GetStatusPage(ctx, id)
|
||||
}
|
||||
func (s *Service) DeleteStatusPage(ctx context.Context, id int64) error {
|
||||
_, e := s.db.ExecContext(ctx, `DELETE FROM status_pages WHERE id=?`, id)
|
||||
return e
|
||||
}
|
||||
func (s *Service) PublicStatusPage(ctx context.Context, slug string) (StatusPage, error) {
|
||||
p := StatusPage{Services: []ProbeGroup{}, ServiceIDs: []int64{}}
|
||||
e := s.db.QueryRowContext(ctx, `SELECT id,name,slug,description,enabled,created_at,updated_at FROM status_pages WHERE slug=? AND enabled=1`, slug).Scan(&p.ID, &p.Name, &p.Slug, &p.Description, &p.Enabled, &p.CreatedAt, &p.UpdatedAt)
|
||||
if e != nil {
|
||||
return p, e
|
||||
}
|
||||
rows, e := s.db.QueryContext(ctx, `SELECT ms.id,ms.name,ms.description,ms.created_at,ms.updated_at FROM monitor_services ms JOIN status_page_services ps ON ps.service_id=ms.id WHERE ps.page_id=? ORDER BY ps.sort_order,ms.name`, p.ID)
|
||||
if e != nil {
|
||||
return p, e
|
||||
}
|
||||
for rows.Next() {
|
||||
var g ProbeGroup
|
||||
if e = rows.Scan(&g.ID, &g.Name, &g.Description, &g.CreatedAt, &g.UpdatedAt); e != nil {
|
||||
_ = rows.Close()
|
||||
return p, e
|
||||
}
|
||||
g.Monitors = []Monitor{}
|
||||
p.Services = append(p.Services, g)
|
||||
p.ServiceIDs = append(p.ServiceIDs, g.ID)
|
||||
}
|
||||
if e = rows.Err(); e != nil {
|
||||
_ = rows.Close()
|
||||
return p, e
|
||||
}
|
||||
if e = rows.Close(); e != nil {
|
||||
return p, e
|
||||
}
|
||||
|
||||
all, e := s.List(ctx)
|
||||
if e != nil {
|
||||
return p, e
|
||||
}
|
||||
byService := map[int64][]Monitor{}
|
||||
for _, m := range all {
|
||||
if m.ServiceID != nil {
|
||||
byService[*m.ServiceID] = append(byService[*m.ServiceID], m)
|
||||
}
|
||||
}
|
||||
for gi := range p.Services {
|
||||
p.Services[gi].Monitors = byService[p.Services[gi].ID]
|
||||
if p.Services[gi].Monitors == nil {
|
||||
p.Services[gi].Monitors = []Monitor{}
|
||||
}
|
||||
p.Services[gi].Status = aggregateStatus(p.Services[gi].Monitors)
|
||||
}
|
||||
p.Status = aggregateServiceStatuses(p.Services)
|
||||
return p, nil
|
||||
}
|
||||
|
||||
func aggregateServiceStatuses(groups []ProbeGroup) string {
|
||||
if len(groups) == 0 {
|
||||
return "unknown"
|
||||
}
|
||||
hasMaint, hasPending, hasPaused := false, false, false
|
||||
for _, g := range groups {
|
||||
switch g.Status {
|
||||
case "down":
|
||||
return "down"
|
||||
case "maintenance":
|
||||
hasMaint = true
|
||||
case "pending", "unknown":
|
||||
hasPending = true
|
||||
case "paused":
|
||||
hasPaused = true
|
||||
}
|
||||
}
|
||||
if hasMaint {
|
||||
return "maintenance"
|
||||
}
|
||||
if hasPending {
|
||||
return "pending"
|
||||
}
|
||||
if hasPaused {
|
||||
return "paused"
|
||||
}
|
||||
return "up"
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package monitor
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"errors"
|
||||
"io"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
var registerSingleConnDriver sync.Once
|
||||
|
||||
func openSingleConnRegressionDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
registerSingleConnDriver.Do(func() { sql.Register("dockwatch-singleconn-regression", singleConnDriver{}) })
|
||||
db, err := sql.Open("dockwatch-singleconn-regression", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db.SetMaxOpenConns(1)
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
return db
|
||||
}
|
||||
|
||||
type singleConnDriver struct{}
|
||||
|
||||
func (singleConnDriver) Open(string) (driver.Conn, error) { return &singleConn{}, nil }
|
||||
|
||||
type singleConn struct{}
|
||||
|
||||
func (*singleConn) Prepare(string) (driver.Stmt, error) {
|
||||
return nil, errors.New("prepare not supported")
|
||||
}
|
||||
func (*singleConn) Close() error { return nil }
|
||||
func (*singleConn) Begin() (driver.Tx, error) { return nil, errors.New("tx not supported") }
|
||||
func (*singleConn) QueryContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Rows, error) {
|
||||
switch {
|
||||
case strings.Contains(query, "FROM monitor_services ORDER BY name"):
|
||||
return newStaticRows([]string{"id", "name", "description", "created_at", "updated_at"}, [][]driver.Value{{int64(1), "Website", "", int64(1), int64(1)}}), nil
|
||||
case strings.Contains(query, "FROM status_pages ORDER BY name"):
|
||||
return newStaticRows([]string{"id", "name", "slug", "description", "enabled", "created_at", "updated_at"}, [][]driver.Value{{int64(1), "Public", "public", "", int64(1), int64(1), int64(1)}}), nil
|
||||
case strings.Contains(query, "FROM status_pages WHERE slug="):
|
||||
return newStaticRows([]string{"id", "name", "slug", "description", "enabled", "created_at", "updated_at"}, [][]driver.Value{{int64(1), "Public", "public", "", int64(1), int64(1), int64(1)}}), nil
|
||||
case strings.Contains(query, "FROM status_page_services WHERE page_id="):
|
||||
return newStaticRows([]string{"service_id"}, [][]driver.Value{{int64(1)}}), nil
|
||||
case strings.Contains(query, "JOIN status_page_services"):
|
||||
return newStaticRows([]string{"id", "name", "description", "created_at", "updated_at"}, [][]driver.Value{{int64(1), "Website", "", int64(1), int64(1)}}), nil
|
||||
case strings.Contains(query, "FROM monitors m"):
|
||||
// The regression only needs to verify that this second query can start
|
||||
// after the service rows were released. No monitor row is required.
|
||||
return newStaticRows([]string{"monitor"}, nil), nil
|
||||
default:
|
||||
return nil, errors.New("unexpected query: " + query)
|
||||
}
|
||||
}
|
||||
|
||||
var _ driver.QueryerContext = (*singleConn)(nil)
|
||||
|
||||
type staticRows struct {
|
||||
cols []string
|
||||
data [][]driver.Value
|
||||
pos int
|
||||
}
|
||||
|
||||
func newStaticRows(cols []string, data [][]driver.Value) *staticRows {
|
||||
return &staticRows{cols: cols, data: data}
|
||||
}
|
||||
func (r *staticRows) Columns() []string { return r.cols }
|
||||
func (r *staticRows) Close() error { return nil }
|
||||
func (r *staticRows) Next(dest []driver.Value) error {
|
||||
if r.pos >= len(r.data) {
|
||||
return io.EOF
|
||||
}
|
||||
copy(dest, r.data[r.pos])
|
||||
r.pos++
|
||||
return nil
|
||||
}
|
||||
|
||||
func shortContext(t *testing.T) context.Context {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
|
||||
t.Cleanup(cancel)
|
||||
return ctx
|
||||
}
|
||||
|
||||
func TestListGroupsDoesNotNestQueriesOnSingleConnection(t *testing.T) {
|
||||
s := &Service{db: openSingleConnRegressionDB(t)}
|
||||
groups, err := s.ListGroups(shortContext(t))
|
||||
if err != nil {
|
||||
t.Fatalf("ListGroups: %v", err)
|
||||
}
|
||||
if len(groups) != 1 || groups[0].Name != "Website" {
|
||||
t.Fatalf("unexpected groups: %#v", groups)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListStatusPagesDoesNotNestQueriesOnSingleConnection(t *testing.T) {
|
||||
s := &Service{db: openSingleConnRegressionDB(t)}
|
||||
pages, err := s.ListStatusPages(shortContext(t))
|
||||
if err != nil {
|
||||
t.Fatalf("ListStatusPages: %v", err)
|
||||
}
|
||||
if len(pages) != 1 || len(pages[0].ServiceIDs) != 1 || pages[0].ServiceIDs[0] != 1 {
|
||||
t.Fatalf("unexpected pages: %#v", pages)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPublicStatusPageDoesNotNestQueriesOnSingleConnection(t *testing.T) {
|
||||
s := &Service{db: openSingleConnRegressionDB(t)}
|
||||
page, err := s.PublicStatusPage(shortContext(t), "public")
|
||||
if err != nil {
|
||||
t.Fatalf("PublicStatusPage: %v", err)
|
||||
}
|
||||
if len(page.Services) != 1 || page.Services[0].Name != "Website" {
|
||||
t.Fatalf("unexpected public page: %#v", page)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,659 @@
|
||||
package monitor
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os/exec"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/buildinfo"
|
||||
"git.send.nrw/sendnrw/dockwatch/internal/nodes"
|
||||
)
|
||||
|
||||
type Monitor struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Target string `json:"target"`
|
||||
NodeID *int64 `json:"node_id,omitempty"`
|
||||
ServiceID *int64 `json:"service_id,omitempty"`
|
||||
IntervalSeconds int `json:"interval_seconds"`
|
||||
TimeoutMS int `json:"timeout_ms"`
|
||||
ExpectedMin int `json:"expected_min"`
|
||||
ExpectedMax int `json:"expected_max"`
|
||||
Method string `json:"method"`
|
||||
HeadersJSON string `json:"headers_json"`
|
||||
Body string `json:"body"`
|
||||
Keyword string `json:"keyword"`
|
||||
InvertKeyword bool `json:"invert_keyword"`
|
||||
IgnoreTLS bool `json:"ignore_tls"`
|
||||
RequireHealthy bool `json:"require_healthy"`
|
||||
Enabled bool `json:"enabled"`
|
||||
Status string `json:"status"`
|
||||
MaintenanceUntil *int64 `json:"maintenance_until,omitempty"`
|
||||
MaintenanceNote string `json:"maintenance_note"`
|
||||
LastCheckedAt *int64 `json:"last_checked_at,omitempty"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
UpdatedAt int64 `json:"updated_at"`
|
||||
Uptime24h float64 `json:"uptime_24h"`
|
||||
LastLatencyMS int64 `json:"last_latency_ms"`
|
||||
LastMessage string `json:"last_message"`
|
||||
LastStatusCode int `json:"last_status_code"`
|
||||
}
|
||||
|
||||
type Input struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Target string `json:"target"`
|
||||
NodeID *int64 `json:"node_id"`
|
||||
ServiceID *int64 `json:"service_id"`
|
||||
IntervalSeconds int `json:"interval_seconds"`
|
||||
TimeoutMS int `json:"timeout_ms"`
|
||||
ExpectedMin int `json:"expected_min"`
|
||||
ExpectedMax int `json:"expected_max"`
|
||||
Method string `json:"method"`
|
||||
HeadersJSON string `json:"headers_json"`
|
||||
Body string `json:"body"`
|
||||
Keyword string `json:"keyword"`
|
||||
InvertKeyword bool `json:"invert_keyword"`
|
||||
IgnoreTLS bool `json:"ignore_tls"`
|
||||
RequireHealthy bool `json:"require_healthy"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type MaintenanceInput struct {
|
||||
Until *int64 `json:"until"`
|
||||
Note string `json:"note"`
|
||||
}
|
||||
type Check struct {
|
||||
ID int64 `json:"id,omitempty"`
|
||||
MonitorID int64 `json:"monitor_id,omitempty"`
|
||||
OK bool `json:"ok"`
|
||||
StatusCode int `json:"status_code"`
|
||||
LatencyMS int64 `json:"latency_ms"`
|
||||
Message string `json:"message"`
|
||||
CheckedAt int64 `json:"checked_at"`
|
||||
}
|
||||
type Event struct {
|
||||
MonitorID int64 `json:"monitor_id"`
|
||||
Name string `json:"name"`
|
||||
Target string `json:"target"`
|
||||
From string `json:"from"`
|
||||
To string `json:"to"`
|
||||
Check Check `json:"check"`
|
||||
}
|
||||
|
||||
type Service struct {
|
||||
db *sql.DB
|
||||
nodes *nodes.Manager
|
||||
workers chan struct{}
|
||||
mu sync.Mutex
|
||||
running map[int64]bool
|
||||
retentionDays int
|
||||
eventSink func(context.Context, Event)
|
||||
}
|
||||
|
||||
func New(db *sql.DB, nm *nodes.Manager, c, r int) *Service {
|
||||
return &Service{db: db, nodes: nm, workers: make(chan struct{}, c), running: map[int64]bool{}, retentionDays: r}
|
||||
}
|
||||
|
||||
func (s *Service) SetEventSink(fn func(context.Context, Event)) { s.eventSink = fn }
|
||||
|
||||
const selectMonitor = `WITH stats AS (
|
||||
SELECT monitor_id,100.0*AVG(ok) AS uptime_24h FROM monitor_checks WHERE checked_at>=? GROUP BY monitor_id
|
||||
), latest AS (
|
||||
SELECT monitor_id,MAX(id) AS id FROM monitor_checks GROUP BY monitor_id
|
||||
)
|
||||
SELECT m.id,m.name,m.type,m.target,m.node_id,m.service_id,m.interval_seconds,m.timeout_ms,m.expected_min,m.expected_max,m.method,m.headers_json,m.body,m.keyword,m.invert_keyword,m.ignore_tls,m.require_healthy,m.enabled,m.status,m.maintenance_until,m.maintenance_note,m.last_checked_at,m.created_at,m.updated_at,
|
||||
COALESCE(stats.uptime_24h,0),COALESCE(c.latency_ms,0),COALESCE(c.message,''),COALESCE(c.status_code,0)
|
||||
FROM monitors m
|
||||
LEFT JOIN stats ON stats.monitor_id=m.id
|
||||
LEFT JOIN latest ON latest.monitor_id=m.id
|
||||
LEFT JOIN monitor_checks c ON c.id=latest.id`
|
||||
|
||||
func scanMonitor(sc interface{ Scan(...any) error }) (Monitor, error) {
|
||||
var m Monitor
|
||||
var node, serviceID, last, maint sql.NullInt64
|
||||
err := sc.Scan(&m.ID, &m.Name, &m.Type, &m.Target, &node, &serviceID, &m.IntervalSeconds, &m.TimeoutMS, &m.ExpectedMin, &m.ExpectedMax, &m.Method, &m.HeadersJSON, &m.Body, &m.Keyword, &m.InvertKeyword, &m.IgnoreTLS, &m.RequireHealthy, &m.Enabled, &m.Status, &maint, &m.MaintenanceNote, &last, &m.CreatedAt, &m.UpdatedAt, &m.Uptime24h, &m.LastLatencyMS, &m.LastMessage, &m.LastStatusCode)
|
||||
if node.Valid {
|
||||
m.NodeID = &node.Int64
|
||||
}
|
||||
if serviceID.Valid {
|
||||
m.ServiceID = &serviceID.Int64
|
||||
}
|
||||
if last.Valid {
|
||||
m.LastCheckedAt = &last.Int64
|
||||
}
|
||||
if maint.Valid {
|
||||
m.MaintenanceUntil = &maint.Int64
|
||||
}
|
||||
return m, err
|
||||
}
|
||||
func (s *Service) List(ctx context.Context) ([]Monitor, error) {
|
||||
rows, e := s.db.QueryContext(ctx, selectMonitor+` ORDER BY m.name`, time.Now().Add(-24*time.Hour).Unix())
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Monitor{}
|
||||
for rows.Next() {
|
||||
m, e := scanMonitor(rows)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
out = append(out, m)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func (s *Service) Get(ctx context.Context, id int64) (Monitor, error) {
|
||||
return scanMonitor(s.db.QueryRowContext(ctx, selectMonitor+` WHERE m.id=?`, time.Now().Add(-24*time.Hour).Unix(), id))
|
||||
}
|
||||
|
||||
const selectMonitorSchedule = `SELECT m.id,m.name,m.type,m.target,m.node_id,m.service_id,m.interval_seconds,m.timeout_ms,m.expected_min,m.expected_max,m.method,m.headers_json,m.body,m.keyword,m.invert_keyword,m.ignore_tls,m.require_healthy,m.enabled,m.status,m.maintenance_until,m.maintenance_note,m.last_checked_at,m.created_at,m.updated_at,0.0,0,'',0 FROM monitors m ORDER BY m.id`
|
||||
|
||||
func (s *Service) listForSchedule(ctx context.Context) ([]Monitor, error) {
|
||||
rows, err := s.db.QueryContext(ctx, selectMonitorSchedule)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Monitor{}
|
||||
for rows.Next() {
|
||||
m, err := scanMonitor(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, m)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func normalize(in *Input, requireName bool) error {
|
||||
in.Name = strings.TrimSpace(in.Name)
|
||||
in.Type = strings.ToLower(strings.TrimSpace(in.Type))
|
||||
in.Target = strings.TrimSpace(in.Target)
|
||||
if in.Target == "" {
|
||||
return errors.New("target required")
|
||||
}
|
||||
if requireName && in.Name == "" {
|
||||
return errors.New("name required")
|
||||
}
|
||||
if strings.ContainsAny(in.Name, "\r\n") {
|
||||
return errors.New("invalid monitor name")
|
||||
}
|
||||
if len(in.Name) > 200 || len(in.Target) > 4096 {
|
||||
return errors.New("name or target too long")
|
||||
}
|
||||
if in.Type != "http" && in.Type != "tcp" && in.Type != "dns" && in.Type != "docker" {
|
||||
return errors.New("type must be http, tcp, dns or docker")
|
||||
}
|
||||
if in.NodeID != nil && *in.NodeID < 1 {
|
||||
return errors.New("node_id must be positive")
|
||||
}
|
||||
if in.ServiceID != nil && *in.ServiceID < 1 {
|
||||
return errors.New("service_id must be positive")
|
||||
}
|
||||
if in.IntervalSeconds == 0 {
|
||||
in.IntervalSeconds = 60
|
||||
}
|
||||
if in.IntervalSeconds < 10 || in.IntervalSeconds > 86400 {
|
||||
return errors.New("interval_seconds must be 10..86400")
|
||||
}
|
||||
if in.TimeoutMS == 0 {
|
||||
in.TimeoutMS = 5000
|
||||
}
|
||||
if in.TimeoutMS < 100 || in.TimeoutMS > 60000 {
|
||||
return errors.New("timeout_ms must be 100..60000")
|
||||
}
|
||||
if len(in.HeadersJSON) > 64<<10 || len(in.Body) > 256<<10 || len(in.Keyword) > 4096 {
|
||||
return errors.New("monitor headers/body/keyword too large")
|
||||
}
|
||||
in.Method = strings.ToUpper(strings.TrimSpace(in.Method))
|
||||
if in.Method == "" {
|
||||
in.Method = "GET"
|
||||
}
|
||||
if in.HeadersJSON == "" {
|
||||
in.HeadersJSON = "{}"
|
||||
}
|
||||
var h map[string]string
|
||||
if err := json.Unmarshal([]byte(in.HeadersJSON), &h); err != nil {
|
||||
return errors.New("headers_json must be a JSON object with string values")
|
||||
}
|
||||
for k, v := range h {
|
||||
if strings.TrimSpace(k) == "" || strings.ContainsAny(k, "\r\n") || strings.ContainsAny(v, "\r\n") {
|
||||
return errors.New("HTTP headers must not contain empty names or newlines")
|
||||
}
|
||||
}
|
||||
switch in.Type {
|
||||
case "http":
|
||||
u, err := url.Parse(in.Target)
|
||||
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.User != nil {
|
||||
return errors.New("HTTP target must be an absolute http(s) URL without embedded credentials")
|
||||
}
|
||||
allowed := map[string]bool{"GET": true, "HEAD": true, "POST": true, "PUT": true, "PATCH": true, "DELETE": true, "OPTIONS": true}
|
||||
if !allowed[in.Method] {
|
||||
return errors.New("unsupported HTTP method")
|
||||
}
|
||||
if in.ExpectedMin == 0 {
|
||||
in.ExpectedMin = 200
|
||||
}
|
||||
if in.ExpectedMax == 0 {
|
||||
in.ExpectedMax = 399
|
||||
}
|
||||
if in.ExpectedMin < 100 || in.ExpectedMax > 599 || in.ExpectedMin > in.ExpectedMax {
|
||||
return errors.New("expected HTTP status range must be within 100..599")
|
||||
}
|
||||
case "tcp":
|
||||
host, port, err := net.SplitHostPort(in.Target)
|
||||
if err != nil || strings.TrimSpace(host) == "" || strings.TrimSpace(port) == "" {
|
||||
return errors.New("TCP target must be host:port")
|
||||
}
|
||||
if p, err := strconv.Atoi(port); err != nil || p < 1 || p > 65535 {
|
||||
return errors.New("TCP target port must be 1..65535")
|
||||
}
|
||||
case "dns":
|
||||
if strings.ContainsAny(in.Target, " /\\") {
|
||||
return errors.New("DNS target must be a hostname or IP address")
|
||||
}
|
||||
case "docker":
|
||||
if len(in.Target) > 255 || strings.ContainsAny(in.Target, "\r\n") || strings.HasPrefix(in.Target, "-") {
|
||||
return errors.New("invalid Docker container target")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) Create(ctx context.Context, in Input, userID int64) (Monitor, error) {
|
||||
if e := normalize(&in, true); e != nil {
|
||||
return Monitor{}, e
|
||||
}
|
||||
enabled := true
|
||||
if in.Enabled != nil {
|
||||
enabled = *in.Enabled
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
res, e := s.db.ExecContext(ctx, `INSERT INTO monitors(name,type,target,node_id,service_id,interval_seconds,timeout_ms,expected_min,expected_max,method,headers_json,body,keyword,invert_keyword,ignore_tls,require_healthy,enabled,status,created_by,created_at,updated_at) VALUES(?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,'pending',?,?,?)`, in.Name, in.Type, in.Target, in.NodeID, in.ServiceID, in.IntervalSeconds, in.TimeoutMS, in.ExpectedMin, in.ExpectedMax, in.Method, in.HeadersJSON, in.Body, in.Keyword, in.InvertKeyword, in.IgnoreTLS, in.RequireHealthy, enabled, userID, now, now)
|
||||
if e != nil {
|
||||
return Monitor{}, e
|
||||
}
|
||||
id, _ := res.LastInsertId()
|
||||
return s.Get(ctx, id)
|
||||
}
|
||||
func (s *Service) Update(ctx context.Context, id int64, in Input) (Monitor, error) {
|
||||
if e := normalize(&in, true); e != nil {
|
||||
return Monitor{}, e
|
||||
}
|
||||
enabled := true
|
||||
if in.Enabled != nil {
|
||||
enabled = *in.Enabled
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
res, e := s.db.ExecContext(ctx, `UPDATE monitors SET name=?,type=?,target=?,node_id=?,service_id=?,interval_seconds=?,timeout_ms=?,expected_min=?,expected_max=?,method=?,headers_json=?,body=?,keyword=?,invert_keyword=?,ignore_tls=?,require_healthy=?,enabled=?,updated_at=? WHERE id=?`, in.Name, in.Type, in.Target, in.NodeID, in.ServiceID, in.IntervalSeconds, in.TimeoutMS, in.ExpectedMin, in.ExpectedMax, in.Method, in.HeadersJSON, in.Body, in.Keyword, in.InvertKeyword, in.IgnoreTLS, in.RequireHealthy, enabled, now, id)
|
||||
if e != nil {
|
||||
return Monitor{}, e
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return Monitor{}, sql.ErrNoRows
|
||||
}
|
||||
if !enabled {
|
||||
_, _ = s.db.ExecContext(ctx, `UPDATE monitors SET status='paused' WHERE id=?`, id)
|
||||
} else {
|
||||
_, _ = s.db.ExecContext(ctx, `UPDATE monitors SET status=CASE WHEN status='paused' THEN 'pending' ELSE status END WHERE id=?`, id)
|
||||
}
|
||||
return s.Get(ctx, id)
|
||||
}
|
||||
func (s *Service) SetPaused(ctx context.Context, id int64, paused bool) error {
|
||||
enabled := !paused
|
||||
status := "pending"
|
||||
if paused {
|
||||
status = "paused"
|
||||
}
|
||||
res, e := s.db.ExecContext(ctx, `UPDATE monitors SET enabled=?,status=?,updated_at=? WHERE id=?`, enabled, status, time.Now().Unix(), id)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return sql.ErrNoRows
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) SetMaintenance(ctx context.Context, id int64, in MaintenanceInput) error {
|
||||
now := time.Now().Unix()
|
||||
if in.Until != nil && *in.Until <= now {
|
||||
return errors.New("maintenance end must be in the future")
|
||||
}
|
||||
in.Note = strings.TrimSpace(in.Note)
|
||||
if len(in.Note) > 2000 {
|
||||
return errors.New("maintenance note too long")
|
||||
}
|
||||
res, e := s.db.ExecContext(ctx, `UPDATE monitors SET maintenance_until=?,maintenance_note=?,status='maintenance',updated_at=? WHERE id=?`, in.Until, in.Note, now, id)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return sql.ErrNoRows
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) ClearMaintenance(ctx context.Context, id int64) error {
|
||||
r, e := s.db.ExecContext(ctx, `UPDATE monitors SET maintenance_until=NULL,maintenance_note='',status=CASE WHEN enabled=1 THEN 'pending' ELSE 'paused' END,updated_at=? WHERE id=?`, time.Now().Unix(), id)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
n, _ := r.RowsAffected()
|
||||
if n == 0 {
|
||||
return sql.ErrNoRows
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) Delete(ctx context.Context, id int64) error {
|
||||
_, e := s.db.ExecContext(ctx, `DELETE FROM monitors WHERE id=?`, id)
|
||||
return e
|
||||
}
|
||||
func (s *Service) Checks(ctx context.Context, id int64, limit int) ([]Check, error) {
|
||||
if limit < 1 {
|
||||
limit = 100
|
||||
}
|
||||
if limit > 1000 {
|
||||
limit = 1000
|
||||
}
|
||||
rows, e := s.db.QueryContext(ctx, `SELECT id,monitor_id,ok,status_code,latency_ms,message,checked_at FROM monitor_checks WHERE monitor_id=? ORDER BY checked_at DESC LIMIT ?`, id, limit)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Check{}
|
||||
for rows.Next() {
|
||||
var c Check
|
||||
if e := rows.Scan(&c.ID, &c.MonitorID, &c.OK, &c.StatusCode, &c.LatencyMS, &c.Message, &c.CheckedAt); e != nil {
|
||||
return nil, e
|
||||
}
|
||||
out = append(out, c)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func (s *Service) Run(ctx context.Context) {
|
||||
tick := time.NewTicker(2 * time.Second)
|
||||
cleanup := time.NewTicker(6 * time.Hour)
|
||||
defer tick.Stop()
|
||||
defer cleanup.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-tick.C:
|
||||
s.schedule(ctx)
|
||||
case <-cleanup.C:
|
||||
s.cleanup(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
func (s *Service) schedule(ctx context.Context) {
|
||||
ms, e := s.listForSchedule(ctx)
|
||||
if e != nil {
|
||||
return
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
for _, m := range ms {
|
||||
if !m.Enabled {
|
||||
continue
|
||||
}
|
||||
if m.Status == "maintenance" && m.MaintenanceUntil == nil {
|
||||
continue
|
||||
}
|
||||
if m.MaintenanceUntil != nil {
|
||||
if *m.MaintenanceUntil == 0 || *m.MaintenanceUntil > now {
|
||||
if m.Status != "maintenance" {
|
||||
_, _ = s.db.ExecContext(ctx, `UPDATE monitors SET status='maintenance' WHERE id=?`, m.ID)
|
||||
}
|
||||
continue
|
||||
}
|
||||
_ = s.ClearMaintenance(ctx, m.ID)
|
||||
}
|
||||
if m.LastCheckedAt != nil && now-*m.LastCheckedAt < int64(m.IntervalSeconds) {
|
||||
continue
|
||||
}
|
||||
s.mu.Lock()
|
||||
if s.running[m.ID] {
|
||||
s.mu.Unlock()
|
||||
continue
|
||||
}
|
||||
s.running[m.ID] = true
|
||||
s.mu.Unlock()
|
||||
go func(mon Monitor) {
|
||||
select {
|
||||
case s.workers <- struct{}{}:
|
||||
case <-ctx.Done():
|
||||
s.mu.Lock()
|
||||
delete(s.running, mon.ID)
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
defer func() { <-s.workers; s.mu.Lock(); delete(s.running, mon.ID); s.mu.Unlock() }()
|
||||
c := s.probe(ctx, mon)
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
x, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
_ = s.recordCheck(x, mon, c)
|
||||
}(m)
|
||||
}
|
||||
}
|
||||
func (s *Service) CheckNow(ctx context.Context, id int64) (Check, error) {
|
||||
s.mu.Lock()
|
||||
if s.running[id] {
|
||||
s.mu.Unlock()
|
||||
return Check{}, errors.New("monitor check already running")
|
||||
}
|
||||
s.running[id] = true
|
||||
s.mu.Unlock()
|
||||
defer func() {
|
||||
s.mu.Lock()
|
||||
delete(s.running, id)
|
||||
s.mu.Unlock()
|
||||
}()
|
||||
m, err := s.Get(ctx, id)
|
||||
if err != nil {
|
||||
return Check{}, err
|
||||
}
|
||||
c := s.probe(ctx, m)
|
||||
if err := ctx.Err(); err != nil {
|
||||
return c, err
|
||||
}
|
||||
// A manual diagnostic check must not implicitly resume a paused monitor or
|
||||
// end maintenance. We still keep the check in history, but preserve the
|
||||
// lifecycle state until the user explicitly changes it.
|
||||
updateState := m.Enabled && m.Status != "maintenance"
|
||||
if err = s.recordCheckResult(ctx, m, c, updateState); err != nil {
|
||||
return c, err
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
func (s *Service) recordCheck(ctx context.Context, m Monitor, c Check) error {
|
||||
return s.recordCheckResult(ctx, m, c, true)
|
||||
}
|
||||
|
||||
func (s *Service) recordCheckResult(ctx context.Context, m Monitor, c Check, updateState bool) error {
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
if _, err = tx.ExecContext(ctx, `INSERT INTO monitor_checks(monitor_id,ok,status_code,latency_ms,message,checked_at) VALUES(?,?,?,?,?,?)`, m.ID, c.OK, c.StatusCode, c.LatencyMS, c.Message, c.CheckedAt); err != nil {
|
||||
return err
|
||||
}
|
||||
status := "down"
|
||||
if c.OK {
|
||||
status = "up"
|
||||
}
|
||||
if updateState {
|
||||
if _, err = tx.ExecContext(ctx, `UPDATE monitors SET status=?,last_checked_at=?,updated_at=? WHERE id=?`, status, c.CheckedAt, c.CheckedAt, m.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
} else {
|
||||
if _, err = tx.ExecContext(ctx, `UPDATE monitors SET last_checked_at=?,updated_at=? WHERE id=?`, c.CheckedAt, c.CheckedAt, m.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err = tx.Commit(); err != nil {
|
||||
return err
|
||||
}
|
||||
if updateState && s.eventSink != nil && (m.Status == "up" || m.Status == "down") && m.Status != status {
|
||||
s.eventSink(context.Background(), Event{MonitorID: m.ID, Name: m.Name, Target: m.Target, From: m.Status, To: status, Check: c})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) probe(ctx context.Context, m Monitor) Check {
|
||||
in := Input{Name: m.Name, Type: m.Type, Target: m.Target, TimeoutMS: m.TimeoutMS, ExpectedMin: m.ExpectedMin, ExpectedMax: m.ExpectedMax, Method: m.Method, HeadersJSON: m.HeadersJSON, Body: m.Body, Keyword: m.Keyword, InvertKeyword: m.InvertKeyword, IgnoreTLS: m.IgnoreTLS, RequireHealthy: m.RequireHealthy}
|
||||
if m.NodeID != nil {
|
||||
b, _, e := s.nodes.Do(ctx, *m.NodeID, http.MethodPost, "/agent/v1/probe", in)
|
||||
if e != nil {
|
||||
return Check{Message: e.Error(), CheckedAt: time.Now().Unix()}
|
||||
}
|
||||
var c Check
|
||||
if json.Unmarshal(b, &c) != nil {
|
||||
return Check{Message: "invalid agent response", CheckedAt: time.Now().Unix()}
|
||||
}
|
||||
return c
|
||||
}
|
||||
return Probe(ctx, in)
|
||||
}
|
||||
func Probe(ctx context.Context, in Input) Check {
|
||||
start := time.Now()
|
||||
c := Check{CheckedAt: start.Unix()}
|
||||
if err := normalize(&in, false); err != nil {
|
||||
c.Message = err.Error()
|
||||
return c
|
||||
}
|
||||
pctx, cancel := context.WithTimeout(ctx, time.Duration(in.TimeoutMS)*time.Millisecond)
|
||||
defer cancel()
|
||||
switch in.Type {
|
||||
case "http":
|
||||
tr := &http.Transport{TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12, InsecureSkipVerify: in.IgnoreTLS}}
|
||||
defer tr.CloseIdleConnections()
|
||||
client := &http.Client{Transport: tr}
|
||||
req, e := http.NewRequestWithContext(pctx, in.Method, in.Target, strings.NewReader(in.Body))
|
||||
if e != nil {
|
||||
c.Message = e.Error()
|
||||
return c
|
||||
}
|
||||
var headers map[string]string
|
||||
_ = json.Unmarshal([]byte(in.HeadersJSON), &headers)
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
req.Header.Set("User-Agent", "Dockwatch/"+buildinfo.Current().Version)
|
||||
resp, e := client.Do(req)
|
||||
if e != nil {
|
||||
c.Message = e.Error()
|
||||
return c
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, _ := io.ReadAll(io.LimitReader(resp.Body, 512<<10))
|
||||
c.StatusCode = resp.StatusCode
|
||||
c.OK = resp.StatusCode >= in.ExpectedMin && resp.StatusCode <= in.ExpectedMax
|
||||
if c.OK && in.Keyword != "" {
|
||||
found := strings.Contains(string(body), in.Keyword)
|
||||
if in.InvertKeyword {
|
||||
found = !found
|
||||
}
|
||||
c.OK = found
|
||||
if !found {
|
||||
c.Message = "keyword assertion failed"
|
||||
}
|
||||
}
|
||||
if !c.OK && c.Message == "" {
|
||||
c.Message = "unexpected HTTP status"
|
||||
}
|
||||
case "tcp":
|
||||
conn, e := (&net.Dialer{}).DialContext(pctx, "tcp", in.Target)
|
||||
if e != nil {
|
||||
c.Message = e.Error()
|
||||
return c
|
||||
}
|
||||
_ = conn.Close()
|
||||
c.OK = true
|
||||
case "docker":
|
||||
cmd := exec.CommandContext(pctx, "docker", "inspect", "--format", "{{json .State}}", in.Target)
|
||||
b, e := cmd.Output()
|
||||
if e != nil {
|
||||
c.Message = "docker inspect: " + e.Error()
|
||||
return c
|
||||
}
|
||||
var st struct {
|
||||
Running bool `json:"Running"`
|
||||
Status string `json:"Status"`
|
||||
Health *struct {
|
||||
Status string `json:"Status"`
|
||||
} `json:"Health"`
|
||||
}
|
||||
if e = json.Unmarshal(bytesTrimSpace(b), &st); e != nil {
|
||||
c.Message = "invalid docker state: " + e.Error()
|
||||
return c
|
||||
}
|
||||
if !st.Running {
|
||||
c.Message = "container is not running (" + st.Status + ")"
|
||||
return c
|
||||
}
|
||||
if in.RequireHealthy {
|
||||
if st.Health == nil {
|
||||
c.Message = "container has no healthcheck"
|
||||
return c
|
||||
}
|
||||
if strings.ToLower(st.Health.Status) != "healthy" {
|
||||
c.Message = "container health is " + st.Health.Status
|
||||
return c
|
||||
}
|
||||
}
|
||||
c.OK = true
|
||||
c.Message = "container running"
|
||||
if in.RequireHealthy {
|
||||
c.Message = "container running and healthy"
|
||||
}
|
||||
case "dns":
|
||||
_, e := net.DefaultResolver.LookupHost(pctx, in.Target)
|
||||
if e != nil {
|
||||
c.Message = e.Error()
|
||||
return c
|
||||
}
|
||||
c.OK = true
|
||||
default:
|
||||
c.Message = "unsupported monitor type"
|
||||
}
|
||||
c.LatencyMS = time.Since(start).Milliseconds()
|
||||
return c
|
||||
}
|
||||
func (s *Service) cleanup(ctx context.Context) {
|
||||
// Expired authentication sessions are always disposable, even when heartbeat
|
||||
// retention is configured as unlimited (0).
|
||||
_, _ = s.db.ExecContext(ctx, `DELETE FROM sessions WHERE expires_at<?`, time.Now().Unix())
|
||||
if s.retentionDays <= 0 {
|
||||
return
|
||||
}
|
||||
cut := time.Now().Add(-time.Duration(s.retentionDays) * 24 * time.Hour).Unix()
|
||||
_, _ = s.db.ExecContext(ctx, `DELETE FROM monitor_checks WHERE checked_at<?`, cut)
|
||||
}
|
||||
func ParseID(v string) (int64, error) {
|
||||
id, e := strconv.ParseInt(v, 10, 64)
|
||||
if e != nil || id < 1 {
|
||||
return 0, fmt.Errorf("invalid id")
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func bytesTrimSpace(b []byte) []byte { return []byte(strings.TrimSpace(string(b))) }
|
||||
@@ -0,0 +1,84 @@
|
||||
package monitor
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestHTTPProbe(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(http.StatusNoContent) }))
|
||||
defer srv.Close()
|
||||
c := Probe(context.Background(), Input{Type: "http", Target: srv.URL, TimeoutMS: 1000, ExpectedMin: 200, ExpectedMax: 299})
|
||||
if !c.OK || c.StatusCode != http.StatusNoContent {
|
||||
t.Fatalf("unexpected check: %+v", c)
|
||||
}
|
||||
}
|
||||
func TestTCPProbe(t *testing.T) {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer ln.Close()
|
||||
c := Probe(context.Background(), Input{Type: "tcp", Target: ln.Addr().String(), TimeoutMS: 1000})
|
||||
if !c.OK {
|
||||
t.Fatalf("unexpected check: %+v", c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAggregateServiceStatus(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
monitors []Monitor
|
||||
want string
|
||||
}{
|
||||
{"all up", []Monitor{{Status: "up"}, {Status: "up"}}, "up"},
|
||||
{"one fault", []Monitor{{Status: "up"}, {Status: "down"}}, "down"},
|
||||
{"maintenance", []Monitor{{Status: "up"}, {Status: "maintenance"}}, "maintenance"},
|
||||
{"empty", nil, "unknown"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := aggregateStatus(tc.monitors); got != tc.want {
|
||||
t.Fatalf("got %s want %s", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDockerProbeRunningAndHealthy(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("shell fixture")
|
||||
}
|
||||
d := t.TempDir()
|
||||
path := filepath.Join(d, "docker")
|
||||
if err := os.WriteFile(path, []byte("#!/bin/sh\nprintf '%s\\n' '{\"Running\":true,\"Status\":\"running\",\"Health\":{\"Status\":\"healthy\"}}'\n"), 0755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("PATH", d+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
c := Probe(context.Background(), Input{Type: "docker", Target: "app", TimeoutMS: 1000, RequireHealthy: true})
|
||||
if !c.OK {
|
||||
t.Fatalf("expected healthy docker probe, got %+v", c)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDockerProbeHealthRequiredWithoutHealthcheck(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("shell fixture")
|
||||
}
|
||||
d := t.TempDir()
|
||||
path := filepath.Join(d, "docker")
|
||||
if err := os.WriteFile(path, []byte("#!/bin/sh\nprintf '%s\\n' '{\"Running\":true,\"Status\":\"running\",\"Health\":null}'\n"), 0755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("PATH", d+string(os.PathListSeparator)+os.Getenv("PATH"))
|
||||
c := Probe(context.Background(), Input{Type: "docker", Target: "app", TimeoutMS: 1000, RequireHealthy: true})
|
||||
if c.OK || c.Message != "container has no healthcheck" {
|
||||
t.Fatalf("unexpected docker check: %+v", c)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,306 @@
|
||||
package nodes
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
type Node struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
BaseURL string `json:"base_url"`
|
||||
Enabled bool `json:"enabled"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
UpdatedAt int64 `json:"updated_at"`
|
||||
}
|
||||
type storedNode struct {
|
||||
Node
|
||||
Token string
|
||||
}
|
||||
type Manager struct {
|
||||
db *sql.DB
|
||||
key []byte
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
func New(db *sql.DB, key []byte) *Manager {
|
||||
return &Manager{db: db, key: key, client: &http.Client{Timeout: 30 * time.Second}}
|
||||
}
|
||||
func (m *Manager) List(ctx context.Context) ([]Node, error) {
|
||||
rows, err := m.db.QueryContext(ctx, `SELECT id,name,base_url,enabled,created_at,updated_at FROM nodes ORDER BY name`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Node{}
|
||||
for rows.Next() {
|
||||
var n Node
|
||||
if err := rows.Scan(&n.ID, &n.Name, &n.BaseURL, &n.Enabled, &n.CreatedAt, &n.UpdatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, n)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func validateBaseURL(baseURL string) error {
|
||||
u, err := url.Parse(baseURL)
|
||||
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" {
|
||||
return errors.New("base_url must be an absolute http(s) URL without credentials, query or fragment")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *Manager) Create(ctx context.Context, name, baseURL, token string) (Node, error) {
|
||||
name = strings.TrimSpace(name)
|
||||
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
|
||||
if name == "" || len(name) > 120 || strings.ContainsAny(name, "\r\n") || len(token) < 24 {
|
||||
return Node{}, errors.New("valid name and token (>=24 chars) required")
|
||||
}
|
||||
if err := validateBaseURL(baseURL); err != nil {
|
||||
return Node{}, err
|
||||
}
|
||||
enc, err := m.encrypt([]byte(token))
|
||||
if err != nil {
|
||||
return Node{}, err
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
res, err := m.db.ExecContext(ctx, `INSERT INTO nodes(name,base_url,token_enc,enabled,created_at,updated_at) VALUES(?,?,?,?,?,?)`, name, baseURL, enc, 1, now, now)
|
||||
if err != nil {
|
||||
return Node{}, err
|
||||
}
|
||||
id, _ := res.LastInsertId()
|
||||
return Node{ID: id, Name: name, BaseURL: baseURL, Enabled: true, CreatedAt: now, UpdatedAt: now}, nil
|
||||
}
|
||||
|
||||
func (m *Manager) Update(ctx context.Context, id int64, name, baseURL, token string, enabled *bool) (Node, error) {
|
||||
old, err := m.get(ctx, id)
|
||||
if err != nil {
|
||||
return Node{}, err
|
||||
}
|
||||
name = strings.TrimSpace(name)
|
||||
baseURL = strings.TrimRight(strings.TrimSpace(baseURL), "/")
|
||||
if name == "" {
|
||||
name = old.Name
|
||||
}
|
||||
if len(name) > 120 || strings.ContainsAny(name, "\r\n") {
|
||||
return Node{}, errors.New("invalid node name")
|
||||
}
|
||||
if baseURL == "" {
|
||||
baseURL = old.BaseURL
|
||||
}
|
||||
if err := validateBaseURL(baseURL); err != nil {
|
||||
return Node{}, err
|
||||
}
|
||||
enc := []byte(nil)
|
||||
if strings.TrimSpace(token) != "" {
|
||||
if len(token) < 24 {
|
||||
return Node{}, errors.New("token must be at least 24 characters")
|
||||
}
|
||||
enc, err = m.encrypt([]byte(token))
|
||||
if err != nil {
|
||||
return Node{}, err
|
||||
}
|
||||
}
|
||||
en := old.Enabled
|
||||
if enabled != nil {
|
||||
en = *enabled
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
var res sql.Result
|
||||
if enc != nil {
|
||||
res, err = m.db.ExecContext(ctx, `UPDATE nodes SET name=?,base_url=?,token_enc=?,enabled=?,updated_at=? WHERE id=?`, name, baseURL, enc, en, now, id)
|
||||
} else {
|
||||
res, err = m.db.ExecContext(ctx, `UPDATE nodes SET name=?,base_url=?,enabled=?,updated_at=? WHERE id=?`, name, baseURL, en, now, id)
|
||||
}
|
||||
if err != nil {
|
||||
return Node{}, err
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return Node{}, sql.ErrNoRows
|
||||
}
|
||||
return Node{ID: id, Name: name, BaseURL: baseURL, Enabled: en, CreatedAt: old.CreatedAt, UpdatedAt: now}, nil
|
||||
}
|
||||
func (m *Manager) Delete(ctx context.Context, id int64) error {
|
||||
_, err := m.db.ExecContext(ctx, `DELETE FROM nodes WHERE id=?`, id)
|
||||
return err
|
||||
}
|
||||
func (m *Manager) get(ctx context.Context, id int64) (storedNode, error) {
|
||||
var n storedNode
|
||||
var enc []byte
|
||||
err := m.db.QueryRowContext(ctx, `SELECT id,name,base_url,token_enc,enabled,created_at,updated_at FROM nodes WHERE id=?`, id).Scan(&n.ID, &n.Name, &n.BaseURL, &enc, &n.Enabled, &n.CreatedAt, &n.UpdatedAt)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
plain, err := m.decrypt(enc)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
n.Token = string(plain)
|
||||
return n, nil
|
||||
}
|
||||
func (m *Manager) Do(ctx context.Context, id int64, method, path string, body any) ([]byte, int, error) {
|
||||
n, err := m.get(ctx, id)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if !n.Enabled {
|
||||
return nil, 0, errors.New("node disabled")
|
||||
}
|
||||
var rdr io.Reader
|
||||
if body != nil {
|
||||
b, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
rdr = bytes.NewReader(b)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, n.BaseURL+path, rdr)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+n.Token)
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
resp, err := m.client.Do(req)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
b, err := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
|
||||
if err != nil {
|
||||
return nil, resp.StatusCode, err
|
||||
}
|
||||
if resp.StatusCode >= 300 {
|
||||
return b, resp.StatusCode, fmt.Errorf("agent returned %s: %s", resp.Status, strings.TrimSpace(string(b)))
|
||||
}
|
||||
return b, resp.StatusCode, nil
|
||||
}
|
||||
|
||||
func (m *Manager) Stream(ctx context.Context, id int64, method, path string, body any, w http.ResponseWriter) error {
|
||||
n, err := m.get(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !n.Enabled {
|
||||
return errors.New("node disabled")
|
||||
}
|
||||
var rdr io.Reader
|
||||
if body != nil {
|
||||
b, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rdr = bytes.NewReader(b)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, method, n.BaseURL+path, rdr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+n.Token)
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
client := *m.client
|
||||
client.Timeout = 0
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
for k, vs := range resp.Header {
|
||||
for _, v := range vs {
|
||||
w.Header().Add(k, v)
|
||||
}
|
||||
}
|
||||
w.WriteHeader(resp.StatusCode)
|
||||
_, err = io.Copy(w, resp.Body)
|
||||
return err
|
||||
}
|
||||
|
||||
func (m *Manager) encrypt(p []byte) ([]byte, error) {
|
||||
b, err := aes.NewCipher(m.key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
g, err := cipher.NewGCM(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nonce := make([]byte, g.NonceSize())
|
||||
if _, err := rand.Read(nonce); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return g.Seal(nonce, nonce, p, nil), nil
|
||||
}
|
||||
func (m *Manager) decrypt(v []byte) ([]byte, error) {
|
||||
b, err := aes.NewCipher(m.key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
g, err := cipher.NewGCM(b)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(v) < g.NonceSize() {
|
||||
return nil, errors.New("invalid encrypted token")
|
||||
}
|
||||
return g.Open(nil, v[:g.NonceSize()], v[g.NonceSize():], nil)
|
||||
}
|
||||
|
||||
// DialWebSocket opens an authenticated websocket to a remote agent.
|
||||
func (m *Manager) DialWebSocket(ctx context.Context, id int64, path string) (*websocket.Conn, *http.Response, error) {
|
||||
n, err := m.get(ctx, id)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if !n.Enabled {
|
||||
return nil, nil, errors.New("node disabled")
|
||||
}
|
||||
wsURL, err := agentWebSocketURL(n.BaseURL, path)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
h := http.Header{}
|
||||
h.Set("Authorization", "Bearer "+n.Token)
|
||||
d := websocket.Dialer{HandshakeTimeout: 15 * time.Second}
|
||||
return d.DialContext(ctx, wsURL, h)
|
||||
}
|
||||
|
||||
func agentWebSocketURL(baseURL, path string) (string, error) {
|
||||
u, err := url.Parse(baseURL)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
rel, err := url.Parse(path)
|
||||
if err != nil || !strings.HasPrefix(rel.Path, "/") {
|
||||
return "", errors.New("invalid agent websocket path")
|
||||
}
|
||||
if u.Scheme == "https" {
|
||||
u.Scheme = "wss"
|
||||
} else if u.Scheme == "http" {
|
||||
u.Scheme = "ws"
|
||||
} else {
|
||||
return "", errors.New("invalid agent websocket base URL")
|
||||
}
|
||||
u.Path = strings.TrimRight(u.Path, "/") + rel.Path
|
||||
u.RawQuery = rel.RawQuery
|
||||
u.Fragment = ""
|
||||
return u.String(), nil
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package nodes
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestAgentWebSocketURLPreservesQuery(t *testing.T) {
|
||||
got, err := agentWebSocketURL("https://agent.example/base", "/agent/v1/stacks/demo/terminal?service=web&shell=sh")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := "wss://agent.example/base/agent/v1/stacks/demo/terminal?service=web&shell=sh"
|
||||
if got != want {
|
||||
t.Fatalf("got %q want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateBaseURLRejectsQueryFragmentAndCredentials(t *testing.T) {
|
||||
bad := []string{
|
||||
"https://user:pass@agent.example",
|
||||
"https://agent.example?token=oops",
|
||||
"https://agent.example/#frag",
|
||||
"ftp://agent.example",
|
||||
}
|
||||
for _, u := range bad {
|
||||
if err := validateBaseURL(u); err == nil {
|
||||
t.Fatalf("expected %q to be rejected", u)
|
||||
}
|
||||
}
|
||||
if err := validateBaseURL("https://agent.example/base"); err != nil {
|
||||
t.Fatalf("valid URL rejected: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,533 @@
|
||||
package notify
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"crypto/tls"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/mail"
|
||||
"net/smtp"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Channel struct {
|
||||
ID int64 `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Config map[string]string `json:"config"`
|
||||
Enabled bool `json:"enabled"`
|
||||
CreatedAt int64 `json:"created_at"`
|
||||
UpdatedAt int64 `json:"updated_at"`
|
||||
}
|
||||
type Input struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type"`
|
||||
Config map[string]string `json:"config"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
type Message struct {
|
||||
Title string
|
||||
Body string
|
||||
Status string
|
||||
MonitorID int64
|
||||
}
|
||||
type Service struct {
|
||||
db *sql.DB
|
||||
key []byte
|
||||
client *http.Client
|
||||
}
|
||||
|
||||
func New(db *sql.DB, key []byte) *Service {
|
||||
return &Service{db: db, key: key, client: &http.Client{Timeout: 15 * time.Second}}
|
||||
}
|
||||
func sanitize(c map[string]string) map[string]string {
|
||||
out := map[string]string{}
|
||||
for k, v := range c {
|
||||
if isSecretKey(k) {
|
||||
if v != "" {
|
||||
out[k] = "••••••••"
|
||||
}
|
||||
} else {
|
||||
out[k] = v
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
func isSecretKey(k string) bool {
|
||||
lk := strings.ToLower(k)
|
||||
return strings.Contains(lk, "password") || strings.Contains(lk, "token") || strings.Contains(lk, "secret")
|
||||
}
|
||||
func (s *Service) encryptString(v string) (string, error) {
|
||||
if v == "" || strings.HasPrefix(v, "enc:v1:") {
|
||||
return v, nil
|
||||
}
|
||||
b, err := aes.NewCipher(s.key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
g, err := cipher.NewGCM(b)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
nonce := make([]byte, g.NonceSize())
|
||||
if _, err = rand.Read(nonce); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := g.Seal(nonce, nonce, []byte(v), nil)
|
||||
return "enc:v1:" + base64.RawStdEncoding.EncodeToString(out), nil
|
||||
}
|
||||
func (s *Service) decryptString(v string) (string, error) {
|
||||
if !strings.HasPrefix(v, "enc:v1:") {
|
||||
return v, nil
|
||||
}
|
||||
raw, err := base64.RawStdEncoding.DecodeString(strings.TrimPrefix(v, "enc:v1:"))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
b, err := aes.NewCipher(s.key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
g, err := cipher.NewGCM(b)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if len(raw) < g.NonceSize() {
|
||||
return "", errors.New("invalid encrypted notification config")
|
||||
}
|
||||
plain, err := g.Open(nil, raw[:g.NonceSize()], raw[g.NonceSize():], nil)
|
||||
return string(plain), err
|
||||
}
|
||||
func (s *Service) encodeConfig(c map[string]string) (string, error) {
|
||||
out := map[string]string{}
|
||||
for k, v := range c {
|
||||
if isSecretKey(k) {
|
||||
e, err := s.encryptString(v)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
out[k] = e
|
||||
} else {
|
||||
out[k] = v
|
||||
}
|
||||
}
|
||||
b, err := json.Marshal(out)
|
||||
return string(b), err
|
||||
}
|
||||
func (s *Service) decodeConfig(raw string) (map[string]string, error) {
|
||||
out := map[string]string{}
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
return out, nil
|
||||
}
|
||||
if err := json.Unmarshal([]byte(raw), &out); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for k, v := range out {
|
||||
if isSecretKey(k) {
|
||||
d, err := s.decryptString(v)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out[k] = d
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func normalize(in *Input) error {
|
||||
in.Name = strings.TrimSpace(in.Name)
|
||||
in.Type = strings.ToLower(strings.TrimSpace(in.Type))
|
||||
if in.Name == "" || len(in.Name) > 120 || strings.ContainsAny(in.Name, "\r\n") {
|
||||
return errors.New("valid notification name required")
|
||||
}
|
||||
if in.Config == nil {
|
||||
in.Config = map[string]string{}
|
||||
}
|
||||
endpoint := func(key string) error {
|
||||
raw := strings.TrimSpace(in.Config[key])
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil || (u.Scheme != "http" && u.Scheme != "https") || u.Host == "" || u.User != nil {
|
||||
return fmt.Errorf("%s must be an absolute http(s) URL without credentials", key)
|
||||
}
|
||||
in.Config[key] = strings.TrimRight(raw, "/")
|
||||
return nil
|
||||
}
|
||||
switch in.Type {
|
||||
case "webhook":
|
||||
if err := endpoint("url"); err != nil {
|
||||
return err
|
||||
}
|
||||
case "ntfy":
|
||||
if err := endpoint("server"); err != nil {
|
||||
return err
|
||||
}
|
||||
if topic := strings.TrimSpace(in.Config["topic"]); topic == "" || len(topic) > 200 || strings.ContainsAny(topic, "\r\n/?#") {
|
||||
return errors.New("valid ntfy topic required")
|
||||
}
|
||||
case "gotify":
|
||||
if err := endpoint("server"); err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(in.Config["token"]) == "" {
|
||||
return errors.New("gotify token required")
|
||||
}
|
||||
case "smtp":
|
||||
host := strings.TrimSpace(in.Config["host"])
|
||||
if host == "" || strings.ContainsAny(host, "\r\n/") {
|
||||
return errors.New("valid smtp host required")
|
||||
}
|
||||
security := strings.ToLower(strings.TrimSpace(in.Config["security"]))
|
||||
if security == "" {
|
||||
security = "starttls"
|
||||
}
|
||||
if security == "ssl" {
|
||||
security = "tls"
|
||||
}
|
||||
if security != "none" && security != "starttls" && security != "tls" {
|
||||
return errors.New("smtp security must be none, starttls or tls")
|
||||
}
|
||||
in.Config["security"] = security
|
||||
port := strings.TrimSpace(in.Config["port"])
|
||||
if port != "" {
|
||||
n, err := strconv.Atoi(port)
|
||||
if err != nil || n < 1 || n > 65535 {
|
||||
return errors.New("smtp port must be between 1 and 65535")
|
||||
}
|
||||
}
|
||||
if _, err := mail.ParseAddress(strings.TrimSpace(in.Config["from"])); err != nil {
|
||||
return errors.New("valid smtp from address required")
|
||||
}
|
||||
if _, err := mail.ParseAddressList(strings.TrimSpace(in.Config["to"])); err != nil {
|
||||
return errors.New("valid smtp recipient list required")
|
||||
}
|
||||
if boolConfig(in.Config["auth"]) && strings.TrimSpace(in.Config["username"]) == "" {
|
||||
return errors.New("smtp username required when authentication is enabled")
|
||||
}
|
||||
default:
|
||||
return errors.New("type must be webhook, ntfy, gotify or smtp")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) List(ctx context.Context) ([]Channel, error) {
|
||||
rows, e := s.db.QueryContext(ctx, `SELECT id,name,type,config_json,enabled,created_at,updated_at FROM notification_channels ORDER BY name`)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []Channel{}
|
||||
for rows.Next() {
|
||||
var c Channel
|
||||
var raw string
|
||||
if e := rows.Scan(&c.ID, &c.Name, &c.Type, &raw, &c.Enabled, &c.CreatedAt, &c.UpdatedAt); e != nil {
|
||||
return nil, e
|
||||
}
|
||||
c.Config, e = s.decodeConfig(raw)
|
||||
if e != nil {
|
||||
return nil, e
|
||||
}
|
||||
c.Config = sanitize(c.Config)
|
||||
out = append(out, c)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
func (s *Service) Create(ctx context.Context, in Input) (Channel, error) {
|
||||
if e := normalize(&in); e != nil {
|
||||
return Channel{}, e
|
||||
}
|
||||
en := true
|
||||
if in.Enabled != nil {
|
||||
en = *in.Enabled
|
||||
}
|
||||
raw, e := s.encodeConfig(in.Config)
|
||||
if e != nil {
|
||||
return Channel{}, e
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
r, e := s.db.ExecContext(ctx, `INSERT INTO notification_channels(name,type,config_json,enabled,created_at,updated_at) VALUES(?,?,?,?,?,?)`, in.Name, in.Type, raw, en, now, now)
|
||||
if e != nil {
|
||||
return Channel{}, e
|
||||
}
|
||||
id, _ := r.LastInsertId()
|
||||
return Channel{ID: id, Name: in.Name, Type: in.Type, Config: sanitize(in.Config), Enabled: en, CreatedAt: now, UpdatedAt: now}, nil
|
||||
}
|
||||
func (s *Service) Update(ctx context.Context, id int64, in Input) (Channel, error) {
|
||||
if e := normalize(&in); e != nil {
|
||||
return Channel{}, e
|
||||
}
|
||||
old, e := s.get(ctx, id)
|
||||
if e != nil {
|
||||
return Channel{}, e
|
||||
}
|
||||
for k, v := range in.Config {
|
||||
if v == "••••••••" {
|
||||
in.Config[k] = old.Config[k]
|
||||
}
|
||||
}
|
||||
en := true
|
||||
if in.Enabled != nil {
|
||||
en = *in.Enabled
|
||||
}
|
||||
raw, e := s.encodeConfig(in.Config)
|
||||
if e != nil {
|
||||
return Channel{}, e
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
_, e = s.db.ExecContext(ctx, `UPDATE notification_channels SET name=?,type=?,config_json=?,enabled=?,updated_at=? WHERE id=?`, in.Name, in.Type, raw, en, now, id)
|
||||
if e != nil {
|
||||
return Channel{}, e
|
||||
}
|
||||
return Channel{ID: id, Name: in.Name, Type: in.Type, Config: sanitize(in.Config), Enabled: en, CreatedAt: old.CreatedAt, UpdatedAt: now}, nil
|
||||
}
|
||||
func (s *Service) Delete(ctx context.Context, id int64) error {
|
||||
_, e := s.db.ExecContext(ctx, `DELETE FROM notification_channels WHERE id=?`, id)
|
||||
return e
|
||||
}
|
||||
func (s *Service) get(ctx context.Context, id int64) (Channel, error) {
|
||||
var c Channel
|
||||
var raw string
|
||||
e := s.db.QueryRowContext(ctx, `SELECT id,name,type,config_json,enabled,created_at,updated_at FROM notification_channels WHERE id=?`, id).Scan(&c.ID, &c.Name, &c.Type, &raw, &c.Enabled, &c.CreatedAt, &c.UpdatedAt)
|
||||
if e != nil {
|
||||
return c, e
|
||||
}
|
||||
c.Config, e = s.decodeConfig(raw)
|
||||
return c, e
|
||||
}
|
||||
func (s *Service) Test(ctx context.Context, id int64) error {
|
||||
c, e := s.get(ctx, id)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
return s.send(ctx, c, Message{Title: "Dockwatch test notification", Body: "Your notification provider is configured correctly.", Status: "test"})
|
||||
}
|
||||
func (s *Service) Broadcast(ctx context.Context, m Message) {
|
||||
rows, e := s.db.QueryContext(ctx, `SELECT id,name,type,config_json,enabled,created_at,updated_at FROM notification_channels WHERE enabled=1`)
|
||||
if e != nil {
|
||||
return
|
||||
}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
var c Channel
|
||||
var raw string
|
||||
if rows.Scan(&c.ID, &c.Name, &c.Type, &raw, &c.Enabled, &c.CreatedAt, &c.UpdatedAt) == nil {
|
||||
c.Config, e = s.decodeConfig(raw)
|
||||
if e != nil {
|
||||
continue
|
||||
}
|
||||
go func(c Channel) {
|
||||
x, k := context.WithTimeout(context.Background(), 20*time.Second)
|
||||
defer k()
|
||||
_ = s.send(x, c, m)
|
||||
}(c)
|
||||
}
|
||||
}
|
||||
}
|
||||
func (s *Service) send(ctx context.Context, c Channel, m Message) error {
|
||||
switch c.Type {
|
||||
case "webhook":
|
||||
return s.webhook(ctx, c, m)
|
||||
case "ntfy":
|
||||
return s.ntfy(ctx, c, m)
|
||||
case "gotify":
|
||||
return s.gotify(ctx, c, m)
|
||||
case "smtp":
|
||||
return s.smtp(ctx, c, m)
|
||||
}
|
||||
return errors.New("unsupported notification type")
|
||||
}
|
||||
func (s *Service) webhook(ctx context.Context, c Channel, m Message) error {
|
||||
u := c.Config["url"]
|
||||
if u == "" {
|
||||
return errors.New("webhook url required")
|
||||
}
|
||||
b, _ := json.Marshal(map[string]any{"title": m.Title, "body": m.Body, "status": m.Status, "monitor_id": m.MonitorID, "timestamp": time.Now().Unix()})
|
||||
req, e := http.NewRequestWithContext(ctx, "POST", u, bytes.NewReader(b))
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if t := c.Config["bearer_token"]; t != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+t)
|
||||
}
|
||||
r, e := s.client.Do(req)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
defer r.Body.Close()
|
||||
if r.StatusCode >= 300 {
|
||||
return fmt.Errorf("webhook: %s", r.Status)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) ntfy(ctx context.Context, c Channel, m Message) error {
|
||||
u := strings.TrimRight(c.Config["server"], "/") + "/" + c.Config["topic"]
|
||||
if c.Config["server"] == "" || c.Config["topic"] == "" {
|
||||
return errors.New("ntfy server and topic required")
|
||||
}
|
||||
req, e := http.NewRequestWithContext(ctx, "POST", u, strings.NewReader(m.Body))
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
req.Header.Set("Title", m.Title)
|
||||
req.Header.Set("Tags", map[string]string{"down": "rotating_light", "up": "white_check_mark"}[m.Status])
|
||||
if t := c.Config["token"]; t != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+t)
|
||||
}
|
||||
r, e := s.client.Do(req)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
defer r.Body.Close()
|
||||
if r.StatusCode >= 300 {
|
||||
return fmt.Errorf("ntfy: %s", r.Status)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) gotify(ctx context.Context, c Channel, m Message) error {
|
||||
server := strings.TrimRight(c.Config["server"], "/")
|
||||
token := c.Config["token"]
|
||||
if server == "" || token == "" {
|
||||
return errors.New("gotify server and token required")
|
||||
}
|
||||
u := server + "/message?token=" + url.QueryEscape(token)
|
||||
b, _ := json.Marshal(map[string]any{"title": m.Title, "message": m.Body, "priority": 5})
|
||||
req, e := http.NewRequestWithContext(ctx, "POST", u, bytes.NewReader(b))
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r, e := s.client.Do(req)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
defer r.Body.Close()
|
||||
if r.StatusCode >= 300 {
|
||||
return fmt.Errorf("gotify: %s", r.Status)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func boolConfig(v string) bool {
|
||||
v = strings.ToLower(strings.TrimSpace(v))
|
||||
return v == "1" || v == "true" || v == "yes" || v == "on"
|
||||
}
|
||||
|
||||
func (s *Service) smtp(ctx context.Context, c Channel, m Message) error {
|
||||
host := strings.TrimSpace(c.Config["host"])
|
||||
port := strings.TrimSpace(c.Config["port"])
|
||||
security := strings.ToLower(strings.TrimSpace(c.Config["security"]))
|
||||
if security == "" {
|
||||
security = "starttls"
|
||||
}
|
||||
if security == "ssl" {
|
||||
security = "tls"
|
||||
}
|
||||
if port == "" {
|
||||
if security == "tls" {
|
||||
port = "465"
|
||||
} else {
|
||||
port = "587"
|
||||
}
|
||||
}
|
||||
fromRaw := strings.TrimSpace(c.Config["from"])
|
||||
toRaw := strings.TrimSpace(c.Config["to"])
|
||||
if host == "" || fromRaw == "" || toRaw == "" {
|
||||
return errors.New("smtp host, from and to required")
|
||||
}
|
||||
fromAddr, err := mail.ParseAddress(fromRaw)
|
||||
if err != nil {
|
||||
return errors.New("invalid smtp from address")
|
||||
}
|
||||
toAddrs, err := mail.ParseAddressList(toRaw)
|
||||
if err != nil || len(toAddrs) == 0 {
|
||||
return errors.New("invalid smtp recipient list")
|
||||
}
|
||||
if security != "none" && security != "starttls" && security != "tls" {
|
||||
return errors.New("smtp security must be none, starttls or tls")
|
||||
}
|
||||
authEnabled := boolConfig(c.Config["auth"])
|
||||
// Backward compatibility: existing configurations with a username implied auth.
|
||||
if c.Config["auth"] == "" && strings.TrimSpace(c.Config["username"]) != "" {
|
||||
authEnabled = true
|
||||
}
|
||||
if authEnabled && strings.TrimSpace(c.Config["username"]) == "" {
|
||||
return errors.New("smtp username required when authentication is enabled")
|
||||
}
|
||||
addr := net.JoinHostPort(host, port)
|
||||
dialer := &net.Dialer{Timeout: 15 * time.Second}
|
||||
tlsCfg := &tls.Config{ServerName: host, MinVersion: tls.VersionTLS12, InsecureSkipVerify: boolConfig(c.Config["skip_verify"])}
|
||||
var conn net.Conn
|
||||
if security == "tls" {
|
||||
conn, err = tls.DialWithDialer(dialer, "tcp", addr, tlsCfg)
|
||||
} else {
|
||||
conn, err = dialer.DialContext(ctx, "tcp", addr)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer conn.Close()
|
||||
if deadline, ok := ctx.Deadline(); ok {
|
||||
_ = conn.SetDeadline(deadline)
|
||||
} else {
|
||||
_ = conn.SetDeadline(time.Now().Add(20 * time.Second))
|
||||
}
|
||||
cl, err := smtp.NewClient(conn, host)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer cl.Close()
|
||||
if security == "starttls" {
|
||||
ok, _ := cl.Extension("STARTTLS")
|
||||
if !ok {
|
||||
return errors.New("smtp server does not support STARTTLS")
|
||||
}
|
||||
if err = cl.StartTLS(tlsCfg); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if authEnabled {
|
||||
if ok, _ := cl.Extension("AUTH"); !ok {
|
||||
return errors.New("smtp server does not advertise AUTH")
|
||||
}
|
||||
auth := smtp.PlainAuth("", c.Config["username"], c.Config["password"], host)
|
||||
if err = cl.Auth(auth); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err = cl.Mail(fromAddr.Address); err != nil {
|
||||
return err
|
||||
}
|
||||
tos := make([]string, 0, len(toAddrs))
|
||||
for _, a := range toAddrs {
|
||||
tos = append(tos, a.String())
|
||||
if err = cl.Rcpt(a.Address); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
w, err := cl.Data()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
subject := strings.NewReplacer("\r", " ", "\n", " ").Replace(m.Title)
|
||||
body := strings.ReplaceAll(strings.ReplaceAll(m.Body, "\r\n", "\n"), "\r", "\n")
|
||||
body = strings.ReplaceAll(body, "\n", "\r\n")
|
||||
msg := []byte("To: " + strings.Join(tos, ", ") + "\r\nFrom: " + fromAddr.String() + "\r\nSubject: " + subject + "\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\n\r\n" + body + "\r\n")
|
||||
if _, err = w.Write(msg); err != nil {
|
||||
_ = w.Close()
|
||||
return err
|
||||
}
|
||||
if err = w.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
return cl.Quit()
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
package notify
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestConfigSecretsEncryptedAtRest(t *testing.T) {
|
||||
s := New(nil, []byte("0123456789abcdef0123456789abcdef"))
|
||||
raw, err := s.encodeConfig(map[string]string{"url": "https://example.invalid/hook", "token": "super-secret-token", "password": "pw123"})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(raw, "super-secret-token") || strings.Contains(raw, "pw123") {
|
||||
t.Fatalf("secret leaked in persisted config: %s", raw)
|
||||
}
|
||||
cfg, err := s.decodeConfig(raw)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if cfg["token"] != "super-secret-token" || cfg["password"] != "pw123" {
|
||||
t.Fatalf("roundtrip mismatch: %#v", cfg)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user