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 ". 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() }