141 lines
3.9 KiB
Go
141 lines
3.9 KiB
Go
package agent
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"time"
|
|
|
|
"gopkg.in/yaml.v3"
|
|
)
|
|
|
|
// Client talks to the tomswallapi control plane for one device.
|
|
type Client struct {
|
|
BaseURL string
|
|
Device string
|
|
Token string
|
|
HTTP *http.Client
|
|
}
|
|
|
|
// NewClient builds a Client with a sane default timeout.
|
|
func NewClient(baseURL, device, token string) *Client {
|
|
return &Client{
|
|
BaseURL: baseURL,
|
|
Device: device,
|
|
Token: token,
|
|
// No keep-alives: every request, the post-apply check included, opens a
|
|
// fresh connection that must pass the current ruleset.
|
|
HTTP: &http.Client{Timeout: 30 * time.Second, Transport: noKeepAlive()},
|
|
}
|
|
}
|
|
|
|
// FetchConfig retrieves the device's rendered config. It returns both the parsed
|
|
// document and the raw bytes (so callers can cache exactly what was served).
|
|
func (c *Client) FetchConfig(ctx context.Context) (*RenderedConfig, []byte, error) {
|
|
url := fmt.Sprintf("%s/api/v1/devices/%s/config", c.BaseURL, c.Device)
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
req.Header.Set("Authorization", "Bearer "+c.Token)
|
|
req.Header.Set("Accept", "application/yaml")
|
|
|
|
resp, err := c.HTTP.Do(req)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
defer resp.Body.Close()
|
|
body, err := io.ReadAll(resp.Body)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
if resp.StatusCode != http.StatusOK {
|
|
return nil, nil, fmt.Errorf("fetch config: status %d: %s", resp.StatusCode, bytes.TrimSpace(body))
|
|
}
|
|
|
|
cfg, err := ParseRendered(body)
|
|
if err != nil {
|
|
return nil, nil, err
|
|
}
|
|
return cfg, body, nil
|
|
}
|
|
|
|
// ParseRendered decodes a rendered config document (YAML, JSON is a subset).
|
|
func ParseRendered(body []byte) (*RenderedConfig, error) {
|
|
var cfg RenderedConfig
|
|
if err := yaml.Unmarshal(body, &cfg); err != nil {
|
|
return nil, fmt.Errorf("parsing rendered config: %w", err)
|
|
}
|
|
return &cfg, nil
|
|
}
|
|
|
|
// ReportRoutes reports the device's reachable prefixes (FIB) so the control
|
|
// plane can scope which routers enforce a rule.
|
|
func (c *Client) ReportRoutes(ctx context.Context, prefixes []string) error {
|
|
url := fmt.Sprintf("%s/api/v1/devices/%s/routes", c.BaseURL, c.Device)
|
|
payload, _ := json.Marshal(map[string][]string{"prefixes": prefixes})
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
req.Header.Set("Authorization", "Bearer "+c.Token)
|
|
req.Header.Set("Content-Type", "application/json")
|
|
|
|
resp, err := c.HTTP.Do(req)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer resp.Body.Close()
|
|
io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<16))
|
|
if resp.StatusCode >= 400 {
|
|
return fmt.Errorf("report routes: %d", resp.StatusCode)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Status values reported to POST /api/v1/devices/{name}/status.
|
|
const (
|
|
StatusApplied = "applied"
|
|
StatusReverted = "reverted"
|
|
StatusFailed = "failed"
|
|
)
|
|
|
|
// Status is the outcome of applying one generation.
|
|
type Status struct {
|
|
Status string `json:"status"`
|
|
Generation int64 `json:"generation"`
|
|
Error string `json:"error,omitempty"`
|
|
}
|
|
|
|
// ReportStatus tells the control plane the outcome of applying a generation.
|
|
func (c *Client) ReportStatus(ctx context.Context, st Status) error {
|
|
url := fmt.Sprintf("%s/api/v1/devices/%s/status", c.BaseURL, c.Device)
|
|
payload, _ := json.Marshal(st)
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
|
|
if err != nil {
|
|
return err
|
|
}
|
|
req.Header.Set("Authorization", "Bearer "+c.Token)
|
|
req.Header.Set("Content-Type", "application/json")
|
|
|
|
resp, err := c.HTTP.Do(req)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer resp.Body.Close()
|
|
io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<16))
|
|
if resp.StatusCode >= 400 {
|
|
return fmt.Errorf("report status: %d", resp.StatusCode)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func noKeepAlive() http.RoundTripper {
|
|
t := http.DefaultTransport.(*http.Transport).Clone()
|
|
t.DisableKeepAlives = true
|
|
return t
|
|
}
|