proxy-pool/internal/gateway/server/handler.go

877 lines
25 KiB
Go

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
}