package server import ( "bufio" "context" "errors" "fmt" "io" "net" "net/http" "slices" "strings" "sync" "sync/atomic" "time" outcomeDomain "proxy-pool/internal/domain/outcome" 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 StickySession StickySessionConfig } 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) } // OutcomeRecorder accepts best-effort proxy observations. Implementations must // return immediately; request forwarding never waits for telemetry delivery. type OutcomeRecorder interface { Record(outcomeDomain.Event) } // RequestMetricsObserver records local Gateway request lifecycle signals. // Implementations must be concurrency-safe and must not block forwarding. type RequestMetricsObserver interface { ObserveRequestStarted(protocol string) ObserveRequestFinished(protocol string) ObserveTunnelOpened() ObserveTunnelClosed() } 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 Outcomes OutcomeRecorder Metrics RequestMetricsObserver } type Handler struct { config Config guards [3]Guard targets TargetPolicy router Router dispatcher Dispatcher transport ProxyTransport outcomes OutcomeRecorder metrics RequestMetricsObserver sticky *stickySession 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 } sticky, err := newStickySession(config.StickySession) if err != nil { return nil, fmt.Errorf("create gateway handler: %w", err) } 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, outcomes: dependencies.Outcomes, metrics: dependencies.Metrics, sticky: sticky, 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() protocol := requestMetricProtocol(request) handler.observeRequestStarted(protocol) defer handler.observeRequestFinished(protocol) 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 } } defer releaseCredentialReservation(request) 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 } if err := authorizeClientRouting(request, route); err != nil { writeGatewayError(writer, err) return } if route.Action == routing.ActionDirect { handler.connectDirect(writer, request, target) return } binding, err := handler.prepareStickySession(request, &route) if err != nil { writeGatewayError(writer, err) return } handler.connect(writer, request, target, route, binding) 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 } if err := authorizeClientRouting(request, route); err != nil { writeGatewayError(writer, err) return } if route.Action == routing.ActionDirect { handler.forwardHTTPDirect(writer, request, target) return } binding, err := handler.prepareStickySession(request, &route) if err != nil { writeGatewayError(writer, err) return } handler.forwardHTTP(writer, request, target, route, binding) } func (handler *Handler) connect( writer http.ResponseWriter, request *http.Request, target policy.Authority, route dispatch.Request, binding *stickyBinding, ) { 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.acquireRouteWithStickySession(request.Context(), &route, binding) if err != nil { lastErr = err if errors.Is(err, dispatch.ErrNoCandidate) && route.OnUnavailable == routing.OnUnavailableDirect { upstream, directErr := handler.openDirectTunnel(request.Context(), target.DialAddress()) if directErr != nil { lastErr = directErr break } handler.serveTunnel(writer, request, nil, upstream, "", "") return } break } started := time.Now().UTC() upstream, err := handler.transport.OpenTunnel(request.Context(), lease.Proxy, target.DialAddress()) if err != nil { handler.recordOutcome(lease.Proxy.ID, route.RoutingName, outcomeDomain.StageProxyHandshake, false, err, started) finishLease(lease, false) binding.Clear() route.PreferredProxyID = "" 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) binding.Clear() route.PreferredProxyID = "" lastErr = err break } binding.Bind(lease.Proxy, handler.config.SafetyMargin) handler.recordOutcome(lease.Proxy.ID, route.RoutingName, outcomeDomain.StageProxyHandshake, true, nil, started) handler.serveTunnel(writer, request, func() { finishLease(lease, true) }, upstream, lease.Proxy.ID, route.RoutingName) return } writeGatewayError(writer, fmt.Errorf("establish gateway CONNECT: %w", lastErr)) } func (handler *Handler) connectDirect(writer http.ResponseWriter, request *http.Request, target policy.Authority) { upstream, err := handler.openDirectTunnel(request.Context(), target.DialAddress()) if err != nil { writeGatewayError(writer, fmt.Errorf("establish direct CONNECT: %w", err)) return } handler.serveTunnel(writer, request, nil, upstream, "", "") } 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, proxyID string, routingName string, ) { if finish != nil { defer finish() } defer upstream.Close() hijacker, ok := writer.(http.Hijacker) if !ok { handler.recordOutcome(proxyID, routingName, outcomeDomain.StageTunnel, false, errors.New("gateway response does not support connection hijacking"), time.Now().UTC()) http.Error(writer, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) return } client, readWriter, err := hijacker.Hijack() if err != nil { handler.recordOutcome(proxyID, routingName, outcomeDomain.StageTunnel, false, err, time.Now().UTC()) 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) started := time.Now().UTC() if _, err := readWriter.WriteString("HTTP/1.1 200 Connection Established\r\n\r\n"); err != nil { handler.recordOutcome(proxyID, routingName, outcomeDomain.StageTunnel, false, err, started) return } if err := readWriter.Flush(); err != nil { handler.recordOutcome(proxyID, routingName, outcomeDomain.StageTunnel, false, err, started) return } bufferedClient := &bufferedClientConn{Conn: client, reader: readWriter.Reader} err = handler.transport.Relay(request.Context(), bufferedClient, upstream) handler.recordOutcome(proxyID, routingName, outcomeDomain.StageTunnel, err == nil, err, started) } 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{}{} if handler.metrics != nil { handler.metrics.ObserveTunnelOpened() } return true } func (handler *Handler) unregisterTunnel(tunnel *activeTunnel) { handler.tunnelMu.Lock() _, exists := handler.tunnels[tunnel] if exists { delete(handler.tunnels, tunnel) } handler.tunnelMu.Unlock() if exists && handler.metrics != nil { handler.metrics.ObserveTunnelClosed() } } func (handler *Handler) observeRequestStarted(protocol string) { if handler.metrics != nil { handler.metrics.ObserveRequestStarted(protocol) } } func (handler *Handler) observeRequestFinished(protocol string) { if handler.metrics != nil { handler.metrics.ObserveRequestFinished(protocol) } } func requestMetricProtocol(request *http.Request) string { if request != nil && request.Method == http.MethodConnect { return "CONNECT" } return "HTTP" } func (handler *Handler) forwardHTTP( writer http.ResponseWriter, request *http.Request, target policy.Authority, route dispatch.Request, binding *stickyBinding, ) { 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.acquireRouteWithStickySession(request.Context(), &route, binding) if err != nil { lastErr = err if errors.Is(err, dispatch.ErrNoCandidate) && route.OnUnavailable == routing.OnUnavailableDirect { response, directErr := handler.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) binding.Bind(lease.Proxy, handler.config.SafetyMargin) return nil } started := time.Now().UTC() response, err := handler.transport.RoundTrip(request.Context(), lease.Proxy, attemptRequest, commit) if err != nil { handler.recordOutcome(lease.Proxy.ID, route.RoutingName, outcomeDomain.StageDial, false, err, started) finishLease(lease, committed.Load()) binding.Clear() route.PreferredProxyID = "" excluded[lease.Proxy.ID] = struct{}{} lastErr = err continue } if !committed.Load() { if err := commit(); err != nil { _ = response.Body.Close() finishLease(lease, false) binding.Clear() route.PreferredProxyID = "" handler.recordOutcome(lease.Proxy.ID, route.RoutingName, outcomeDomain.StageResponseHeaders, false, err, started) lastErr = err break } } if response.StatusCode == http.StatusProxyAuthRequired { binding.Clear() route.PreferredProxyID = "" handler.recordOutcome(lease.Proxy.ID, route.RoutingName, outcomeDomain.StageResponseHeaders, false, &transportDomain.ProxyResponseError{StatusCode: response.StatusCode}, started) } else { handler.recordOutcome(lease.Proxy.ID, route.RoutingName, outcomeDomain.StageResponseHeaders, true, nil, started) } handler.writeResponse(writer, response) finishLease(lease, true) return } writeGatewayError(writer, fmt.Errorf("forward gateway request: %w", lastErr)) } func (handler *Handler) forwardHTTPDirect(writer http.ResponseWriter, request *http.Request, target policy.Authority) { attemptRequest, err := requestForAttempt(request, 0) if err != nil { writeGatewayError(writer, fmt.Errorf("prepare direct request: %w", err)) return } pinHTTPDestination(attemptRequest, target) removeHopByHop(attemptRequest.Header) response, err := handler.roundTripDirect(request.Context(), attemptRequest) if err != nil { writeGatewayError(writer, fmt.Errorf("forward direct request: %w", err)) return } handler.writeResponse(writer, response) } func (handler *Handler) recordOutcome( proxyID, routingName string, stage outcomeDomain.Stage, success bool, err error, started time.Time, ) { if handler == nil || handler.outcomes == nil || proxyID == "" { return } observedAt := time.Now().UTC() if started.IsZero() { started = observedAt } latency := observedAt.Sub(started) if latency < 0 { latency = 0 } event := outcomeDomain.Event{ ProxyID: proxyID, RoutingName: routingName, Stage: stage, Success: success, Latency: latency, ObservedAt: observedAt, } if !success { event.ErrorClass = classifyOutcomeError(stage, err) } handler.outcomes.Record(event) } func classifyOutcomeError(stage outcomeDomain.Stage, err error) outcomeDomain.ErrorClass { switch { case errors.Is(err, context.Canceled): return outcomeDomain.ErrorClassCanceled case errors.Is(err, context.DeadlineExceeded): return outcomeDomain.ErrorClassTimeout } var responseError *transportDomain.ProxyResponseError if errors.As(err, &responseError) { return outcomeDomain.ErrorClassProxyResponse } var networkError net.Error if errors.As(err, &networkError) && networkError.Timeout() { return outcomeDomain.ErrorClassTimeout } switch stage { case outcomeDomain.StageDial: return outcomeDomain.ErrorClassDial case outcomeDomain.StageProxyHandshake: return outcomeDomain.ErrorClassHandshake case outcomeDomain.StageTunnel: return outcomeDomain.ErrorClassRelay default: return outcomeDomain.ErrorClassInternal } } 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) prepareStickySession(request *http.Request, route *dispatch.Request) (*stickyBinding, error) { if handler == nil || route == nil || handler.sticky == nil { return nil, nil } binding, preferredProxyID, err := handler.sticky.prepare(request, route.RoutingName) if err != nil { status := http.StatusBadRequest if errors.Is(err, ErrMissingSessionClient) { status = http.StatusForbidden } return nil, &HTTPError{StatusCode: status, Cause: err} } route.PreferredProxyID = preferredProxyID return binding, nil } func (handler *Handler) acquireRouteWithStickySession( ctx context.Context, route *dispatch.Request, binding *stickyBinding, ) (*dispatch.Lease, error) { if route == nil { return nil, dispatch.ErrNoCandidate } lease, err := handler.acquireRoute(ctx, *route) if !errors.Is(err, dispatch.ErrPreferredUnavailable) || binding == nil { return lease, err } binding.Clear() route.PreferredProxyID = "" return handler.acquireRoute(ctx, *route) } func (handler *Handler) directTransport() (DirectTransport, error) { direct, ok := handler.transport.(DirectTransport) if !ok { return nil, ErrDirectRouteUnsupported } return direct, nil } func (handler *Handler) openDirectTunnel(ctx context.Context, target string) (net.Conn, error) { direct, err := handler.directTransport() if err != nil { return nil, err } return direct.OpenDirectTunnel(ctx, target) } func (handler *Handler) roundTripDirect(ctx context.Context, request *http.Request) (*http.Response, error) { direct, err := handler.directTransport() if err != nil { return nil, err } return direct.RoundTripDirect(ctx, request) } 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 }