Files
dockwatch/internal/nodes/nodes.go
T
jbergner 45ca18b74e
release-tag / release-image (push) Failing after 1m20s
init
2026-08-31 17:09:21 +02:00

307 lines
8.0 KiB
Go

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
}