@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user