diff --git a/pkg/aqhttp/aqhttp.go b/pkg/aqhttp/aqhttp.go index 1fbc59eb1..2e1c7a5e2 100644 --- a/pkg/aqhttp/aqhttp.go +++ b/pkg/aqhttp/aqhttp.go @@ -1,30 +1,49 @@ package aqhttp import ( + "context" "net/http" "time" ) -var Client http.Client var UserAgent string = "streamplace/unknown" -type AddHeaderTransport struct { - T http.RoundTripper -} +// Client is the default HTTP client with SSRF protection. +// Uses DNS-over-HTTPS to validate destination IPs and blocks private/bogon ranges. +// For trusted infrastructure endpoints, use TrustedClient instead. +var Client http.Client -func (adt *AddHeaderTransport) RoundTrip(req *http.Request) (*http.Response, error) { - req.Header.Add("User-Agent", UserAgent) - return adt.T.RoundTrip(req) -} +// TrustedClient is an HTTP client without SSRF protection. +// Use this only for trusted infrastructure endpoints (e.g., livepeer.com) +// where the validation overhead is problematic. +var TrustedClient http.Client func init() { Client = http.Client{ - Transport: &AddHeaderTransport{T: &http.Transport{}}, - // do not follow redirects automatically + Transport: NewUntrustedTransport(), + CheckRedirect: func(req *http.Request, via []*http.Request) error { + return http.ErrUseLastResponse + }, + Timeout: 30 * time.Second, + } + + TrustedClient = http.Client{ + Transport: NewTrustedTransport(), CheckRedirect: func(req *http.Request, via []*http.Request) error { return http.ErrUseLastResponse }, - // add reasonable timeout Timeout: 30 * time.Second, } } + +// Do executes an HTTP request with SSRF protection (secure by default). +// Most callsites should use this function. +func Do(ctx context.Context, req *http.Request) (*http.Response, error) { + return Client.Do(req.WithContext(ctx)) +} + +// DoTrusted executes an HTTP request without SSRF protection. +// Use this only for trusted infrastructure endpoints. +func DoTrusted(ctx context.Context, req *http.Request) (*http.Response, error) { + return TrustedClient.Do(req.WithContext(ctx)) +} diff --git a/pkg/aqhttp/resolv.go b/pkg/aqhttp/resolv.go index d5f8a4abd..b9314e615 100644 --- a/pkg/aqhttp/resolv.go +++ b/pkg/aqhttp/resolv.go @@ -7,6 +7,7 @@ import ( "net" "net/http" "net/url" + "sync" "time" ) @@ -15,10 +16,17 @@ const ( TypeAAAA = 28 // IPv6 ) +type dnsRecord struct { + ips []string + expiresAt time.Time +} + type DoHResolver struct { Server string Client *http.Client invalidRanges []*net.IPNet + cache map[string]*dnsRecord + mu sync.RWMutex } func NewDoHResolver(server string) *DoHResolver { @@ -60,6 +68,7 @@ func NewDoHResolver(server string) *DoHResolver { Timeout: 10 * time.Second, }, invalidRanges: invalidRanges, + cache: make(map[string]*dnsRecord), } } @@ -74,6 +83,17 @@ type DoHResponse struct { } func (r *DoHResolver) Resolve(domain string, recordType int) ([]string, error) { + cacheKey := fmt.Sprintf("%s:%d", domain, recordType) + + r.mu.RLock() + if record, ok := r.cache[cacheKey]; ok { + if time.Now().Before(record.expiresAt) { + defer r.mu.RUnlock() + return record.ips, nil + } + } + r.mu.RUnlock() + reqURL := fmt.Sprintf("%s?name=%s&type=%d", r.Server, url.QueryEscape(domain), recordType) req, err := http.NewRequest("GET", reqURL, nil) @@ -104,10 +124,23 @@ func (r *DoHResolver) Resolve(domain string, recordType int) ([]string, error) { } var results []string + var minTTL = 3600 for _, answer := range dohResp.Answer { if answer.Type == recordType { results = append(results, answer.Data) + if answer.TTL < minTTL { + minTTL = answer.TTL + } + } + } + + if len(results) > 0 { + r.mu.Lock() + r.cache[cacheKey] = &dnsRecord{ + ips: results, + expiresAt: time.Now().Add(time.Duration(minTTL) * time.Second), } + r.mu.Unlock() } return results, nil diff --git a/pkg/aqhttp/transport.go b/pkg/aqhttp/transport.go new file mode 100644 index 000000000..676d7962b --- /dev/null +++ b/pkg/aqhttp/transport.go @@ -0,0 +1,101 @@ +package aqhttp + +import ( + "context" + "fmt" + "net" + "net/http" + "time" +) + +// TrustedTransport is a basic transport that adds User-Agent headers. +// Use this for trusted infrastructure endpoints where SSRF is not a concern. +type TrustedTransport struct { + Base http.RoundTripper +} + +func (t *TrustedTransport) RoundTrip(req *http.Request) (*http.Response, error) { + req.Header.Add("User-Agent", UserAgent) + return t.Base.RoundTrip(req) +} + +// NewTrustedTransport creates a transport for trusted endpoints. +func NewTrustedTransport() *TrustedTransport { + return &TrustedTransport{ + Base: &http.Transport{ + MaxIdleConns: 100, + IdleConnTimeout: 90 * time.Second, + TLSHandshakeTimeout: 10 * time.Second, + }, + } +} + +// UntrustedTransport validates destination IPs using DNS-over-HTTPS before connecting. +// Prevents SSRF attacks by blocking private, loopback, and bogon IP ranges. +type UntrustedTransport struct { + Base http.RoundTripper + resolver *DoHResolver +} + +func (t *UntrustedTransport) RoundTrip(req *http.Request) (*http.Response, error) { + req.Header.Add("User-Agent", UserAgent) + return t.Base.RoundTrip(req) +} + +// NewUntrustedTransport creates a transport that validates all destination IPs. +func NewUntrustedTransport() *UntrustedTransport { + resolver := NewDoHResolver("") + + dialer := &net.Dialer{ + Timeout: 30 * time.Second, + KeepAlive: 30 * time.Second, + } + + transport := &http.Transport{ + DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) { + host, port, err := net.SplitHostPort(addr) + if err != nil { + return nil, fmt.Errorf("failed to parse address: %w", err) + } + + // Resolve IPv4 addresses using DoH + ipv4Addrs, _ := resolver.Resolve(host, TypeA) + var validIP string + + // Check IPv4 addresses first + for _, ip := range ipv4Addrs { + if !resolver.IsInvalidIP(ip) { + validIP = ip + break + } + } + + // Fall back to IPv6 if no valid IPv4 + if validIP == "" { + ipv6Addrs, _ := resolver.Resolve(host, TypeAAAA) + for _, ip := range ipv6Addrs { + if !resolver.IsInvalidIP(ip) { + validIP = ip + break + } + } + } + + if validIP == "" { + return nil, fmt.Errorf("all resolved IPs for %s are private/invalid", host) + } + + // Dial using the validated IP + targetAddr := net.JoinHostPort(validIP, port) + return dialer.DialContext(ctx, network, targetAddr) + }, + MaxIdleConns: 100, + IdleConnTimeout: 90 * time.Second, + TLSHandshakeTimeout: 10 * time.Second, + } + + return &UntrustedTransport{ + Base: transport, + resolver: resolver, + } +} diff --git a/pkg/integrations/discord/send-chat.go b/pkg/integrations/discord/send-chat.go index 897f52241..93e9f5a1b 100644 --- a/pkg/integrations/discord/send-chat.go +++ b/pkg/integrations/discord/send-chat.go @@ -6,11 +6,9 @@ import ( "encoding/json" "fmt" "io" - "net" "net/http" "strings" - "golang.org/x/net/context/ctxhttp" "stream.place/streamplace/pkg/aqhttp" "stream.place/streamplace/pkg/integrations/discord/discordtypes" "stream.place/streamplace/pkg/log" @@ -19,12 +17,6 @@ import ( func SendChat(ctx context.Context, w *discordtypes.Webhook, did string, scm *streamplace.ChatDefs_MessageView) error { - resolv := aqhttp.NewDoHResolver("") - targetIP, parsedURL, err := resolv.ValidateAndGetIP(w.URL) - if err != nil { - return fmt.Errorf("webhook URL validation failed: %w", err) - } - msg, ok := scm.Record.Val.(*streamplace.ChatMessage) if !ok { return fmt.Errorf("failed to cast chat message to streamplace chat message") @@ -61,26 +53,13 @@ func SendChat(ctx context.Context, w *discordtypes.Webhook, did string, scm *str log.Warn(ctx, "sending chat to discord", "payload", string(jsonPayload)) - port := parsedURL.Port() - if port == "" { - // Use default port based on scheme - if parsedURL.Scheme == "https" { - port = "443" - } else { - port = "80" - } - } - requestURL := fmt.Sprintf("%s://%s%s", parsedURL.Scheme, net.JoinHostPort(targetIP, port), parsedURL.RequestURI()) - - req, err := http.NewRequestWithContext(ctx, "POST", requestURL, bytes.NewReader(jsonPayload)) + req, err := http.NewRequestWithContext(ctx, "POST", w.URL, bytes.NewReader(jsonPayload)) if err != nil { return fmt.Errorf("failed to create request: %w", err) } req.Header.Set("Content-Type", "application/json") - // set Host header to the original hostname - req.Host = parsedURL.Host - resp, err := ctxhttp.Do(ctx, &aqhttp.Client, req) + resp, err := aqhttp.Do(ctx, req) if err != nil { return fmt.Errorf("failed to send request: %w", err) } diff --git a/pkg/integrations/discord/send-livestream.go b/pkg/integrations/discord/send-livestream.go index e509b9bca..a865f61f2 100644 --- a/pkg/integrations/discord/send-livestream.go +++ b/pkg/integrations/discord/send-livestream.go @@ -6,14 +6,12 @@ import ( "encoding/json" "fmt" "io" - "net" "net/http" "net/url" "strconv" "strings" "github.com/bluesky-social/indigo/api/bsky" - "golang.org/x/net/context/ctxhttp" "stream.place/streamplace/pkg/aqhttp" "stream.place/streamplace/pkg/integrations/discord/discordtypes" "stream.place/streamplace/pkg/log" @@ -23,13 +21,6 @@ import ( func SendLivestream(ctx context.Context, w *discordtypes.Webhook, pdsURL string, lsv *streamplace.Livestream_LivestreamView, postView *bsky.FeedDefs_PostView, spcp *streamplace.ChatProfile) error { - // get safe IP - resolv := aqhttp.NewDoHResolver("") - targetIP, parsedURL, err := resolv.ValidateAndGetIP(w.URL) - if err != nil { - return fmt.Errorf("webhook URL validation failed: %w", err) - } - ctx = log.WithLogValues(ctx, "func", "SendLivestream") ls, ok := lsv.Record.Val.(*streamplace.Livestream) if !ok { @@ -104,24 +95,13 @@ func SendLivestream(ctx context.Context, w *discordtypes.Webhook, pdsURL string, log.Warn(ctx, "sending livestream to discord", "payload", string(jsonPayload)) - port := parsedURL.Port() - if port == "" { - if parsedURL.Scheme == "https" { - port = "443" - } else { - port = "80" - } - } - requestURL := fmt.Sprintf("%s://%s%s", parsedURL.Scheme, net.JoinHostPort(targetIP, port), parsedURL.RequestURI()) - - req, err := http.NewRequestWithContext(ctx, "POST", requestURL, bytes.NewReader(jsonPayload)) + req, err := http.NewRequestWithContext(ctx, "POST", w.URL, bytes.NewReader(jsonPayload)) if err != nil { return fmt.Errorf("failed to create request: %w", err) } req.Header.Set("Content-Type", "application/json") - req.Host = parsedURL.Host - resp, err := ctxhttp.Do(ctx, &aqhttp.Client, req) + resp, err := aqhttp.Do(ctx, req) if err != nil { return fmt.Errorf("failed to send request: %w", err) } diff --git a/pkg/livepeer/livepeer.go b/pkg/livepeer/livepeer.go index 6b3609f3b..91caf1a4b 100644 --- a/pkg/livepeer/livepeer.go +++ b/pkg/livepeer/livepeer.go @@ -14,7 +14,6 @@ import ( "strings" "time" - "golang.org/x/net/context/ctxhttp" "stream.place/streamplace/pkg/aqhttp" "stream.place/streamplace/pkg/config" "stream.place/streamplace/pkg/log" @@ -136,7 +135,7 @@ func (ls *LivepeerSession) PostSegmentToGateway(ctx context.Context, buf []byte, log.Log(ctx, "wrote debug file", "file", debugFile) } - resp, err := ctxhttp.Do(ctx, &aqhttp.Client, req) + resp, err := aqhttp.DoTrusted(ctx, req) if err != nil { <-ls.Guard return nil, fmt.Errorf("failed to send segment to gateway (config %s): %w", string(bs), err) diff --git a/pkg/spxrpc/og.go b/pkg/spxrpc/og.go index 8b9e70691..9211d1991 100644 --- a/pkg/spxrpc/og.go +++ b/pkg/spxrpc/og.go @@ -19,7 +19,6 @@ import ( imagedraw "image/draw" "golang.org/x/image/draw" - "golang.org/x/net/context/ctxhttp" "github.com/bluesky-social/indigo/api/bsky" "github.com/bluesky-social/indigo/xrpc" @@ -208,7 +207,7 @@ func downloadImage(ctx context.Context, url string) ([]byte, error) { return nil, fmt.Errorf("failed to create request: %w", err) } - resp, err := ctxhttp.Do(ctx, &aqhttp.Client, req) + resp, err := aqhttp.Do(ctx, req) if err != nil { return nil, fmt.Errorf("HTTP request failed: %w", err) }