577 lines
15 KiB
Go
577 lines
15 KiB
Go
package server
|
|
|
|
import (
|
|
"bufio"
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"net/http"
|
|
"slices"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"time"
|
|
|
|
proxyDomain "github.com/proxy-pool/proxy-pool/internal/domain/proxy"
|
|
"github.com/proxy-pool/proxy-pool/internal/gateway/dispatch"
|
|
"github.com/proxy-pool/proxy-pool/internal/gateway/policy"
|
|
transportDomain "github.com/proxy-pool/proxy-pool/internal/gateway/transport"
|
|
"github.com/proxy-pool/proxy-pool/internal/platform/httpsecurity"
|
|
)
|
|
|
|
type Config struct {
|
|
MaxAttempts int
|
|
RetryMethods []string
|
|
SafetyMargin time.Duration
|
|
CopyBufferSize int
|
|
MaxConcurrentRequests int
|
|
}
|
|
|
|
const requestClosingBit uint64 = 1 << 63
|
|
|
|
type Guard interface {
|
|
Check(context.Context, *http.Request) error
|
|
}
|
|
|
|
type GuardFunc func(context.Context, *http.Request) error
|
|
|
|
func (guard GuardFunc) Check(ctx context.Context, request *http.Request) error {
|
|
return guard(ctx, request)
|
|
}
|
|
|
|
type TargetPolicy interface {
|
|
EvaluateURL(context.Context, string) (policy.Authority, error)
|
|
EvaluateConnectAuthority(context.Context, string) (policy.Authority, error)
|
|
}
|
|
|
|
type Router interface {
|
|
Route(*http.Request) (dispatch.Request, error)
|
|
}
|
|
|
|
type RouteFunc func(*http.Request) (dispatch.Request, error)
|
|
|
|
func (route RouteFunc) Route(request *http.Request) (dispatch.Request, error) {
|
|
return route(request)
|
|
}
|
|
|
|
type Dispatcher interface {
|
|
Acquire(dispatch.Request) (*dispatch.Lease, error)
|
|
}
|
|
|
|
type DispatcherFunc func(dispatch.Request) (*dispatch.Lease, error)
|
|
|
|
func (acquire DispatcherFunc) Acquire(request dispatch.Request) (*dispatch.Lease, error) {
|
|
return acquire(request)
|
|
}
|
|
|
|
type ProxyTransport interface {
|
|
RoundTrip(context.Context, proxyDomain.Proxy, *http.Request, ...func() error) (*http.Response, error)
|
|
OpenTunnel(context.Context, proxyDomain.Proxy, string) (net.Conn, error)
|
|
Relay(context.Context, net.Conn, net.Conn) error
|
|
}
|
|
|
|
type Dependencies struct {
|
|
Auth Guard
|
|
Access Guard
|
|
Admission Guard
|
|
Targets TargetPolicy
|
|
Router Router
|
|
Dispatcher Dispatcher
|
|
Transport ProxyTransport
|
|
}
|
|
|
|
type Handler struct {
|
|
config Config
|
|
guards [3]Guard
|
|
targets TargetPolicy
|
|
router Router
|
|
dispatcher Dispatcher
|
|
transport ProxyTransport
|
|
buffers sync.Pool
|
|
inFlight chan struct{}
|
|
forceClose atomic.Bool
|
|
requestState atomic.Uint64
|
|
shutdownStart sync.Once
|
|
drainOnce sync.Once
|
|
shutdownDone chan struct{}
|
|
tunnelMu sync.Mutex
|
|
tunnels map[*activeTunnel]struct{}
|
|
}
|
|
|
|
func New(config Config, dependencies Dependencies) (*Handler, error) {
|
|
if dependencies.Targets == nil {
|
|
return nil, errors.New("create gateway handler: target policy is required")
|
|
}
|
|
if dependencies.Router == nil {
|
|
return nil, errors.New("create gateway handler: router is required")
|
|
}
|
|
if dependencies.Dispatcher == nil {
|
|
return nil, errors.New("create gateway handler: dispatcher is required")
|
|
}
|
|
if dependencies.Transport == nil {
|
|
return nil, errors.New("create gateway handler: transport is required")
|
|
}
|
|
if config.MaxAttempts <= 0 {
|
|
config.MaxAttempts = 1
|
|
}
|
|
if config.CopyBufferSize <= 0 {
|
|
config.CopyBufferSize = 32 << 10
|
|
}
|
|
handler := &Handler{
|
|
config: config,
|
|
guards: [3]Guard{dependencies.Auth, dependencies.Access, dependencies.Admission},
|
|
targets: dependencies.Targets,
|
|
router: dependencies.Router,
|
|
dispatcher: dependencies.Dispatcher,
|
|
transport: dependencies.Transport,
|
|
tunnels: make(map[*activeTunnel]struct{}),
|
|
shutdownDone: make(chan struct{}),
|
|
}
|
|
if config.MaxConcurrentRequests > 0 {
|
|
handler.inFlight = make(chan struct{}, config.MaxConcurrentRequests)
|
|
}
|
|
handler.buffers.New = func() any { return make([]byte, config.CopyBufferSize) }
|
|
return handler, nil
|
|
}
|
|
|
|
func (handler *Handler) ServeHTTP(writer http.ResponseWriter, request *http.Request) {
|
|
if !handler.beginRequest() {
|
|
http.Error(writer, http.StatusText(http.StatusServiceUnavailable), http.StatusServiceUnavailable)
|
|
return
|
|
}
|
|
defer handler.finishRequest()
|
|
if handler.inFlight != nil {
|
|
select {
|
|
case handler.inFlight <- struct{}{}:
|
|
defer func() { <-handler.inFlight }()
|
|
default:
|
|
http.Error(writer, http.StatusText(http.StatusServiceUnavailable), http.StatusServiceUnavailable)
|
|
return
|
|
}
|
|
}
|
|
for _, guard := range handler.guards {
|
|
if guard == nil {
|
|
continue
|
|
}
|
|
if err := guard.Check(request.Context(), request); err != nil {
|
|
writeGatewayError(writer, err)
|
|
return
|
|
}
|
|
}
|
|
|
|
if request.Method == http.MethodConnect {
|
|
target, err := handler.targets.EvaluateConnectAuthority(request.Context(), request.Host)
|
|
if err != nil {
|
|
writeGatewayError(writer, err)
|
|
return
|
|
}
|
|
route, err := handler.router.Route(request)
|
|
if err != nil {
|
|
writeGatewayError(writer, fmt.Errorf("route gateway CONNECT: %w", err))
|
|
return
|
|
}
|
|
handler.connect(writer, request, target, route)
|
|
return
|
|
}
|
|
if request.URL == nil || !request.URL.IsAbs() {
|
|
http.Error(writer, http.StatusText(http.StatusBadRequest), http.StatusBadRequest)
|
|
return
|
|
}
|
|
target, err := handler.targets.EvaluateURL(request.Context(), request.URL.String())
|
|
if err != nil {
|
|
writeGatewayError(writer, err)
|
|
return
|
|
}
|
|
route, err := handler.router.Route(request)
|
|
if err != nil {
|
|
writeGatewayError(writer, fmt.Errorf("route gateway request: %w", err))
|
|
return
|
|
}
|
|
handler.forwardHTTP(writer, request, target, route)
|
|
}
|
|
|
|
func (handler *Handler) connect(
|
|
writer http.ResponseWriter,
|
|
request *http.Request,
|
|
target policy.Authority,
|
|
route dispatch.Request,
|
|
) {
|
|
attempts := handler.attemptLimit(request)
|
|
excluded := cloneSet(route.Exclude)
|
|
var lastErr error
|
|
|
|
for attempt := 0; attempt < attempts; attempt++ {
|
|
route.Now = time.Now().UTC()
|
|
route.Exclude = excluded
|
|
route.SafetyMargin = handler.config.SafetyMargin
|
|
lease, err := handler.dispatcher.Acquire(route)
|
|
if err != nil {
|
|
lastErr = err
|
|
break
|
|
}
|
|
|
|
upstream, err := handler.transport.OpenTunnel(request.Context(), lease.Proxy, target.DialAddress())
|
|
if err != nil {
|
|
finishLease(lease, false)
|
|
excluded[lease.Proxy.ID] = struct{}{}
|
|
var responseError *transportDomain.ProxyResponseError
|
|
if errors.As(err, &responseError) {
|
|
lastErr = responseError
|
|
if responseError.Retryable() && attempt+1 < attempts {
|
|
continue
|
|
}
|
|
handler.writeConnectError(writer, responseError)
|
|
return
|
|
}
|
|
lastErr = err
|
|
continue
|
|
}
|
|
|
|
if err := lease.Commit(); err != nil {
|
|
_ = upstream.Close()
|
|
finishLease(lease, false)
|
|
lastErr = err
|
|
break
|
|
}
|
|
handler.serveTunnel(writer, request, lease, upstream)
|
|
return
|
|
}
|
|
|
|
writeGatewayError(writer, fmt.Errorf("establish gateway CONNECT: %w", lastErr))
|
|
}
|
|
|
|
func (handler *Handler) writeConnectError(writer http.ResponseWriter, responseError *transportDomain.ProxyResponseError) {
|
|
header := responseError.Header.Clone()
|
|
removeHopByHop(header)
|
|
copyHeaders(writer.Header(), header)
|
|
writer.WriteHeader(responseError.StatusCode)
|
|
_, _ = writer.Write(responseError.Body)
|
|
}
|
|
|
|
func (handler *Handler) serveTunnel(
|
|
writer http.ResponseWriter,
|
|
request *http.Request,
|
|
lease *dispatch.Lease,
|
|
upstream net.Conn,
|
|
) {
|
|
defer finishLease(lease, true)
|
|
defer upstream.Close()
|
|
|
|
hijacker, ok := writer.(http.Hijacker)
|
|
if !ok {
|
|
http.Error(writer, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
client, readWriter, err := hijacker.Hijack()
|
|
if err != nil {
|
|
http.Error(writer, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError)
|
|
return
|
|
}
|
|
defer client.Close()
|
|
tunnel := &activeTunnel{client: client, upstream: upstream}
|
|
if !handler.registerTunnel(tunnel) {
|
|
return
|
|
}
|
|
defer handler.unregisterTunnel(tunnel)
|
|
|
|
if _, err := readWriter.WriteString("HTTP/1.1 200 Connection Established\r\n\r\n"); err != nil {
|
|
return
|
|
}
|
|
if err := readWriter.Flush(); err != nil {
|
|
return
|
|
}
|
|
bufferedClient := &bufferedClientConn{Conn: client, reader: readWriter.Reader}
|
|
_ = handler.transport.Relay(request.Context(), bufferedClient, upstream)
|
|
}
|
|
|
|
func (handler *Handler) Shutdown(ctx context.Context) error {
|
|
if handler == nil {
|
|
return nil
|
|
}
|
|
handler.shutdownStart.Do(func() {
|
|
for {
|
|
state := handler.requestState.Load()
|
|
if handler.requestState.CompareAndSwap(state, state|requestClosingBit) {
|
|
if state == 0 {
|
|
handler.signalDrained()
|
|
}
|
|
break
|
|
}
|
|
}
|
|
if closer, ok := handler.transport.(interface{ CloseIdleConnections() }); ok {
|
|
closer.CloseIdleConnections()
|
|
}
|
|
})
|
|
select {
|
|
case <-handler.shutdownDone:
|
|
return nil
|
|
case <-ctx.Done():
|
|
handler.forceClose.Store(true)
|
|
handler.closeActiveTunnels()
|
|
return ctx.Err()
|
|
}
|
|
}
|
|
|
|
func (handler *Handler) beginRequest() bool {
|
|
for {
|
|
state := handler.requestState.Load()
|
|
if state&requestClosingBit != 0 || state == requestClosingBit-1 {
|
|
return false
|
|
}
|
|
if handler.requestState.CompareAndSwap(state, state+1) {
|
|
return true
|
|
}
|
|
}
|
|
}
|
|
|
|
func (handler *Handler) finishRequest() {
|
|
state := handler.requestState.Add(^uint64(0))
|
|
if state == requestClosingBit {
|
|
handler.signalDrained()
|
|
}
|
|
}
|
|
|
|
func (handler *Handler) signalDrained() {
|
|
handler.drainOnce.Do(func() { close(handler.shutdownDone) })
|
|
}
|
|
|
|
func (handler *Handler) closeActiveTunnels() {
|
|
handler.tunnelMu.Lock()
|
|
tunnels := make([]*activeTunnel, 0, len(handler.tunnels))
|
|
for tunnel := range handler.tunnels {
|
|
tunnels = append(tunnels, tunnel)
|
|
}
|
|
handler.tunnelMu.Unlock()
|
|
for _, tunnel := range tunnels {
|
|
_ = tunnel.client.Close()
|
|
_ = tunnel.upstream.Close()
|
|
}
|
|
}
|
|
|
|
func (handler *Handler) registerTunnel(tunnel *activeTunnel) bool {
|
|
handler.tunnelMu.Lock()
|
|
defer handler.tunnelMu.Unlock()
|
|
if handler.forceClose.Load() {
|
|
_ = tunnel.client.Close()
|
|
_ = tunnel.upstream.Close()
|
|
return false
|
|
}
|
|
handler.tunnels[tunnel] = struct{}{}
|
|
return true
|
|
}
|
|
|
|
func (handler *Handler) unregisterTunnel(tunnel *activeTunnel) {
|
|
handler.tunnelMu.Lock()
|
|
delete(handler.tunnels, tunnel)
|
|
handler.tunnelMu.Unlock()
|
|
}
|
|
|
|
func (handler *Handler) forwardHTTP(
|
|
writer http.ResponseWriter,
|
|
request *http.Request,
|
|
target policy.Authority,
|
|
route dispatch.Request,
|
|
) {
|
|
attempts := handler.attemptLimit(request)
|
|
excluded := cloneSet(route.Exclude)
|
|
var lastErr error
|
|
|
|
for attempt := 0; attempt < attempts; attempt++ {
|
|
attemptRequest, err := requestForAttempt(request, attempt)
|
|
if err != nil {
|
|
lastErr = err
|
|
break
|
|
}
|
|
pinHTTPDestination(attemptRequest, target)
|
|
removeHopByHop(attemptRequest.Header)
|
|
|
|
route.Now = time.Now().UTC()
|
|
route.Exclude = excluded
|
|
route.SafetyMargin = handler.config.SafetyMargin
|
|
lease, err := handler.dispatcher.Acquire(route)
|
|
if err != nil {
|
|
lastErr = err
|
|
break
|
|
}
|
|
|
|
var committed atomic.Bool
|
|
commit := func() error {
|
|
if err := lease.Commit(); err != nil {
|
|
return err
|
|
}
|
|
committed.Store(true)
|
|
return nil
|
|
}
|
|
response, err := handler.transport.RoundTrip(request.Context(), lease.Proxy, attemptRequest, commit)
|
|
if err != nil {
|
|
finishLease(lease, committed.Load())
|
|
excluded[lease.Proxy.ID] = struct{}{}
|
|
lastErr = err
|
|
continue
|
|
}
|
|
if !committed.Load() {
|
|
if err := commit(); err != nil {
|
|
_ = response.Body.Close()
|
|
finishLease(lease, false)
|
|
lastErr = err
|
|
break
|
|
}
|
|
}
|
|
handler.writeResponse(writer, response)
|
|
finishLease(lease, true)
|
|
return
|
|
}
|
|
|
|
writeGatewayError(writer, fmt.Errorf("forward gateway request: %w", lastErr))
|
|
}
|
|
|
|
func pinHTTPDestination(request *http.Request, target policy.Authority) {
|
|
if request == nil || request.URL == nil || (!target.ResolvedIP.IsValid() && !target.LiteralIP.IsValid()) {
|
|
return
|
|
}
|
|
request.URL.Host = target.DialAddress()
|
|
}
|
|
|
|
func (handler *Handler) writeResponse(writer http.ResponseWriter, response *http.Response) {
|
|
defer response.Body.Close()
|
|
removeHopByHop(response.Header)
|
|
copyHeaders(writer.Header(), response.Header)
|
|
writer.WriteHeader(response.StatusCode)
|
|
buffer := handler.buffers.Get().([]byte)
|
|
defer handler.buffers.Put(buffer)
|
|
_, _ = io.CopyBuffer(writer, response.Body, buffer)
|
|
}
|
|
|
|
func (handler *Handler) attemptLimit(request *http.Request) int {
|
|
if handler.config.MaxAttempts <= 1 || !containsMethod(handler.config.RetryMethods, request.Method) {
|
|
return 1
|
|
}
|
|
if request.Body != nil && request.Body != http.NoBody && request.GetBody == nil {
|
|
return 1
|
|
}
|
|
return handler.config.MaxAttempts
|
|
}
|
|
|
|
func requestForAttempt(request *http.Request, attempt int) (*http.Request, error) {
|
|
clone := request.Clone(request.Context())
|
|
clone.RequestURI = ""
|
|
clone.Header = request.Header.Clone()
|
|
if attempt == 0 || request.Body == nil || request.Body == http.NoBody {
|
|
return clone, nil
|
|
}
|
|
body, err := request.GetBody()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("replay gateway request body: %w", err)
|
|
}
|
|
clone.Body = body
|
|
return clone, nil
|
|
}
|
|
|
|
func finishLease(lease *dispatch.Lease, committed bool) {
|
|
if committed {
|
|
_ = lease.Release()
|
|
return
|
|
}
|
|
_ = lease.Cancel()
|
|
}
|
|
|
|
func containsMethod(methods []string, method string) bool {
|
|
return slices.ContainsFunc(methods, func(candidate string) bool {
|
|
return strings.EqualFold(candidate, method)
|
|
})
|
|
}
|
|
|
|
func cloneSet(source map[string]struct{}) map[string]struct{} {
|
|
cloned := make(map[string]struct{}, len(source)+1)
|
|
for key := range source {
|
|
cloned[key] = struct{}{}
|
|
}
|
|
return cloned
|
|
}
|
|
|
|
func copyHeaders(destination, source http.Header) {
|
|
for name, values := range source {
|
|
for _, value := range values {
|
|
destination.Add(name, value)
|
|
}
|
|
}
|
|
}
|
|
|
|
func removeHopByHop(header http.Header) {
|
|
for _, value := range header.Values("Connection") {
|
|
for name := range strings.SplitSeq(value, ",") {
|
|
header.Del(strings.TrimSpace(name))
|
|
}
|
|
}
|
|
for _, name := range []string{
|
|
"Connection",
|
|
"Keep-Alive",
|
|
"Proxy-Connection",
|
|
"TE",
|
|
"Trailer",
|
|
"Transfer-Encoding",
|
|
"Upgrade",
|
|
} {
|
|
header.Del(name)
|
|
}
|
|
}
|
|
|
|
func writeGatewayError(writer http.ResponseWriter, err error) {
|
|
status := http.StatusBadGateway
|
|
var httpError *HTTPError
|
|
if errors.As(err, &httpError) {
|
|
status = httpError.StatusCode
|
|
copyHeaders(writer.Header(), httpError.Header)
|
|
} else {
|
|
var securityError *httpsecurity.HTTPError
|
|
if errors.As(err, &securityError) {
|
|
status = securityError.StatusCode
|
|
copyHeaders(writer.Header(), securityError.Header)
|
|
} else {
|
|
switch {
|
|
case errors.Is(err, policy.ErrInvalidAuthority):
|
|
status = http.StatusBadRequest
|
|
case errors.Is(err, policy.ErrTargetDenied):
|
|
status = http.StatusForbidden
|
|
case errors.Is(err, ErrRouteRejected):
|
|
status = http.StatusForbidden
|
|
case errors.Is(err, dispatch.ErrNoCandidate), errors.Is(err, ErrRouteNotFound):
|
|
status = http.StatusServiceUnavailable
|
|
case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded):
|
|
status = http.StatusGatewayTimeout
|
|
}
|
|
}
|
|
}
|
|
http.Error(writer, http.StatusText(status), status)
|
|
}
|
|
|
|
type bufferedClientConn struct {
|
|
net.Conn
|
|
reader *bufio.Reader
|
|
}
|
|
|
|
type activeTunnel struct {
|
|
client net.Conn
|
|
upstream net.Conn
|
|
}
|
|
|
|
func (connection *bufferedClientConn) Read(buffer []byte) (int, error) {
|
|
return connection.reader.Read(buffer)
|
|
}
|
|
|
|
func (connection *bufferedClientConn) CloseWrite() error {
|
|
if halfCloser, ok := connection.Conn.(interface{ CloseWrite() error }); ok {
|
|
return halfCloser.CloseWrite()
|
|
}
|
|
return connection.Close()
|
|
}
|
|
|
|
func (connection *bufferedClientConn) CloseRead() error {
|
|
if halfCloser, ok := connection.Conn.(interface{ CloseRead() error }); ok {
|
|
return halfCloser.CloseRead()
|
|
}
|
|
return nil
|
|
}
|