316 lines
8.4 KiB
Go
316 lines
8.4 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) {
|
|
return m.DoWithTimeout(ctx, id, method, path, body, m.client.Timeout)
|
|
}
|
|
|
|
// DoWithTimeout performs an authenticated agent request with an operation-specific
|
|
// timeout. Long-running host package operations use this rather than weakening
|
|
// the normal 30-second control-plane timeout for every request.
|
|
func (m *Manager) DoWithTimeout(ctx context.Context, id int64, method, path string, body any, timeout time.Duration) ([]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")
|
|
}
|
|
client := *m.client
|
|
client.Timeout = timeout
|
|
resp, err := 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
|
|
}
|