package server import ( "bufio" "context" "errors" "fmt" "io" "net" "net/http" "slices" "strings" "sync" "sync/atomic" "time" proxyDomain "proxy-pool/internal/domain/proxy" "proxy-pool/internal/domain/routing" "proxy-pool/internal/gateway/dispatch" "proxy-pool/internal/gateway/policy" transportDomain "proxy-pool/internal/gateway/transport" "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 DirectTransport interface { RoundTripDirect(context.Context, *http.Request) (*http.Response, error) OpenDirectTunnel(context.Context, string) (net.Conn, error) } type waitingDispatcher interface { AcquireWait(context.Context, dispatch.Request, time.Duration) (*dispatch.Lease, 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.acquireRoute(request.Context(), route) if err != nil { lastErr = err if errors.Is(err, dispatch.ErrNoCandidate) && route.OnUnavailable == routing.OnUnavailableDirect { direct, directErr := handler.directTransport() if directErr != nil { lastErr = directErr break } upstream, directErr := direct.OpenDirectTunnel(request.Context(), target.DialAddress()) if directErr != nil { lastErr = directErr break } handler.serveTunnel(writer, request, nil, upstream) return } 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, func() { finishLease(lease, true) }, 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, finish func(), upstream net.Conn, ) { if finish != nil { defer finish() } 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.acquireRoute(request.Context(), route) if err != nil { lastErr = err if errors.Is(err, dispatch.ErrNoCandidate) && route.OnUnavailable == routing.OnUnavailableDirect { direct, directErr := handler.directTransport() if directErr != nil { lastErr = directErr break } response, directErr := direct.RoundTripDirect(request.Context(), attemptRequest) if directErr != nil { lastErr = directErr break } handler.writeResponse(writer, response) return } 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 (handler *Handler) acquireRoute(ctx context.Context, route dispatch.Request) (*dispatch.Lease, error) { lease, err := handler.dispatcher.Acquire(route) if !errors.Is(err, dispatch.ErrNoCandidate) || route.OnUnavailable != routing.OnUnavailableWait || route.WaitTimeout <= 0 { return lease, err } waiter, ok := handler.dispatcher.(waitingDispatcher) if !ok { return nil, err } return waiter.AcquireWait(ctx, route, route.WaitTimeout) } func (handler *Handler) directTransport() (DirectTransport, error) { direct, ok := handler.transport.(DirectTransport) if !ok { return nil, ErrDirectRouteUnsupported } return direct, nil } 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 }