534 lines
15 KiB
Go
534 lines
15 KiB
Go
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()
|
|
}
|