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 }