Files
summercms/modules/fetchguard/client.go
Jakub Zych e6a67134d1 feat(14-01): fetchguard client covers PUT, multipart, bearer and a trusted mode
- TrustedMode (declared after PublicOnlyMode) lifts the scheme, host and dial checks for Client only
- PutJSON, PostMultipart with FormField/FormFile, Bearer
- tests for modes, redirects, multipart order, body cap and the scheme guard
- README, root modules row and outbound HTTP docs describe the client and its test seam
2026-10-03 19:42:37 +02:00

279 lines
8.5 KiB
Go

package fetchguard
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"net/textproto"
"net/url"
"strings"
"time"
"git.golem15.com/golem15/summercms/modules/compass"
)
// Client is a guarded outbound HTTP client for calling service APIs: any
// method, JSON or multipart bodies, caller-supplied headers. It keeps the
// guarantees of Fetch for the policy it was built with: URL validation
// before any I/O (https only and the host list, except in TrustedMode), the
// dial-time private and reserved address check on every new connection
// (except in TrustedMode), a response body cap and no redirects. Status codes
// are returned, never judged. Request bodies are not capped: callers bound
// their own inputs.
//
// A Client is safe for concurrent use; build one per vendor and reuse it so
// connections are kept alive.
type Client struct {
policy Policy
maxBytes int64
timeout time.Duration
http *http.Client
}
// NewClient resolves the policy's limits once (zero MaxBytes or Timeout fall
// back to cfg, then to Defaults) and builds the client.
func NewClient(policy Policy, cfg *compass.Config) (*Client, error) {
maxBytes, timeout, err := resolveLimits(policy, cfg)
if err != nil {
return nil, err
}
return &Client{
policy: policy,
maxBytes: maxBytes,
timeout: timeout,
http: newHTTPClient(newTransport(policy, timeout, true, policy.Mode != TrustedMode), timeout),
}, nil
}
type transportKey struct{}
// WithTransport returns a context whose requests a Client sends through rt
// instead of its own network transport. It is the test and parity-replay
// seam: a test hands it a fake http.RoundTripper (for example the tide
// upstream fake) so the client's real request is asserted offline.
//
// The override is code-only. No Policy field, Client field, config key,
// environment variable or request header can set it; only Go code holding
// the context can. URL validation still runs before the override is
// consulted, and redirects are still not followed.
func WithTransport(ctx context.Context, rt http.RoundTripper) context.Context {
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, transportKey{}, rt)
}
func transportFrom(ctx context.Context) http.RoundTripper {
if ctx == nil {
return nil
}
rt, _ := ctx.Value(transportKey{}).(http.RoundTripper)
return rt
}
// Do validates req's URL against the client's policy and sends it. The
// response body is capped at the policy's MaxBytes: reading past it fails
// with an *Error whose Reason is ReasonTooLarge. The caller closes the body.
// A 3xx response is returned as-is; redirects are never followed.
//
// A non-nil error is always *Error.
func (c *Client) Do(req *http.Request) (*http.Response, error) {
if req == nil || req.URL == nil {
return nil, &Error{Reason: ReasonInvalidURL}
}
if err := c.check(req.URL); err != nil {
return nil, err
}
hc := c.http
if rt := transportFrom(req.Context()); rt != nil {
hc = newHTTPClient(rt, c.timeout)
}
resp, err := hc.Do(req)
if err != nil {
return nil, mapTransportError(err)
}
resp.Body = &cappedBody{rc: resp.Body, max: c.maxBytes}
return resp, nil
}
// Send sends req through Do and reads the whole capped body into a Result.
func (c *Client) Send(req *http.Request) (*Result, error) {
resp, err := c.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
if err != nil {
var fe *Error
if errors.As(err, &fe) {
return nil, fe
}
return nil, &Error{Reason: ReasonNetworkError, Err: err}
}
return &Result{
Body: data,
ContentType: resp.Header.Get("Content-Type"),
StatusCode: resp.StatusCode,
Header: resp.Header,
}, nil
}
// PostJSON marshals body as JSON and POSTs it to rawURL. Content-Type is set
// to application/json first, then header is copied over it, so a caller can
// override any header including Content-Type.
func (c *Client) PostJSON(ctx context.Context, rawURL string, header http.Header, body any) (*Result, error) {
return c.sendJSON(ctx, http.MethodPost, rawURL, header, body)
}
// PutJSON is PostJSON with the PUT method.
func (c *Client) PutJSON(ctx context.Context, rawURL string, header http.Header, body any) (*Result, error) {
return c.sendJSON(ctx, http.MethodPut, rawURL, header, body)
}
// FormField is one plain multipart/form-data field.
type FormField struct {
Name string
Value string
}
// FormFile is one multipart/form-data file part. ContentType defaults to
// application/octet-stream.
type FormFile struct {
Field string
Filename string
ContentType string
Body io.Reader
}
// PostMultipart POSTs a multipart/form-data body to rawURL: every field in
// slice order, then every file in slice order. header is copied first; the
// multipart Content-Type with its boundary always wins.
func (c *Client) PostMultipart(ctx context.Context, rawURL string, header http.Header, fields []FormField, files []FormFile) (*Result, error) {
var buf bytes.Buffer
mw := multipart.NewWriter(&buf)
for _, f := range fields {
if err := mw.WriteField(f.Name, f.Value); err != nil {
return nil, &Error{Reason: ReasonInvalidURL, Err: err}
}
}
for _, f := range files {
ct := f.ContentType
if ct == "" {
ct = "application/octet-stream"
}
h := textproto.MIMEHeader{}
h.Set("Content-Disposition", fmt.Sprintf(`form-data; name="%s"; filename="%s"`,
quoteEscaper.Replace(f.Field), quoteEscaper.Replace(f.Filename)))
h.Set("Content-Type", ct)
part, err := mw.CreatePart(h)
if err != nil {
return nil, &Error{Reason: ReasonInvalidURL, Err: err}
}
if f.Body != nil {
if _, err := io.Copy(part, f.Body); err != nil {
return nil, &Error{Reason: ReasonNetworkError, Err: err}
}
}
}
if err := mw.Close(); err != nil {
return nil, &Error{Reason: ReasonInvalidURL, Err: err}
}
last := http.Header{"Content-Type": {mw.FormDataContentType()}}
return c.sendWith(ctx, http.MethodPost, rawURL, header, last, &buf)
}
var quoteEscaper = strings.NewReplacer("\\", "\\\\", `"`, "\\\"")
// Bearer returns the Authorization header value "Bearer <token>".
func Bearer(token string) string {
return "Bearer " + token
}
// Get sends a GET to rawURL with header.
func (c *Client) Get(ctx context.Context, rawURL string, header http.Header) (*Result, error) {
return c.send(ctx, http.MethodGet, rawURL, nil, header, nil)
}
func (c *Client) sendJSON(ctx context.Context, method, rawURL string, header http.Header, body any) (*Result, error) {
raw, err := json.Marshal(body)
if err != nil {
return nil, &Error{Reason: ReasonInvalidURL, Err: err}
}
h := http.Header{}
h.Set("Content-Type", "application/json")
return c.send(ctx, method, rawURL, h, header, bytes.NewReader(raw))
}
func (c *Client) send(ctx context.Context, method, rawURL string, base, header http.Header, body io.Reader) (*Result, error) {
return c.sendWith(ctx, method, rawURL, base, header, body)
}
// sendWith builds a request whose headers are first, then second (second
// wins per header name) and sends it.
func (c *Client) sendWith(ctx context.Context, method, rawURL string, first, second http.Header, body io.Reader) (*Result, error) {
if ctx == nil {
ctx = context.Background()
}
req, err := http.NewRequestWithContext(ctx, method, rawURL, body)
if err != nil {
return nil, &Error{Reason: ReasonInvalidURL, Err: err}
}
for _, h := range []http.Header{first, second} {
for k, vs := range h {
req.Header[http.CanonicalHeaderKey(k)] = append([]string(nil), vs...)
}
}
return c.Send(req)
}
// check applies the policy's URL rules before any I/O.
func (c *Client) check(u *url.URL) error {
if u.Scheme == "" || u.Host == "" {
return &Error{Reason: ReasonInvalidURL}
}
scheme := strings.ToLower(u.Scheme)
if c.policy.Mode == TrustedMode {
if scheme != "https" && scheme != "http" {
return &Error{Reason: ReasonScheme}
}
return nil
}
if scheme != "https" {
return &Error{Reason: ReasonScheme}
}
if c.policy.Mode == AllowHostsMode && !hostAllowed(u.Hostname(), c.policy.AllowHosts) {
return &Error{Reason: ReasonInvalidURL}
}
return nil
}
// cappedBody fails with ReasonTooLarge once more than max bytes were read.
type cappedBody struct {
rc io.ReadCloser
max int64
n int64
}
func (b *cappedBody) Read(p []byte) (int, error) {
if b.n > b.max {
return 0, &Error{Reason: ReasonTooLarge}
}
if room := b.max - b.n + 1; int64(len(p)) > room {
p = p[:room]
}
n, err := b.rc.Read(p)
b.n += int64(n)
if b.n > b.max {
return n - int(b.n-b.max), &Error{Reason: ReasonTooLarge}
}
return n, err
}
func (b *cappedBody) Close() error { return b.rc.Close() }