mediaproxy.mx raw

   1  // Package mediaproxy fetches a remote HTTP/HTTPS resource so it can be
   2  // re-served from the musiquay origin with COEP-compatible CORP headers.
   3  //
   4  // The relay's epoll loop is single-threaded; Fetch blocks the loop for the
   5  // duration of the upstream call. Acceptable for a personal-scale relay; if
   6  // concurrent proxy load becomes a problem this should move to a spawned
   7  // worker domain that returns the response over a channel.
   8  package mediaproxy
   9  
  10  import (
  11  	"bufio"
  12  	"bytes"
  13  	"crypto/tls"
  14  	"fmt"
  15  	"io"
  16  	"net"
  17  	"net/url"
  18  	"strconv"
  19  	"time"
  20  
  21  	"git.smesh.lol/nostr/pkg/ws"
  22  )
  23  
  24  const (
  25  	DefaultMaxBytes int64         = 32 * 1024 * 1024
  26  	DefaultTimeout  time.Duration = 8 * time.Second
  27  	maxRedirects                  = 5
  28  )
  29  
  30  // Fetch performs a GET on rawURL with redirect following (max 5 hops).
  31  // Returns upstream status, response headers (lowercased keys), body.
  32  // Only http and https schemes accepted.
  33  func Fetch(rawURL string, maxBytes int64) (code int32, hdrs map[string]string, out []byte, err error) {
  34  	if maxBytes <= 0 {
  35  		maxBytes = DefaultMaxBytes
  36  	}
  37  	current := rawURL
  38  	for hop := 0; hop < maxRedirects; hop++ {
  39  		status, headers, body, err2 := fetchOnce(current, maxBytes)
  40  		if err2 != nil {
  41  			return 0, nil, nil, err2
  42  		}
  43  		if status >= 300 && status < 400 {
  44  			loc := headers["location"]
  45  			if loc == "" {
  46  				return status, headers, body, nil
  47  			}
  48  			next, err1 := resolveRedirect(current, loc)
  49  			if err1 != nil {
  50  				return 0, nil, nil, err1
  51  			}
  52  			current = next
  53  			continue
  54  		}
  55  		return status, headers, body, nil
  56  	}
  57  	return 0, nil, nil, fmt.Errorf("too many redirects")
  58  }
  59  
  60  func fetchOnce(rawURL string, maxBytes int64) (status int32, hdrs map[string]string, out []byte, err error) {
  61  	u, err9 := url.Parse(rawURL)
  62  	if err9 != nil {
  63  		return 0, nil, nil, fmt.Errorf("parse: %w", err9)
  64  	}
  65  	useTLS := false
  66  	switch u.Scheme {
  67  	case "https":
  68  		useTLS = true
  69  	case "http":
  70  	default:
  71  		return 0, nil, nil, fmt.Errorf("scheme %q not allowed", u.Scheme)
  72  	}
  73  	host := u.Hostname()
  74  	port := u.Port()
  75  	if port == "" {
  76  		if useTLS {
  77  			port = "443"
  78  		} else {
  79  			port = "80"
  80  		}
  81  	}
  82  	ip := host
  83  	if net.ParseIP(host) == nil {
  84  		ip, err9 = ws.ResolveHost(host)
  85  		if err9 != nil {
  86  			return 0, nil, nil, fmt.Errorf("resolve %s: %w", host, err9)
  87  		}
  88  	}
  89  	conn, err8 := net.DialTimeout("tcp", net.JoinHostPort(ip, port), DefaultTimeout)
  90  	if err8 != nil {
  91  		return 0, nil, nil, fmt.Errorf("dial: %w", err8)
  92  	}
  93  	defer conn.Close()
  94  	deadline := time.Now().Add(DefaultTimeout)
  95  	conn.SetDeadline(deadline)
  96  	if useTLS {
  97  		tlsConn := tls.Client(conn, &tls.Config{ServerName: []byte(host)})
  98  		if err4 := tlsConn.Handshake(); err4 != nil {
  99  			return 0, nil, nil, fmt.Errorf("tls: %w", err4)
 100  		}
 101  		conn = tlsConn
 102  	}
 103  	path := u.RequestURI()
 104  	if path == "" {
 105  		path = "/"
 106  	}
 107  	req := "GET " | path | " HTTP/1.1\r\n" |
 108  		"Host: " | host | "\r\n" |
 109  		"User-Agent: musiquay-mediaproxy/1\r\n" |
 110  		"Accept: image/*,video/*,*/*;q=0.5\r\n" |
 111  		"Accept-Encoding: identity\r\n" |
 112  		"Connection: close\r\n" |
 113  		"\r\n"
 114  	if _, err7 := conn.Write([]byte(req)); err7 != nil {
 115  		return 0, nil, nil, fmt.Errorf("write: %w", err7)
 116  	}
 117  	br := bufio.NewReaderSize(conn, 32768)
 118  	statusLine, err6 := br.ReadString('\n')
 119  	if err6 != nil {
 120  		return 0, nil, nil, fmt.Errorf("read status: %w", err6)
 121  	}
 122  	status, err5 := parseStatus(statusLine)
 123  	if err5 != nil {
 124  		return 0, nil, nil, err5
 125  	}
 126  	headers := map[string]string{}
 127  	for {
 128  		line, err3 := br.ReadString('\n')
 129  		if err3 != nil {
 130  			return 0, nil, nil, fmt.Errorf("read header: %w", err3)
 131  		}
 132  		trimmed := bytes.TrimRight(line, "\r\n")
 133  		if len(trimmed) == 0 {
 134  			break
 135  		}
 136  		col := bytes.IndexByte(trimmed, ':')
 137  		if col < 0 {
 138  			continue
 139  		}
 140  		k := string(bytes.ToLower(bytes.TrimSpace(trimmed[:col])))
 141  		v := string(bytes.TrimSpace(trimmed[col+1:]))
 142  		headers[k] = v
 143  	}
 144  	// Redirects don't have meaningful bodies; return early.
 145  	if status >= 300 && status < 400 {
 146  		return status, headers, nil, nil
 147  	}
 148  	var body []byte
 149  	te := headers["transfer-encoding"]
 150  	if te != "" && bytes.Contains(bytes.ToLower([]byte(te)), "chunked") {
 151  		body, err5 = readChunked(br, maxBytes)
 152  		if err5 != nil {
 153  			return 0, nil, nil, err5
 154  		}
 155  	} else if cl := headers["content-length"]; cl != "" {
 156  		n, err2 := strconv.ParseInt(cl, 10, 64)
 157  		if err2 != nil || n < 0 {
 158  			return 0, nil, nil, fmt.Errorf("bad content-length")
 159  		}
 160  		if n > maxBytes {
 161  			return 0, nil, nil, fmt.Errorf("response too large: %d", n)
 162  		}
 163  		body = []byte{:n}
 164  		if _, err1 := io.ReadFull(br, body); err1 != nil {
 165  			return 0, nil, nil, fmt.Errorf("read body: %w", err1)
 166  		}
 167  	} else {
 168  		var rerr error
 169  		body, rerr = readToEOF(br, maxBytes)
 170  		if rerr != nil {
 171  			return 0, nil, nil, rerr
 172  		}
 173  	}
 174  	return status, headers, body, nil
 175  }
 176  
 177  func parseStatus(line string) (code int32, err error) {
 178  	trimmed := bytes.TrimSpace(line)
 179  	sp1 := bytes.IndexByte(trimmed, ' ')
 180  	if sp1 < 0 {
 181  		return 0, fmt.Errorf("bad status line: %s", trimmed)
 182  	}
 183  	rest := bytes.TrimSpace(trimmed[sp1+1:])
 184  	sp2 := bytes.IndexByte(rest, ' ')
 185  	var statusBytes []byte
 186  	if sp2 < 0 {
 187  		statusBytes = rest
 188  	} else {
 189  		statusBytes = rest[:sp2]
 190  	}
 191  	n, err := strconv.Atoi(string(statusBytes))
 192  	if err != nil {
 193  		return 0, fmt.Errorf("bad status code: %s", statusBytes)
 194  	}
 195  	return n, nil
 196  }
 197  
 198  func readToEOF(r io.Reader, maxBytes int64) (out []byte, rerr error) {
 199  	var buf []byte
 200  	chunk := []byte{:32 * 1024}
 201  	for {
 202  		n, err := r.Read(chunk)
 203  		if n > 0 {
 204  			if int64(len(buf))+int64(n) > maxBytes {
 205  				return nil, fmt.Errorf("response too large")
 206  			}
 207  			buf = buf | chunk[:n]
 208  		}
 209  		if err == io.EOF {
 210  			return buf, nil
 211  		}
 212  		if err != nil {
 213  			return nil, err
 214  		}
 215  	}
 216  }
 217  
 218  func readChunked(br *bufio.Reader, maxBytes int64) (out []byte, rerr error) {
 219  	var buf []byte
 220  	for {
 221  		line, err5 := br.ReadString('\n')
 222  		if err5 != nil {
 223  			return nil, fmt.Errorf("chunked size: %w", err5)
 224  		}
 225  		sz := bytes.TrimRight(line, "\r\n")
 226  		if sc := bytes.IndexByte(sz, ';'); sc >= 0 {
 227  			sz = sz[:sc]
 228  		}
 229  		n, err4 := strconv.ParseInt(string(bytes.TrimSpace(sz)), 16, 64)
 230  		if err4 != nil {
 231  			return nil, fmt.Errorf("chunked size parse: %w", err4)
 232  		}
 233  		if n == 0 {
 234  			// Discard trailers up to empty line.
 235  			for {
 236  				t, err1 := br.ReadString('\n')
 237  				if err1 != nil {
 238  					break
 239  				}
 240  				if len(bytes.TrimRight(t, "\r\n")) == 0 {
 241  					break
 242  				}
 243  			}
 244  			return buf, nil
 245  		}
 246  		if int64(len(buf))+n > maxBytes {
 247  			return nil, fmt.Errorf("response too large")
 248  		}
 249  		chunk := []byte{:n}
 250  		if _, err3 := io.ReadFull(br, chunk); err3 != nil {
 251  			return nil, fmt.Errorf("chunked body: %w", err3)
 252  		}
 253  		buf = buf | chunk
 254  		if _, err2 := br.ReadString('\n'); err2 != nil {
 255  			return nil, fmt.Errorf("chunked trailer: %w", err2)
 256  		}
 257  	}
 258  }
 259  
 260  func resolveRedirect(base, loc string) (target string, derr error) {
 261  	if len(loc) == 0 {
 262  		return "", fmt.Errorf("empty location")
 263  	}
 264  	// Absolute URL with scheme.
 265  	if len(loc) >= 7 && (string(loc[:7]) == "http://" ||
 266  		(len(loc) >= 8 && string(loc[:8]) == "https://")) {
 267  		return loc, nil
 268  	}
 269  	bu, err := url.Parse(base)
 270  	if err != nil {
 271  		return "", err
 272  	}
 273  	// Network-path reference: "//host/path"
 274  	if len(loc) >= 2 && loc[0] == '/' && loc[1] == '/' {
 275  		return bu.Scheme | ":" | loc, nil
 276  	}
 277  	// Absolute path: "/path"
 278  	if loc[0] == '/' {
 279  		return bu.Scheme | "://" | bu.Host | loc, nil
 280  	}
 281  	// Relative path - resolve against base path.
 282  	basePath := bu.Path
 283  	if basePath == "" {
 284  		basePath = "/"
 285  	}
 286  	slash := -1
 287  	for i := len(basePath) - 1; i >= 0; i-- {
 288  		if basePath[i] == '/' {
 289  			slash = i
 290  			break
 291  		}
 292  	}
 293  	if slash < 0 {
 294  		basePath = "/"
 295  	} else {
 296  		basePath = basePath[:slash+1]
 297  	}
 298  	return bu.Scheme | "://" | bu.Host | basePath | loc, nil
 299  }
 300