diff --git a/.github/dependabot.yml b/.github/dependabot.yml new file mode 100644 index 00000000..47f97d2f --- /dev/null +++ b/.github/dependabot.yml @@ -0,0 +1,12 @@ +version: 2 +updates: +- package-ecosystem: maven + directory: "/" + schedule: + interval: daily + open-pull-requests-limit: 10 +- package-ecosystem: "github-actions" + directory: "/" + schedule: + interval: daily + open-pull-requests-limit: 10 diff --git a/.github/workflows/main.yml b/.github/workflows/main.yml new file mode 100644 index 00000000..943b4227 --- /dev/null +++ b/.github/workflows/main.yml @@ -0,0 +1,84 @@ +# This workflow will build a Java project with Maven, and cache/restore any dependencies to improve the workflow execution time +# For more information see: https://help.github.com/actions/language-and-framework-guides/building-and-testing-java-with-maven + +name: Java CI with Maven + +on: + push: + branches: [ main ] + paths-ignore: + - '.idea/**' + - '.run/**' + - '**/*.md' + - 'src/site/**' + - '**/.editorconfig' + - '**/.gitattributes' + - '**/.gitignore' + - '/*.txt' + - '/*.bash' + - 'release' + pull_request: + paths-ignore: + - '.idea/**' + - '.run/**' + - '**/*.md' + - 'src/site/**' + - '**/.editorconfig' + - '**/.gitattributes' + - '**/.gitignore' + - '/*.txt' + - '/*.bash' + - 'release' + +concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + build: + strategy: + matrix: + platform: [ubuntu-latest, macos-latest, windows-latest] + fail-fast: false + + runs-on: ${{ matrix.platform }} + timeout-minutes: 7 + + steps: + - uses: actions/checkout@v7 + - name: Set up JDK 17 + uses: actions/setup-java@v6 + with: + java-version: '17' + distribution: 'temurin' + cache: maven + - name: Check Java + run: java -version + - name: Check Maven + run: mvn --batch-mode --version + - name: Compile code + run: mvn --batch-mode compile + - name: Branch name [env] + run: echo running on branch ${GITHUB_REF##*/} + - name: Run smoke tests + run: mvn --batch-mode package -Psmoke-test + - name: Run slow tests + if: success() + run: mvn --batch-mode package -Pslow-tests + + - name: Generate Maven test report + if: failure() + run: mvn --batch-mode surefire-report:report-only + - name: Upload test report + uses: actions/upload-artifact@v7 + if: failure() + with: + name: test-report-${{matrix.platform}} + retention-days: 5 + path: | + target/surefire-reports + target/site + target/*.hprof + target/*.txt + target/*.jks + target/*_cert* diff --git a/.gitignore b/.gitignore index eb7fa26c..146483be 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,3 @@ -log.txt *.swp *.settings *.classpath @@ -9,10 +8,13 @@ log.txt *.DS_Store *.orig target/ -.idea/ +dependency-reduced-pom.xml jmeter.log lib/ LittleProxy.pro /bin -/*.jks performance/site/ +.claude/settings.local.json +/logs +/.opencode +/.sisyphus \ No newline at end of file diff --git a/.idea/.gitignore b/.idea/.gitignore new file mode 100644 index 00000000..f46b8fe8 --- /dev/null +++ b/.idea/.gitignore @@ -0,0 +1,8 @@ +# Default ignored files +/* + +# things we want to be shared +!/dictionaries +!/inspectionProfiles +!/scopes +!/.gitignore diff --git a/.idea/inspectionProfiles/Project_Default.xml b/.idea/inspectionProfiles/Project_Default.xml new file mode 100644 index 00000000..02b2d7f2 --- /dev/null +++ b/.idea/inspectionProfiles/Project_Default.xml @@ -0,0 +1,35 @@ + + + + \ No newline at end of file diff --git a/.run/tests.run.xml b/.run/tests.run.xml new file mode 100644 index 00000000..9d1c328e --- /dev/null +++ b/.run/tests.run.xml @@ -0,0 +1,19 @@ + + + + + + + + + \ No newline at end of file diff --git a/.travis.yml b/.travis.yml deleted file mode 100644 index 34ea22c1..00000000 --- a/.travis.yml +++ /dev/null @@ -1,10 +0,0 @@ -sudo: false - -language: java -dist: trusty -jdk: - - oraclejdk8 - -cache: - directories: - - $HOME/.m2 diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 00000000..cf1605a6 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,64 @@ +# CLAUDE.md + +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. + +## Project + +LittleProxy is a high-performance HTTP/HTTPS proxy library written on top of Netty. +It is consumed as an embedded library (`io.github.littleproxy:littleproxy`) and can also run as a standalone +executable via the shaded jar (`Launcher` main class). + +The active fork is maintained at https://github.com/LittleProxy/LittleProxy. + +## Build & Test Commands + +- Build & run all tests: `mvn test` +- Package (produces the shaded runnable jar): `mvn clean package` +- Skip tests while packaging: `mvn clean package -DskipTests` +- Smoke tests only (fast subset, excludes `slow-test` tag): `mvn package -Psmoke-test` +- Slow tests only (tagged `@Tag("slow-test")`): `mvn package -Pslow-tests` +- Single test class: `mvn test -Dtest=MitmProxyTest` +- Single test method: `mvn test -Dtest=MitmProxyTest#testProxyConnects` +- Run the shaded jar locally: `./run.bash --server --config ./config/littleproxy.properties` (wraps `mvn package -Dmaven.test.skip=true` + `java -jar`) + +CI (`.github/workflows/main.yml`) runs `-Psmoke-test` then `-Pslow-tests` on Ubuntu/macOS/Windows with JDK 17. + +## Toolchain Constraints + +- **Main source targets Java 11** (`11`); **tests compile with Java 17** (`compile-tests-with-java17` execution). New main-source code must not use language features newer than Java 11. +- Source is auto-formatted by **Spotless (google-java-format 1.28.0)** in the `compile` phase — it rewrites files on every build. Don't hand-align imports; just run the build. +- **ErrorProne** runs during compilation with `-Xep:MissingSummary:OFF -Xep:JdkObsolete:OFF -Xep:ReferenceEquality:OFF -Xep:OperatorPrecedence:OFF`. Other checks are active. +- `maven-enforcer-plugin` bans `junit:junit` and `org.hamcrest:hamcrest-core` — use JUnit Jupiter + AssertJ. + +## High-Level Architecture + +The proxy is built around two mirror Netty channel handlers and a small state machine. + +- **`DefaultHttpProxyServer`** (`impl/`) — bootstrap/entry point; owns the `ServerGroup` (shared Netty event loops), `HttpFiltersSource`, optional `MitmManager`, `ChainedProxyManager`, and `ProxyAuthenticator`. `HttpProxyServerBootstrap` is the fluent builder users interact with. +- **`ClientToProxyConnection`** (`impl/`) — one per inbound client channel. Decodes HTTP requests, runs the filter chain, resolves/reuses a per-(host,port) `ProxyToServerConnection`, and writes responses back to the client. Also handles `CONNECT` (tunnel vs. MITM branch) and proxy auth. +- **`ProxyToServerConnection`** (`impl/`) — one per upstream target. Drives `ConnectionFlow` during setup, then proxies request/response chunks. Caches chained-proxy fallback state. +- **`ConnectionFlow` / `ConnectionFlowStep`** (`impl/`) — ordered async step machine used to set up outbound connections: `ConnectChannel` → optional chained-proxy handshake (HTTP CONNECT / SOCKS4 / SOCKS5) → optional `StartTunneling` or `EncryptChannel` (MITM) → `RespondCONNECTSuccessful` → `MitmEncryptClientChannel`. Each step returns a Netty `Future` and advances via callbacks. +- **`ProxyConnection`** (`impl/`) — abstract base for both connections; owns the `ConnectionState` transitions (`AWAITING_INITIAL`, `AWAITING_CHUNK`, `CONNECTING`, `HANDSHAKING`, `NEGOTIATING_CONNECT`, `AWAITING_PROXY_AUTHENTICATION`, `DISCONNECT_REQUESTED`, `DISCONNECTED`). Saturation callbacks implement backpressure (pause client reads when any upstream is saturated, and vice versa). +- **`HttpFilters` / `HttpFiltersSource`** (top-level package) — primary extension point. A new `HttpFilters` instance is created per request via `filterRequest(...)`. Callback order (request → server connect → response) is documented in `LittleProxy_Request_Handling_Architecture.md`; returning a non-null `HttpResponse` from `clientToProxyRequest`/`proxyToServerRequest` short-circuits the request. +- **`MitmManager` + `SslEngineSource`** — supply `SSLEngine`s for HTTPS interception. The in-tree `extras/SelfSignedMitmManager` is demo-grade; production setups typically plug in `LittleProxy-mitm` or the BrowserMob `mitm` module (see README). +- **`ChainedProxyManager` / `ChainedProxy`** — per-request upstream proxy selection. Returning `ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION` or an empty queue fails over to a direct connection; otherwise connection failures walk the queue. +- **Netty pipelines** — client-side: `HAProxyMessageDecoder?` → `HttpRequestDecoder` → optional `HttpObjectAggregator` → monitors → `IdleStateHandler` → `ClientToProxyConnection`. Server-side mirrors this with `HttpRequestEncoder` / `HeadAwareHttpResponseDecoder` plus optional `GlobalTrafficShapingHandler` for throttling. When a `CONNECT` tunnel is established (without MITM), HTTP codecs are removed and data flows as raw bytes via `readRaw`/`write`. +- **`extras/`** contains production-adjacent but optional implementations (`ActivityLogger`, `LogFormat`, `SelfSignedMitmManager`, `HAProxyMessageEncoder`, `TrustingTrustManager`). Core must not depend on `extras`. + +For diagrams and the full lifecycle of CONNECT/MITM/filter callbacks, see `LittleProxy_Request_Handling_Architecture.md`. + +## Testing Notes + +- Always use SLF4J (`org.slf4j.Logger`/`LoggerFactory`) for logging, not log4j directly — this is the logging API used throughout the codebase. +- Tests are JUnit Jupiter + AssertJ + Mockito; Jetty and WireMock are used as real backends (do not mock them away). +- Long-running or timing-sensitive tests are marked `@Tag("slow-test")` and excluded from the default/smoke profile. +- Integration tests start real proxy instances on ephemeral ports; prefer extending `AbstractProxyTest` / `BaseProxyTest` / `BaseChainedProxyTest` rather than duplicating setup. +- Test resources include keystores and `log4j.xml` under `src/test/resources/`. +- Prefer `final` fields for test fixtures (e.g. `private final Foo foo = mock();`) over non-final fields assigned in `@BeforeEach`. +- Prefer initializing fields inline at declaration over assigning them in `@BeforeEach`, when possible (e.g. only fall back to `@BeforeEach` when a value depends on another mock's stubbing or on a checked exception). +- Never use `any()`/`any(Class)` in a `verify(...)` clause — pass the actual expected argument instead, e.g. `verify(throwingTracker).serverDisconnected(fullFlowContext, hostAddress);` not `verify(throwingTracker).serverDisconnected(any(), any());`. `any()` in `verify` only checks that some call happened, not that it happened with the right arguments. +- Avoid reflection to reach a private field/method from a test. Best: test through the class's public API. If that's not practical, widen the member to package-private (test class lives in the same package) rather than reaching in via reflection — less magic, and the IDE/compiler catch renames. + +## Release + +Release steps (version bumps in `pom.xml` + `README.md`, `deploy.bash`, tag, publish on Sonatype Central) are documented in `CONTRIBUTING.md`. Do not bump versions unless a release is being cut. diff --git a/CLAUDE.md b/CLAUDE.md new file mode 100644 index 00000000..25571393 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1,2 @@ +# CLAUDE.md +@AGENTS.md \ No newline at end of file diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md new file mode 100644 index 00000000..de47ecdc --- /dev/null +++ b/CONTRIBUTING.md @@ -0,0 +1,43 @@ + +Contribution guide for LittleProxy + +## Project scope + +LittleProxy is intended to be a LITTLE library, not a kitchen-sink proxy with every conceivable feature. +Before starting work on a new feature, please discuss it first in +[GitHub Discussions](https://github.com/LittleProxy/LittleProxy/discussions) or as a +[GitHub Issue](https://github.com/LittleProxy/LittleProxy/issues), so we can agree whether it fits the +project's scope before you invest time implementing it. + +If you need a richer feature set, consider [BrowserUp Proxy](https://github.com/valfirst/browserup-proxy), +which is built on top of LittleProxy and provides more functionality. + +## How to start + + git clone git@github.com:LittleProxy/LittleProxy.git + cd LittleProxy + mvn test + +## Release + +* Update the release notes (file RELEASE_NOTES.md) +* Change version in the pom.xml file (e.g. "2.4.5-SNAPSHOT" -> "2.4.5") +* Change version in README.md (e.g. "2.4.4" -> "2.4.5") +* Run `mvn clean install` +* Run `deploy.bash` (takes a while, showing "Waiting until Deployment **** is published") +* Log into https://central.sonatype.com/publishing/deployments and ensure that the deployment status is "Published". +* Commit the changes (e.g. with message "release LittleProxy 2.4.5") and make sure the CI build passes +* Run `git tag v2.4.5 && git push --tags` +* Create a new release in GitHub (go to `https://github.com/LittleProxy/LittleProxy/releases/new`) +* Update the version in the pom.xml file to the next SNAPSHOT version (e.g. "2.4.5" -> "2.4.6-SNAPSHOT"). +* Commit the pom.xml change (e.g. with commit message "working on LittleProxy 2.4.6") +* Announce the release in https://groups.google.com/forum/#!forum/littleproxy2 + +## How-to + +### Check the built JAR target version +> mvn clean package -DskipTests +> javap -verbose -classpath target/littleproxy-*SNAPSHOT.jar org.littleshoot.proxy.ProxyAuthenticator | grep "major" + +* major version: 55 -> Java 11 +* major version: 61 -> Java 17 \ No newline at end of file diff --git a/LittleProxy_Request_Handling_Architecture.md b/LittleProxy_Request_Handling_Architecture.md new file mode 100644 index 00000000..e699a3f3 --- /dev/null +++ b/LittleProxy_Request_Handling_Architecture.md @@ -0,0 +1,551 @@ +# LittleProxy Request Handling Architecture + +This document explains with detailed diagrams how LittleProxy handles HTTP/HTTPS requests. + +## Architecture Overview + +```mermaid +flowchart TB + subgraph Client["🖥️ Client (Browser)"] + C[Client Application] + end + + subgraph LittleProxy["🔧 LittleProxy Server"] + direction TB + + subgraph ServerGroup["ServerGroup (Thread Pools)"] + CAT["ClientToProxy Acceptor"] + CWT["ClientToProxy Worker"] + SWT["ProxyToServer Worker"] + end + + subgraph ClientToProxy["ClientToProxyConnection"] + CP_PIPELINE["HTTP Pipeline
Decoder → Encoder"] + CP_HANDLER["Main Handler"] + AUTH["Authentication"] + end + + subgraph ProxyToServer["ProxyToServerConnection"] + PS_PIPELINE["HTTP Pipeline
Encoder → Decoder"] + PS_HANDLER["Main Handler"] + CONN_FLOW["ConnectionFlow"] + end + + subgraph Filters["🎛️ HttpFilters"] + F1["clientToProxyRequest()"] + F2["proxyToServerRequest()"] + F3["serverToProxyResponse()"] + F4["proxyToClientResponse()"] + end + end + + subgraph Upstream["🌐 Upstream"] + direction TB + CHAINED["Chained Proxy
(Optional)"] + TARGET["Target Server"] + end + + C -->|"HTTP/HTTPS Request"| CAT + CAT --> CWT + CWT --> ClientToProxy + CP_PIPELINE --> CP_HANDLER + CP_HANDLER --> F1 + F1 -->|"Short-circuit?"| CP_HANDLER + F1 -->|"Continue"| F2 + F2 --> ProxyToServer + ProxyToServer --> CONN_FLOW + CONN_FLOW -->|"Connect + SSL/SOCKS"| PS_PIPELINE + PS_PIPELINE --> PS_HANDLER + PS_HANDLER -->|"HTTP Proxy"| CHAINED + PS_HANDLER -->|"Direct"| TARGET + CHAINED --> TARGET + TARGET -->|"Response"| PS_HANDLER + PS_HANDLER --> F3 + F3 --> CP_HANDLER + CP_HANDLER --> F4 + F4 -->|"Final Response"| C +``` + +## Main Class Diagram + +```mermaid +classDiagram + class ProxyConnection~I~ { + <> + -ConnectionState currentState + -boolean tunneling + -SSLEngine sslEngine + +read(Object msg) + +write(Object msg) + +connected() + +disconnected() + +encrypt(SSLEngine) + +become(ConnectionState) + } + + class ClientToProxyConnection { + -Map~String,ProxyToServerConnection~ serverConnections + -HttpFilters currentFilters + -HttpRequest currentRequest + -boolean mitming + +readHTTPInitial(HttpRequest) + +readHTTPChunk(HttpContent) + +respond(ProxyToServerConnection, HttpFilters, HttpRequest, HttpResponse, HttpObject) + +serverConnectionSucceeded() + +serverConnectionFailed() + } + + class ProxyToServerConnection { + -ClientToProxyConnection clientConnection + -ChainedProxy chainedProxy + -ConnectionFlow connectionFlow + -Queue~ChainedProxy~ availableChainedProxies + -HttpRequest initialRequest + +readHTTPInitial(HttpResponse) + +write(Object, HttpFilters) + +connectionSucceeded() + +connectionFailed() + +initializeConnectionFlow() + } + + class ConnectionFlow { + -Deque~ConnectionFlowStep~ steps + -ConnectionFlowStep currentStep + +start() + +advance() + +succeed() + +fail() + } + + class ConnectionFlowStep~T~ { + <> + -ProxyConnection connection + -ConnectionState state + +execute() Future + +onSuccess(ConnectionFlow) + +read(ConnectionFlow, Object) + } + + class DefaultHttpProxyServer { + -ServerGroup serverGroup + -HttpFiltersSource filtersSource + -MitmManager mitmManager + -ChainedProxyManager chainProxyManager + -ProxyAuthenticator proxyAuthenticator + +start() + +stop() + } + + class HttpFilters { + <> + +clientToProxyRequest(HttpObject) + +proxyToServerRequest(HttpObject) + +serverToProxyResponse(HttpObject) + +proxyToClientResponse(HttpObject) + } + + ProxyConnection <|-- ClientToProxyConnection + ProxyConnection <|-- ProxyToServerConnection + ClientToProxyConnection "1" --> "0..*" ProxyToServerConnection : manages + ProxyToServerConnection --> ConnectionFlow : uses + ConnectionFlow "1" --> "0..*" ConnectionFlowStep : contains + ClientToProxyConnection --> HttpFilters : applies + ProxyToServerConnection --> HttpFilters : applies + DefaultHttpProxyServer --> ClientToProxyConnection : creates + DefaultHttpProxyServer --> HttpFilters : provides +``` + +## HTTP Request Lifecycle (Non-CONNECT) + +```mermaid +sequenceDiagram + autonumber + actor Client + participant C2P as ClientToProxyConnection + participant Filters as HttpFilters + participant P2S as ProxyToServerConnection + participant CF as ConnectionFlow + participant Target as Target Server + + Note over Client,Target: Standard HTTP Request (GET/POST/...) + + Client->>C2P: HTTP Request + activate C2P + + C2P->>Filters: clientToProxyRequest(request) + alt Short-circuit response + Filters-->>C2P: HttpResponse (short-circuit) + C2P->>C2P: respondWithShortCircuitResponse() + C2P-->>Client: Filtered Response + else Continue processing + Filters-->>C2P: null (continue) + + C2P->>C2P: identifyHostAndPort() + C2P->>C2P: find/reuse ProxyToServerConnection + + C2P->>Filters: proxyToServerRequest(request) + Filters-->>C2P: null (continue) + + C2P->>P2S: write(request, filters) + activate P2S + + alt Existing connection + P2S->>Target: Send request + else New connection + P2S->>CF: initializeConnectionFlow() + CF->>CF: start() → ConnectChannel + CF->>Target: TCP Connection + CF->>CF: succeed() + CF-->>P2S: connectionSucceeded() + P2S->>Target: Send request + end + + P2S->>Filters: proxyToServerRequestSending() + P2S->>Filters: proxyToServerRequestSent() + deactivate P2S + + Target-->>P2S: HTTP Response + activate P2S + P2S->>Filters: serverToProxyResponseReceiving() + P2S->>Filters: serverToProxyResponse(response) + Filters-->>P2S: modified response + + P2S->>C2P: respond(filters, request, response, object) + deactivate P2S + + C2P->>Filters: proxyToClientResponse(response) + Filters-->>C2P: final response + C2P->>C2P: modifyResponseHeaders() + C2P-->>Client: Final Response + deactivate C2P + + P2S->>Filters: serverToProxyResponseReceived() + end +``` + +## CONNECT Request Lifecycle (HTTPS Tunneling) + +```mermaid +sequenceDiagram + autonumber + actor Client + participant C2P as ClientToProxyConnection + participant P2S as ProxyToServerConnection + participant CF as ConnectionFlow + participant Target as Target Server + + Note over Client,Target: HTTPS Tunneling (without MITM) + + Client->>C2P: CONNECT host:443 + activate C2P + C2P->>C2P: doReadHTTPInitial() + C2P->>P2S: create() + write(connectRequest) + activate P2S + + P2S->>CF: initializeConnectionFlow() + CF->>CF: start() + + Note right of CF: ConnectionFlow Steps + CF->>Target: 1. ConnectChannel (TCP) + CF->>Target: 2. StartTunneling (tunnel mode) + + CF-->>C2P: RespondCONNECTSuccessful + C2P-->>Client: 200 Connection established + + CF->>C2P: StartTunneling + CF->>CF: succeed() + CF-->>P2S: connectionSucceeded() + deactivate P2S + + Note over Client,Target: Active Tunnel - raw data + + Client->>C2P: Raw SSL/TLS data + C2P->>P2S: readRaw() → write() + P2S->>Target: Raw SSL/TLS data + + Target-->>P2S: Raw SSL/TLS data + P2S->>C2P: readRaw() → write() + C2P-->>Client: Raw SSL/TLS data + + deactivate C2P +``` + +## Connection Flow (with MITM) + +```mermaid +flowchart TB + subgraph "ConnectionFlow Steps" + START(["Start"]) --> CONNECT["1️⃣ ConnectChannel
TCP Connection"] + CONNECT --> CHAIN_PROXY{"Chained Proxy ?"} + + CHAIN_PROXY -->|"Yes - HTTP"| HTTP_PROXY["HTTPCONNECTWithChainedProxy"] + CHAIN_PROXY -->|"Yes - SOCKS4"| SOCKS4["SOCKS4CONNECTWithChainedProxy"] + CHAIN_PROXY -->|"Yes - SOCKS5"| SOCKS5["SOCKS5InitialRequest
→ SOCKS5SendPasswordCredentials
→ SOCKS5CONNECTRequest"] + CHAIN_PROXY -->|"No"| CHECK_CONNECT{"CONNECT Request ?"} + + HTTP_PROXY --> CHECK_CONNECT + SOCKS4 --> CHECK_CONNECT + SOCKS5 --> CHECK_CONNECT + + CHECK_CONNECT -->|"No"| END_SUCCESS(["Success
AWAITING_INITIAL"]) + + CHECK_CONNECT -->|"Yes"| MITM_CHECK{"MITM Enabled ?"} + + MITM_CHECK -->|"Yes"| ENCRYPT_SERVER["EncryptChannel
(SSL to server)"] + ENCRYPT_SERVER --> RESPOND_OK1["RespondCONNECTSuccessful
200 to client"] + RESPOND_OK1 --> ENCRYPT_CLIENT["MitmEncryptClientChannel
(SSL to client)"] + ENCRYPT_CLIENT --> END_MITM(["MITM Success
AWAITING_INITIAL"]) + + MITM_CHECK -->|"No"| TUNNEL_SERVER["StartTunneling
(server side)"] + TUNNEL_SERVER --> RESPOND_OK2["RespondCONNECTSuccessful
200 to client"] + RESPOND_OK2 --> TUNNEL_CLIENT["StartTunneling
(client side)"] + TUNNEL_CLIENT --> END_TUNNEL(["Tunnel Success
Raw bytes mode"]) + end + + style START fill:#90EE90 + style END_SUCCESS fill:#90EE90 + style END_MITM fill:#FFD700 + style END_TUNNEL fill:#87CEEB +``` + +## Netty Pipeline - Client Side (Inbound) + +```mermaid +flowchart LR + subgraph "ClientToProxy Pipeline" + direction LR + IN["📥 Inbound"] + OUT["📤 Outbound"] + + IN --> BYTES_READ["bytesReadMonitor
📊 Read stats"] + BYTES_READ --> BYTES_WRITE["bytesWrittenMonitor
📊 Write stats"] + BYTES_WRITE --> ENCODER["HttpResponseEncoder
📤 Encode responses"] + + PROXY_DEC["HAProxyMessageDecoder
🌐 Proxy Protocol"] + DECODER["HttpRequestDecoder
📥 Decode requests"] + AGG["HttpObjectAggregator
📦 Buffering (opt)"] + REQ_MON["requestReadMonitor
📊 Request stats"] + RES_MON["responseWrittenMonitor
📊 Response stats"] + IDLE["IdleStateHandler
⏱️ Timeout"] + HANDLER["ClientToProxyConnection
🎯 Main handler"] + + ENCODER --> PROXY_DEC + PROXY_DEC --> DECODER + DECODER --> AGG + AGG --> REQ_MON + REQ_MON --> RES_MON + RES_MON --> IDLE + IDLE --> HANDLER + HANDLER --> OUT + end + + style HANDLER fill:#FFD700,stroke:#FF8C00,stroke-width:3px +``` + +## Netty Pipeline - Server Side (Outbound) + +```mermaid +flowchart LR + subgraph "ProxyToServer Pipeline" + direction LR + IN["📥 Inbound"] + OUT["📤 Outbound"] + + IN --> BYTES_READ["bytesReadMonitor
📊 Read stats"] + BYTES_READ --> BYTES_WRITE["bytesWrittenMonitor
📊 Write stats"] + BYTES_WRITE --> TRAFFIC["GlobalTrafficShapingHandler
🚦 Throttling (opt)"] + + TRAFFIC --> PROXY_ENC["HAProxyMessageEncoder
🌐 Proxy Protocol"] + PROXY_ENC --> ENCODER["HttpRequestEncoder
📤 Encode requests"] + DECODER["HeadAwareHttpResponseDecoder
📥 Decode responses"] + AGG["HttpObjectAggregator
📦 Buffering (opt)"] + RES_MON["responseReadMonitor
📊 Response stats"] + REQ_MON["requestWrittenMonitor
📊 Request stats"] + IDLE["IdleStateHandler
⏱️ Timeout"] + HANDLER["ProxyToServerConnection
🎯 Main handler"] + + ENCODER --> DECODER + DECODER --> AGG + AGG --> RES_MON + RES_MON --> REQ_MON + REQ_MON --> IDLE + IDLE --> HANDLER + HANDLER --> OUT + end + + style HANDLER fill:#87CEEB,stroke:#4682B4,stroke-width:3px +``` + +## Connection State Machine + +```mermaid +stateDiagram-v2 + [*] --> AWAITING_INITIAL: Connection accepted + + AWAITING_INITIAL --> AWAITING_CHUNK: Chunked request received + AWAITING_CHUNK --> AWAITING_CHUNK: Chunk received + AWAITING_CHUNK --> AWAITING_INITIAL: LastHttpContent received + + AWAITING_INITIAL --> CONNECTING: New server connection + CONNECTING --> AWAITING_CONNECT_OK: TCP connection OK + CONNECTING --> DISCONNECTED: Connection failed + + AWAITING_CONNECT_OK --> HANDSHAKING: SSL/SOCKS/CONNECT in progress + AWAITING_CONNECT_OK --> AWAITING_INITIAL: Connect OK (no TLS) + + HANDSHAKING --> AWAITING_INITIAL: SSL handshake succeeded + HANDSHAKING --> DISCONNECTED: Handshake failed + + AWAITING_INITIAL --> NEGOTIATING_CONNECT: CONNECT request received + NEGOTIATING_CONNECT --> AWAITING_INITIAL: Tunnel established + + AWAITING_INITIAL --> AWAITING_PROXY_AUTHENTICATION: Auth required + AWAITING_PROXY_AUTHENTICATION --> AWAITING_INITIAL: Auth succeeded + AWAITING_PROXY_AUTHENTICATION --> DISCONNECT_REQUESTED: Auth failed + + AWAITING_INITIAL --> DISCONNECT_REQUESTED: Close requested + AWAITING_CHUNK --> DISCONNECT_REQUESTED: Close requested + + DISCONNECT_REQUESTED --> DISCONNECTED: Disconnect + DISCONNECTED --> [*]: End +``` + +## Filter Chain (HttpFilters) + +```mermaid +flowchart TB + subgraph "Filter Chain Execution Order" + direction TB + + REQ["📝 Client Request"] --> F1["1. clientToProxyRequest()
📍 Initial interception
Short-circuit possible"] + + F1 -->|"Short-circuit"| RESP_FINAL["🔚 Client Response"] + F1 -->|"Continue"| F2["2. proxyToServerRequest()
📍 Modify before sending to server
Short-circuit possible"] + + F2 -->|"Short-circuit"| RESP_FINAL + F2 -->|"Continue"| CONN["🔌 Connect to Server"] + + CONN --> F3["3. serverToProxyResponse()
📍 Modify server response
Return null = disconnect"] + + F3 -->|"null (force disconnect)"| DISCONNECT["❌ Force disconnect"] + F3 -->|"Continue"| F4["4. proxyToClientResponse()
📍 Final response modification
Return null = disconnect"] + + F4 -->|"null (force disconnect)"| DISCONNECT + F4 -->|"Continue"| RESP_FINAL + end + + subgraph "Callback Notifications (Order)" + direction TB + N1["clientToProxyRequest"] --> N2["proxyToServerConnectionQueued"] + N2 --> N3["proxyToServerResolutionStarted"] + N3 --> N4["proxyToServerResolutionSucceeded/Failed"] + N4 --> N5["proxyToServerRequest"] + N5 --> N6["proxyToServerConnectionStarted"] + N6 --> N7["proxyToServerConnectionSSLHandshakeStarted (if HTTPS)"] + N7 --> N8["proxyToServerConnectionSucceeded/Failed"] + N8 --> N9["proxyToServerRequestSending"] + N9 --> N10["proxyToServerRequestSent"] + N10 --> N11["serverToProxyResponseReceiving"] + N11 --> N12["serverToProxyResponse"] + N12 --> N13["serverToProxyResponseReceived"] + N13 --> N14["proxyToClientResponse"] + end + + style F1 fill:#FFD700 + style F2 fill:#FFD700 + style F3 fill:#87CEEB + style F4 fill:#87CEEB +``` + +## Chained Proxy Management + +```mermaid +flowchart TB + subgraph "ChainedProxy Resolution" + START(["New request"]) --> RESOLVE["ChainedProxyManager.lookupChainedProxies()"] + + RESOLVE --> HAS_PROXY{"Proxies
available ?"} + HAS_PROXY -->|"No"| NULL["Returns null
502 Bad Gateway"] + HAS_PROXY -->|"Yes"| QUEUE["Proxy queue
ConcurrentLinkedQueue"] + + QUEUE --> CREATE["Create ProxyToServerConnection
with first proxy"] + CREATE --> CONNECT["Attempt connection"] + + CONNECT --> SUCCESS{"Connection
OK ?"} + SUCCESS -->|"Yes"| USE["Use this proxy"] + SUCCESS -->|"No"| FALLBACK["ChainedProxy.connectionFailed()"] + + FALLBACK --> NEXT_PROXY{"More proxies
in queue ?"} + NEXT_PROXY -->|"Yes"| NEXT["Take next proxy"] + NEXT --> CONNECT + + NEXT_PROXY -->|"No"| FALLBACK_DIRECT["FALLBACK_TO_DIRECT_CONNECTION"] + FALLBACK_DIRECT --> DIRECT["Direct connection
(no proxy)"] + DIRECT --> DIRECT_SUCCESS{"Direct
OK ?"} + DIRECT_SUCCESS -->|"Yes"| USE + DIRECT_SUCCESS -->|"No"| FAIL["502 Bad Gateway"] + end + + subgraph "Proxy Types" + HTTP["HTTP Proxy
→ CONNECT request"] + SOCKS4["SOCKS4 Proxy
→ CONNECT command"] + SOCKS5["SOCKS5 Proxy
→ Auth + CONNECT command"] + end + + style USE fill:#90EE90 + style FAIL fill:#FF6B6B + style NULL fill:#FF6B6B +``` + +## Backpressure Management + +```mermaid +flowchart TB + subgraph "Backpressure Management" + C2P_SAT["ClientToProxy saturated"] --> STOP_ALL["Stop reading ALL
ProxyToServer connections"] + + P2S_SAT["ProxyToServer saturated"] --> STOP_CLIENT["Stop reading
ClientToProxy"] + + C2P_WRITE["ClientToProxy writable"] --> CHECK_ALL{"All servers
writable ?"} + CHECK_ALL -->|"Yes"| RESUME_CLIENT["Resume reading
ClientToProxy"] + CHECK_ALL -->|"No"| WAIT["Wait"] + + P2S_WRITE["ProxyToServer writable"] --> CHECK_SAT{"All servers
non-saturated ?"} + CHECK_SAT -->|"Yes"| RESUME_ALL["Resume reading ALL
ProxyToServer + Client"] + CHECK_SAT -->|"No"| KEEP_WAIT["Keep waiting"] + end + + style STOP_ALL fill:#FFD700 + style STOP_CLIENT fill:#FFD700 + style RESUME_CLIENT fill:#90EE90 + style RESUME_ALL fill:#90EE90 +``` + +## Key Components Summary + +| Component | Role | File | +|-----------|------|------| +| **DefaultHttpProxyServer** | Entry point, bootstrap, configuration | `DefaultHttpProxyServer.java` | +| **ClientToProxyConnection** | Manages incoming client connections | `ClientToProxyConnection.java` | +| **ProxyToServerConnection** | Manages outgoing server connections | `ProxyToServerConnection.java` | +| **ConnectionFlow** | Orchestration of connection steps | `ConnectionFlow.java` | +| **ConnectionFlowStep** | Individual step in the connection flow | `ConnectionFlowStep.java` | +| **HttpFilters** | Interface for filtering/modifying requests/responses | `HttpFilters.java` | +| **ProxyConnection** | Abstract base class for connections | `ProxyConnection.java` | +| **ServerGroup** | Netty thread pool management | `ServerGroup.java` | + +## Key Architecture Points + +1. **Separation of Concerns**: The [`ClientToProxyConnection`](src/main/java/org/littleshoot/proxy/impl/ClientToProxyConnection.java:88) class manages the client side, while [`ProxyToServerConnection`](src/main/java/org/littleshoot/proxy/impl/ProxyToServerConnection.java:104) manages the server side. + +2. **Connection Reuse**: Only one [`ProxyToServerConnection`](src/main/java/org/littleshoot/proxy/impl/ProxyToServerConnection.java:104) per host:port is maintained and reused for HTTP requests. + +3. **Tunnel Mode**: For CONNECT requests (HTTPS), HTTP encoders/decoders are removed and data passes through in raw bytes mode. + +4. **MITM (Man-In-The-Middle)**: Allows decrypting HTTPS traffic by acting as an SSL server on the client side and SSL client on the server side. + +5. **Chained Proxies**: Support for HTTP, SOCKS4, and SOCKS5 with automatic fallback mechanism. + +6. **Filters**: Extensible filter chain allowing modification of requests/responses at different stages. + +7. **Backpressure**: Saturation management mechanism to prevent memory overload. diff --git a/Netty_4_Upgrade_Notes.md b/Netty_4_Upgrade_Notes.md index 53fab01d..414fc93f 100644 --- a/Netty_4_Upgrade_Notes.md +++ b/Netty_4_Upgrade_Notes.md @@ -21,7 +21,7 @@ Relevant Changes * DefaultChannelGroup? -* SimpleChannelUpstreamHandler -> SimpleChanneInboundHandler +* SimpleChannelUpstreamHandler -> SimpleChannelInboundHandler * InterestOps is gone - what does this mean to setReadable() and channelInterestChanged() ? diff --git a/PERFORMANCE_AND_LOGGING.md b/PERFORMANCE_AND_LOGGING.md new file mode 100644 index 00000000..2a5584a3 --- /dev/null +++ b/PERFORMANCE_AND_LOGGING.md @@ -0,0 +1,666 @@ +# LittleProxy Performance and Logging Guide + +This guide covers logging performance optimization techniques and configuration options for LittleProxy. + +## Table of Contents + +- [Logging Modes](#logging-modes) + - [Synchronous Logging](#synchronous-logging) + - [Asynchronous Logging](#asynchronous-logging) +- [Performance Considerations](#performance-considerations) +- [Logging Configuration](#logging-configuration) +- [Log Filtering and Rate Limiting](#log-filtering-and-rate-limiting) + - [BurstFilter Configuration](#burstfilter-configuration) + - [Custom Filter Implementation](#custom-filter-implementation) +- [Activity Logging](#activity-logging) +- [Best Practices](#best-practices) +- [Troubleshooting](#troubleshooting) + +## Logging Modes + +### Synchronous Logging + +**Default Mode**: Synchronous logging is the default behavior and provides reliable logging with immediate disk writes. + +**Characteristics:** +- ✅ Simple and reliable +- ✅ Logs are immediately written to disk +- ❌ Higher I/O overhead +- ❌ Can impact proxy performance under heavy load +- ❌ Slower response times during peak traffic + +**Configuration File**: `src/main/resources/littleproxy_default_log4j2.xml` + +**Appender Type**: `RollingFile` with immediate flush + +**Example Usage:** +```bash +./run.bash --server --config ./config/littleproxy.properties --port 9092 +``` + +### Asynchronous Logging + +**Performance Mode**: Asynchronous logging significantly improves performance by buffering log events and writing them in batches. + +**Characteristics:** +- ✅ Much lower I/O overhead +- ✅ Better performance under heavy load +- ✅ Reduced disk I/O operations +- ✅ Configurable buffer sizes +- ❌ Slight risk of log loss on JVM crash +- ❌ Logs may be delayed during shutdown + +**Configuration File**: `src/main/resources/littleproxy_async_log4j2.xml` + +**Appender Type**: `RollingRandomAccessFile` with `immediateFlush="false"` + +**Async Features:** +- `AsyncRoot` for root logger +- `AsyncLogger` for all specific loggers +- `includeLocation="true"` for better debugging + +**Performance Optimizations:** +- **Buffer Size**: Configurable via Log4j2 system properties +- **Batch Writing**: Logs are written in batches rather than individually +- **Reduced I/O**: `immediateFlush="false"` reduces disk operations +- **Larger Files**: 250MB file size vs 50MB in sync mode reduces rollover frequency + +**Example Usage:** +```bash +# Using the async_logging_default flag +./run.bash --async_logging_default --server --config ./config/littleproxy.properties --port 9092 + +# Direct Java command +java -server -XX:+HeapDumpOnOutOfMemoryError -Xmx800m \ + -jar ./target/littleproxy-2.9.1-littleproxy-shade.jar \ + --server --config ./config/littleproxy.properties --port 9092 \ + --log_config ./target/classes/littleproxy_async_log4j2.xml +``` + +## Performance Considerations + +### When to Use Synchronous Logging + +- **Development environments** where immediate log visibility is important +- **Debugging scenarios** where you need real-time log output +- **Low-traffic productions** where performance impact is negligible +- **Compliance requirements** that mandate immediate log persistence + +### When to Use Asynchronous Logging + +- **High-traffic productions** where performance is critical +- **Load testing environments** to get accurate performance metrics +- **Resource-constrained systems** where I/O reduction is needed +- **Burst traffic scenarios** where logging can become a bottleneck + +### Performance Impact Comparison + +| Metric | Synchronous | Asynchronous | Improvement | +|--------|-------------|--------------|-------------| +| **Throughput** | Baseline | +30-50% | 1.3-1.5x | +| **Latency** | Baseline | -40-60% | 0.4-0.6x | +| **Disk I/O** | High | Low | 5-10x less | +| **CPU Usage** | Moderate | Lower | 10-20% less | +| **Memory Usage** | Low | Slightly higher | Buffer overhead | + +## Logging Configuration + +### Default Configuration (Synchronous) + +```xml + + + + %d{ISO8601} %-5p [%t] %c{2} (%F:%L).%M() - %m%n + + + + + + +``` + +### Async Configuration + +```xml + + + + %d{ISO8601} %-5p [%t] %c{2} (%F:%L).%M() - %m%n + + + + + + + + + + + + + +``` + +### Advanced Configuration Options + +**Log4j2 System Properties:** + +```bash +# Increase async logger ring buffer size (default: 262144) +java -Dlog4j2.AsyncLogger.RingBufferSize=1048576 ... + +# Increase async logger queue size +java -Dlog4j2.AsyncLoggerConfig.RingBufferSize=1048576 ... + +# Disable location tracking for better performance +java -Dlog4j2.includeLocation=false ... +``` + +## Log Filtering and Rate Limiting + +### BurstFilter Configuration + +Log4j2 provides a `BurstFilter` that can limit the number of log events within a time period to prevent log flooding. + +**Example Configuration:** + +```xml + + + + + + + + + +``` + +**BurstFilter Parameters:** + +- `level`: The log level to filter (INFO, DEBUG, WARN, etc.) +- `rate`: Maximum number of log events per second +- `maxBurst`: Maximum number of events allowed in a burst + +**Example Scenarios:** + +```xml + + + + + + + + +``` + +### Custom Filter Implementation + +For more sophisticated filtering, you can implement custom Log4j2 filters using Java or Groovy scripts. + +#### Java Custom Filter Example + +**Example Custom Filter:** + +```java +package org.littleshoot.proxy.logging; + +import org.apache.logging.log4j.core.LogEvent; +import org.apache.logging.log4j.core.filter.AbstractFilter; + +/** + * Custom filter to exclude specific request patterns + */ +public class RequestPatternFilter extends AbstractFilter { + + private final String[] excludedPatterns; + + public RequestPatternFilter(String[] excludedPatterns) { + this.excludedPatterns = excludedPatterns; + } + + @Override + public Result filter(LogEvent event) { + String message = event.getMessage().getFormattedMessage(); + + // Check if message matches any excluded pattern + for (String pattern : excludedPatterns) { + if (message.contains(pattern)) { + return Result.DENY; // Exclude this log + } + } + + return Result.NEUTRAL; // Allow this log + } + + @Override + public Result filter(org.apache.logging.log4j.core.Logger logger, org.apache.logging.log4j.Level level, + org.apache.logging.log4j.Marker marker, String msg, Object... params) { + return filterLogEvent(level, marker, msg, params); + } + + @Override + public Result filter(org.apache.logging.log4j.core.Logger logger, org.apache.logging.log4j.Level level, + org.apache.logging.log4j.Marker marker, Object msg, Throwable t) { + return filterLogEvent(level, marker, msg, t); + } + + @Override + public Result filter(org.apache.logging.log4j.core.Logger logger, org.apache.logging.log4j.Level level, + org.apache.logging.log4j.Marker marker, Message msg, Throwable t) { + return filterLogEvent(level, marker, msg, t); + } + + private Result filterLogEvent(org.apache.logging.log4j.Level level, org.apache.logging.log4j.Marker marker, + Object msg, Object... params) { + if (msg instanceof String message) { + for (String pattern : excludedPatterns) { + if (message.contains(pattern)) { + return Result.DENY; + } + } + } + return Result.NEUTRAL; + } +} +``` + +#### Groovy Script Filter Example + +Log4j2 supports Groovy scripts for dynamic filtering without compilation. This is perfect for sampling, conditional logging, or complex logic. + +**Example: Sampling Filter (log only 10% of messages)** + +```xml + + + + + + + + + + + + + + + + + +``` + +**Example: Conditional Filter (exclude health checks)** + +```xml + + + + + + + + +``` + +**Example: Time-Based Filter (business hours only)** + +```xml + + + + + + + + +``` + +**Groovy Script Filter Benefits:** + +- ✅ **No compilation needed**: Scripts are interpreted at runtime +- ✅ **Dynamic logic**: Can use complex conditions and external data +- ✅ **Easy to modify**: Change filtering logic without recompiling +- ✅ **Powerful**: Access to full Groovy language features +- ✅ **Performance**: Still efficient for most use cases + +**Script Filter Parameters:** + +- `onMatch`: What to do when script returns true (ACCEPT/DENY/NEUTRAL) +- `onMismatch`: What to do when script returns false (ACCEPT/DENY/NEUTRAL) +- `language`: Script language (groovy, javascript, etc.) + +**Available Variables in Script:** + +- `logEvent`: The current LogEvent object +- `loggerName`: Name of the logger +- `level`: Log level +- `message`: Formatted message +- `marker`: Marker (if any) +- `throwable`: Exception (if any) + +**Performance Considerations:** + +- Script filters add some overhead compared to compiled filters +- Use for complex logic that's hard to implement in Java +- Test script performance before deploying to production +- Consider caching results for repeated patterns + +**XML Configuration for Custom Filter:** + +```xml + + + + + + + + +``` + +### Filter Examples by Use Case + +**1. Exclude Health Check Requests:** +```xml + +``` + +**2. Limit Debug Logs in Production:** +```xml + +``` + +**3. Filter by Request Type:** +```xml + +``` + +**4. Rate Limiting for Specific Loggers:** +```xml + + + +``` + +## Activity Logging + +Activity logging in LittleProxy captures HTTP request/response details. This can be a significant performance factor. + +### Activity Log Formats + +**CLF (Common Log Format):** +```bash +--activity_log_format CLF +``` + +**ELF (Extended Log Format):** +```bash +--activity_log_format ELF +``` + +**JSON:** +```bash +--activity_log_format JSON +``` + +**SQUID:** +```bash +--activity_log_format SQUID +``` + +**W3C:** +```bash +--activity_log_format W3C +``` + +**LTSV (Labeled Tab-Separated Values):** +```bash +--activity_log_format LTSV +``` + +**CSV (Comma-Separated Values):** +```bash +--activity_log_format CSV +``` + +**HAPROXY:** +```bash +--activity_log_format HAPROXY +``` + +### Activity Logging Performance Impact + +| Format | Performance | Use Case | +|--------|-------------|----------| +| **CLF** | ⚡ Fastest | Production, high volume | +| **ELF** | ⚡ Fast | Extended logging needs | +| **JSON** | 🏃 Moderate | Analytics, structured logging | +| **SQUID** | 🏃 Moderate | Squid proxy compatibility | +| **W3C** | 🏃 Moderate | Web standards compliance | +| **LTSV** | 🏃 Moderate | Machine-readable logs | +| **CSV** | 🏃 Moderate | Spreadsheet analysis | +| **HAPROXY** | 🏃 Moderate | HAProxy compatibility | + +### Activity Logging Best Practices + +1. **Disable in Development**: omit the `--activity_log_format` flag when not needed +2. **Use CLF in Production**: Fastest format for high-volume scenarios +3. **Sample Activity Logs**: Consider sampling (every Nth request) +4. **Separate Activity Logs**: Use different files for access logs vs application logs + +**Example with Activity Logging:** +```bash +# Async logging with CLF activity format (best performance) +./run.bash --async_logging_default --server --config ./config/littleproxy.properties \ + --port 9092 --activity_log_format CLF + +# Sync logging with JSON activity format (structured logging) +./run.bash --server --config ./config/littleproxy.properties \ + --port 9092 --activity_log_format JSON +``` + +## Best Practices + +### Logging Configuration + +1. **Use Async for Production**: Always use `--async_logging_default` in production +2. **Keep Sync for Development**: Use default sync logging during development +3. **Monitor Log Growth**: Set appropriate file sizes and rotation policies +4. **Test Configuration**: Validate Log4j2 configuration before deployment + +### Performance Optimization + +1. **Tune Buffer Sizes**: Adjust `log4j2.AsyncLogger.RingBufferSize` based on load +2. **Limit Location Info**: Use `includeLocation="false"` for production +3. **Filter Early**: Apply filters at the appender level +4. **Use Appropriate Levels**: DEBUG for development, INFO for production + +### Monitoring and Maintenance + +1. **Monitor Log Files**: Check disk usage regularly +2. **Rotate Logs**: Configure proper rotation policies +3. **Archive Old Logs**: Implement log archiving strategy +4. **Alert on Errors**: Set up monitoring for ERROR level logs + +## Troubleshooting + +### Common Issues and Solutions + +**Issue: No logs appearing with async mode** +- **Solution**: Check that `littleproxy_async_log4j2.xml` is in the correct location +- **Solution**: Verify file permissions on log directory +- **Solution**: Check for Log4j2 configuration errors + +**Issue: High CPU usage with logging** +- **Solution**: Switch to async logging mode +- **Solution**: Reduce log level from DEBUG to INFO +- **Solution**: Apply BurstFilter to limit log volume + +**Issue: Disk full due to logs** +- **Solution**: Configure proper rotation policies +- **Solution**: Increase file size limits +- **Solution**: Implement log archiving + +**Issue: Log4j2 configuration errors** +- **Solution**: Check XML syntax +- **Solution**: Validate with `status="TRACE"` in configuration +- **Solution**: Ensure all referenced appenders exist + +### Debugging Log4j2 Configuration + +Add debug output to Log4j2: + +```xml + + + +``` + +**Debug Levels:** +- `OFF`: No internal logging +- `ERROR`: Only errors +- `WARN`: Warnings and errors +- `INFO`: Informational messages +- `DEBUG`: Debug information +- `TRACE`: Verbose debugging + +### Checking Async Logger Status + +```bash +# Check async logger buffer status +java -Dlog4j2.AsyncLoggerConfig.StatusLogger.level=INFO \ + -jar ./target/littleproxy-2.9.1-littleproxy-shade.jar \ + --server --config ./config/littleproxy.properties --port 9092 +``` + +## Advanced Topics + +### Custom Appender Implementation + +For specialized logging needs, implement custom appenders: + +```java +@Plugin(name = "CustomAppender", category = "Core", elementType = "appender", printObject = true) +public class CustomAppender extends AbstractAppender { + + protected CustomAppender(String name, Filter filter, Layout layout) { + super(name, filter, layout); + } + + @Override + public void append(LogEvent event) { + // Custom logging logic + byte[] bytes = getLayout().toByteArray(event); + // Send to custom destination (database, network, etc.) + } +} +``` + +### Dynamic Log Level Adjustment + +Change log levels at runtime: + +```java +// Get the logger context +LoggerContext context = (LoggerContext) LogManager.getContext(false); +Configuration config = context.getConfiguration(); + +// Adjust log level +LoggerConfig loggerConfig = config.getLoggerConfig(LogManager.ROOT_LOGGER_NAME); +loggerConfig.setLevel(Level.DEBUG); + +// Update configuration +context.updateLoggers(); +``` + +### Log Enrichment + +Add contextual information to logs: + +```java +// Use ThreadContext to add contextual data +ThreadContext.put("requestId", UUID.randomUUID().toString()); +ThreadContext.put("clientIp", clientAddress); + +try { + // Process request - logs will include context data + logger.info("Processing request"); +} finally { + ThreadContext.clear(); +} +``` + +**Pattern Layout with Context:** +```xml + +``` + +## Summary + +This guide provides comprehensive information on optimizing LittleProxy logging performance: + +- **Synchronous vs Asynchronous**: Choose based on your performance needs +- **Configuration Options**: Default and async configurations provided +- **Filtering**: BurstFilter and custom filters for rate limiting +- **Activity Logging**: Format options and performance considerations +- **Best Practices**: Production-ready recommendations +- **Troubleshooting**: Common issues and solutions + +For most production environments, **asynchronous logging with CLF activity format** provides the best balance of performance and functionality: + +```bash +./run.bash --async_logging_default --server --config ./config/littleproxy.properties \ + --port 9092 --activity_log_format CLF +``` \ No newline at end of file diff --git a/README.md b/README.md index 6e0eb4e5..2111d9ae 100644 --- a/README.md +++ b/README.md @@ -1,94 +1,447 @@ -[![Build Status](https://travis-ci.com/mrog/LittleProxy.svg?branch=master)](https://travis-ci.com/mrog/LittleProxy) -[![DepShield Badge](https://depshield.sonatype.org/badges/mrog/LittleProxy/depshield.svg)](https://depshield.github.io) - This is an updated fork of adamfisk's LittleProxy. The original project appears -to have been abondoned. Because it's so incredibly useful, it's being brought +to have been abandoned. Because it's so incredibly useful, it's being brought back to life in this repository. LittleProxy is a high performance HTTP proxy written in Java atop Trustin Lee's excellent [Netty](http://netty.io) event-based networking library. It's quite -stable, performs well, and is easy to integrate into your projects. +stable, performs well, and is easy to integrate into your projects. + +# Usage + +## Command Line -One option is to clone LittleProxy and run it from the command line. This is as simple as: +One option is to clone LittleProxy and run it from the command line. This is as simple as running the following commands : ``` -$ git clone git@github.com:mrog/LittleProxy.git +$ git clone git@github.com:LittleProxy/LittleProxy.git $ cd LittleProxy $ ./run.bash ``` -You can embed LittleProxy in your own projects through Maven with the following: +### Options + +Multiple options can be passed to the script as arguments. The following options are supported : + +#### Config File + +This will start LittleProxy with the configuration (path relative to the working directory or absolute) +specified in the given file. + +```bash +$ ./run.bash --config path/to/config/littleproxy.properties +``` + +You can, for example, run the shell script at the root project directory as a server, pointing +to the provided _littleproxy.properties_ file : + +```bash +$ ./run.bash --server --config ./config/littleproxy.properties +``` + +##### config file description + +The config file is a properties file with the following properties : +- `dnssec` : boolean value to enable/disable DNSSEC validation (default : `false`) +- `transparent` : boolean value to enable/disable transparent proxy mode (default : `false`) +- `idleConnectionTimeout` : integer value to set the idle connection timeout in seconds (default : `-1`, i.e. no timeout) +- `connect_timeout` : integer value to set the connect timeout in seconds (default : `0`, i.e. no timeout) +- `max_initial_line_length` : integer value to set the max initial line length in bytes (default : `8192`) +- `max_header_size` : integer value to set the max header size in bytes (default : `16384`) +- `max_chunk_size` : integer value to set the max chunk size in bytes (default : `16384`) +- `server_connection_pool_type` : pool implementation used by the shared server connection pool (`CONCURRENT_MAP`) (default : `CONCURRENT_MAP`) -- only effective when `use_shared_server_connection_pool=true` +- `max_total_connections` : integer value to set the maximum total pooled server connections (default : `200`) -- only effective when `use_shared_server_connection_pool=true` +- `max_connections_per_host` : integer value to set the maximum pooled server connections per host:port (default : `10`) -- only effective when `use_shared_server_connection_pool=true` +- `name` : string value to set the proxy server name (default : `LittleProxy`) +- `address` : string value to set the proxy server address (default : `0.0.0.0:8080`) +- `port` : integer value to set the proxy server port (default : `8080`) +- `nic` : string value to set the network interface card (default : `0.0.0.0`) +- `proxy_alias` : string value to set the proxy alias (default : hostname of the machine) +- `allow_local_only` : boolean value to allow only local connections (default : `false`) +- `authenticate_ssl_clients` : boolean value to enable/disable SSL client authentication (default : `false`) +- `ssl_clients_trust_all_servers` : boolean value to trust all servers (default : `false`) +- `ssl_clients_send_certs` : boolean value to send certificates (default : `false`) +- `ssl_clients_key_store_file_path` : string value to set the key store file path (default : `null`) +- `ssl_clients_key_store_alias` : string value to set the key store alias (default : `null`) +- `ssl_clients_key_store_password` : string value to set the key store password (default : `null`) +- `throttle_read_bytes_per_second` : integer value to set the throttle read bytes per second (default : `0`) +- `throttle_write_bytes_per_second` : integer value to set the throttle write bytes per second (default : `0`) +- `allow_requests_to_origin_server` : boolean value to allow requests to origin server (default : `false`) +- `allow_proxy_protocol` : boolean value to allow proxy protocol (default : `false`) +- `send_proxy_protocol` : boolean value to send proxy protocol header (default : `false`) +- `activity_log_format` : string value to set the activity log format (CLF, ELF, JSON, LTSV, CSV, SQUID, HAPROXY) (default: disabled) + +Options set from the command line, override the ones set in the config file. + +> **Note**: For advanced logging configuration and performance optimization, see our [Performance and Logging Guide](PERFORMANCE_AND_LOGGING.md). + +##### littleproxy.properties Example + +````properties +dnssec=true +transparent=false +idleConnectionTimeout=60 +connect_timeout=30 +max_initial_line_length=8192 +max_header_size=16384 +max_chunk_size=16384 +server_connection_pool_type=CONCURRENT_MAP +max_total_connections=200 +max_connections_per_host=10 +name=LittleProxy +address=192.168.1.100:8080 +port=8080 +nic=eth0 +proxy_alias=myproxy +allow_local_only=false +authenticate_ssl_clients=false +ssl_clients_trust_all_servers=false +ssl_clients_send_certs=false +ssl_clients_key_store_file_path=/path/to/keystore.jks +ssl_clients_key_store_alias=myalias +ssl_clients_key_store_password=mypassword +throttle_read_bytes_per_second=1024 +throttle_write_bytes_per_second=1024 +allow_requests_to_origin_server=true +allow_proxy_protocol=true +send_proxy_protocol=true +activity_log_format=CLF +```` +#### DNSSec + +This will start LittleProxy with DNSSEC validation enabled ; i.e, it will use secure DNS lookups for outbound +connections. + + +```bash +$ ./run.bash --dnssec true +``` + +#### Log configuration file + +This will start LittleProxy with the specified log configuration file. +Path of the log configuration file can be relative or absolute. + +If it is relative, it will be resolved relative to the current working directory : +```bash +$ ./run.bash --log_config ./log4j.xml +``` +If it is absolute, it will be resolved as is : + +```bash +$ ./run.bash --log_config /home/user/log4j.xml +``` + +#### Activity Log Format + +This will enable the activity tracker with the specified log format. +Supported formats: `CLF`, `ELF`, `W3C`, `JSON`, `LTSV`, `CSV`, `SQUID`, `HAPROXY`. + +```bash +$ ./run.bash --activity_log_format CLF +``` + +#### Port + +This will start LittleProxy on port `8080` by default. +You can customize the port by passing a port number as an argument to the script : + +```bash +$ ./run.bash --port 9090 +``` + +#### NIC + +This will start LittleProxy on the default network interface. You can customize the network interface by passing +a NIC name (`eth0` in the example below) as an argument to the script : + +```bash +$ ./run.bash --nic eth0 +``` + +#### MITM Manager + +If you pass this option, this will start LittleProxy with the default MITM manager (`SelfSignedMitmManager` implementation). +It will generate a self-signed certificate for each domain you visit. + +```bash +$ ./run.bash --mitm +``` +#### name + +This will start LittleProxy with the specified name. This name will be used to name the threads. + +```bash +$ ./run.bash --name MyProxy +``` + +#### address + +This will start LittleProxy binding to the specified address. IPV4,IPV6 and hostname addresses are supported. + +```bash +$ ./run.bash --address 127.0.0.1:8080 +``` +#### nic + +This will start LittleProxy binding to the specified network interface. + +```bash +$ ./run.bash --nic eth0 +``` + +#### proxy_alias + +This will start LittleProxy with the specified proxy alias. +The alias or pseudonym for this proxy, used when adding the `Via` header. + +```bash +$ ./run.bash --proxy_alias MyProxy +``` + +#### allow_local_only + +This will start LittleProxy allowing only local connections (default is `false`). + +```bash +$ ./run.bash --allow_local_only true +``` + +#### authenticate_ssl_clients + +This will start LittleProxy authenticating SSL clients (default is `false`). + +```bash +$ ./run.bash --authenticate_ssl_clients true``` + +#### trust_all_servers + +This will start LittleProxy authenticating SSL clients and trusting all servers (default is `false`). + +```bash +$ ./run.bash --authenticate_ssl_clients true --trust_all_servers true +``` +#### send_certs + +This will start LittleProxy authenticating SSL clients and sending certificates (default is `false`). + +```bash +$ ./run.bash --authenticate_ssl_clients true --send_certs true``` + +#### ssl_client_keystore_path + +This will start LittleProxy authenticating SSL clients and using the specified keystore path. + +```bash +$ ./run.bash --authenticate_ssl_clients true --ssl_client_keystore_path /path/to/keystore` +``` +#### ssl_client_keystore_alias + +This will start LittleProxy authenticating SSL clients and using the specified keystore alias. + +```bash +$ ./run.bash --authenticate_ssl_clients true --ssl_client_keystore_alias myalias``` +``` +#### ssl_client_keystore_password + +This will start LittleProxy authenticating SSL clients and using the specified keystore password. + +```bash +$ ./run.bash --authenticate_ssl_clients true --ssl_client_keystore_password mypassword``` + +#### throttle_read_bytes_per_second + +This will start LittleProxy throttling the read bytes per second. + +```bash +$ ./run.bash --throttle_read_bytes_per_second 1024 +``` + +#### throttle_write_bytes_per_second + +This will start LittleProxy throttling the write bytes per second. + +```bash +$ ./run.bash --throttle_write_bytes_per_second 1024 +``` + +#### allow_request_to_origin_server + +This will start LittleProxy allowing requests to the origin server. + +```bash +$ ./run.bash --allow_request_to_origin_server true +``` + +#### allow_proxy_protocol + +This will start LittleProxy allowing the PROXY protocol. + +```bash +$ ./run.bash --allow_proxy_protocol true +``` + +#### send_proxy_protocol + +This will start LittleProxy sending the PROXY protocol header. + +```bash +$ ./run.bash --send_proxy_protocol true +``` + +#### client_to_proxy_worker_threads + +This will start LittleProxy with the specified number of client to proxy worker threads. + +```bash +$ ./run.bash --client_to_proxy_worker_threads 10 +``` + +#### proxy_to_server_worker_threads + +This will start LittleProxy with the specified number of proxy to server worker threads. + +```bash +$ ./run.bash --proxy_to_server_worker_threads 10 +``` + +#### acceptor_threads + +This will start LittleProxy with the specified number of acceptor threads. + +```bash +$ ./run.bash --acceptor_threads 10 +``` + + +#### server + +This will start LittleProxy as a server, i.e it will not stop, until you stop the process running it (via a `kill`kill command). + +```bash +$ ./run.bash --server +``` + +#### Help + +This will print the help message: + +```bash +$ ./run.bash --help +``` + +## Embedding in your own projects + +You can embed LittleProxy in your own projects through Maven with the following : ``` - xyz.rogfam + io.github.littleproxy littleproxy - 2.0.0-beta-5 + 2.9.1 ``` Or with Gradle like this -`compile "xyz.rogfam:littleproxy:2.0.0-beta-5"` +`implementation "io.github.littleproxy:littleproxy:2.9.1"` Once you've included LittleProxy, you can start the server with the following: ```java HttpProxyServer server = - DefaultHttpProxyServer.bootstrap() - .withPort(8080) - .start(); + DefaultHttpProxyServer.bootstrap() + .withPort(8080) + .start(); +``` + +### Shared server connection pool configuration + +LittleProxy supports pluggable shared server connection pools. This allows multiple client +connections to reuse upstream connections and helps prevent connection explosion under load. + +```java +HttpProxyServer server = + DefaultHttpProxyServer.bootstrap() + .withPort(8080) + .withSharedServerConnectionPool(true) + .withMaxConnections(500) + .withMaxConnectionsPerHost(50) + .withPoolIdleTimeout(Duration.ofSeconds(30)) + .start(); ``` +Available pool types: + +- `CONCURRENT_MAP`: lightweight default implementation + +#### Pool metrics + +Each server connection pool implementation exposes runtime metrics through `PoolMetrics` +(`totalConnections`, `activeConnections`, `idleConnections`, `borrowCount`, `returnCount`, +`evictionCount`, `validationFailureCount`). + +Implementation details: + +- `CONCURRENT_MAP` + - `totalConnections`: current number of tracked pooled server connections + - `activeConnections`: `totalConnections - idleConnections` + - `idleConnections`: connections currently waiting in the available queue + - `borrowCount` / `returnCount`: incremented on successful borrow/return + - `evictionCount`: incremented when idle-timeout eviction removes connections + - `validationFailureCount`: incremented when validation rejects a pooled connection + +These metrics are useful for capacity tuning, behavior validation during load tests, and +troubleshooting connection reuse. + To intercept and manipulate HTTPS traffic, LittleProxy uses a man-in-the-middle (MITM) manager. LittleProxy's default implementation (`SelfSignedMitmManager`) has a fairly limited feature set. For greater control over certificate impersonation, -browser trust, the TLS handshake, and more, use a the LittleProxy-compatible MITM extension: +browser trust, the TLS handshake, and more, use a LittleProxy-compatible MITM extension: - [LittleProxy-mitm](https://github.com/ganskef/LittleProxy-mitm) - A LittleProxy MITM extension that aims to support every Java platform including Android - [mitm](https://github.com/lightbody/browsermob-proxy/tree/master/mitm) - A LittleProxy MITM extension that supports elliptic curve cryptography and custom trust stores -To filter HTTP traffic, you can add request and response filters using a +To filter HTTP traffic, you can add request and response filters using a `HttpFiltersSource(Adapter)`, for example: ```java HttpProxyServer server = - DefaultHttpProxyServer.bootstrap() - .withPort(8080) - .withFiltersSource(new HttpFiltersSourceAdapter() { - public HttpFilters filterRequest(HttpRequest originalRequest, ChannelHandlerContext ctx) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - // TODO: implement your filtering here - return null; - } + DefaultHttpProxyServer.bootstrap() + .withPort(8080) + .withFiltersSource(new HttpFiltersSourceAdapter() { + public HttpFilters filterRequest(HttpRequest originalRequest, ChannelHandlerContext ctx) { + return new HttpFiltersAdapter(originalRequest) { + @Override + public HttpResponse clientToProxyRequest(HttpObject httpObject) { + // TODO: implement your filtering here + return null; + } - @Override - public HttpObject serverToProxyResponse(HttpObject httpObject) { - // TODO: implement your filtering here - return httpObject; + @Override + public HttpObject serverToProxyResponse(HttpObject httpObject) { + // TODO: implement your filtering here + return httpObject; + } + }; } - }; - } - }) - .start(); + }) + .start(); ``` -Please refer to the Javadoc of `org.littleshoot.proxy.HttpFilters` to see the -methods you can use. +Please refer to the Javadoc of `org.littleshoot.proxy.HttpFilters` to see the +methods you can use. -To enable aggregator and inflater you have to return a value greater than 0 in -your `HttpFiltersSource#get(Request/Response)BufferSizeInBytes()` methods. This -provides to you a `FullHttp(Request/Response)' with the complete content in your -filter uncompressed. Otherwise you have to handle the chunks yourself. +To enable aggregator and inflater you have to return a value greater than 0 in +your `HttpFiltersSource#get(Request/Response)BufferSizeInBytes()` methods. This +provides to you a `FullHttp(Request/Response)` with the complete content in your +filter uncompressed. Otherwise, you have to handle the chunks yourself. ```java @Override - public int getMaximumResponseBufferSizeInBytes() { - return 10 * 1024 * 1024; - } +public int getMaximumResponseBufferSizeInBytes() { + return 10 * 1024 * 1024; +} ``` -This size limit applies to every connection. To disable aggregating by URL at -*.iso or *dmg files for example, you can return in your filters source a filter +This size limit applies to every connection. To disable aggregating by URL at +*.iso or *dmg files for example, you can return in your filters source a filter like this: ```java @@ -106,39 +459,43 @@ return new HttpFiltersAdapter(originalRequest, serverCtx) { } }; ``` -This enables huge downloads in an application, which regular handles size -limited `FullHttpResponse`s to modify its content, HTML for example. +This enables huge downloads in an application, which regular handles size +limited `FullHttpResponse`s to modify its content, HTML for example. -A proxy server like LittleProxy contains always a web server, too. If you get an -URI without scheme, host and port in `originalRequest` it's a direct request to -your proxy. You can return a `HttpFilters` implementation which answers +A proxy server like LittleProxy contains always a web server, too. If you get a +URI without scheme, host and port in `originalRequest` it's a direct request to +your proxy. You can return a `HttpFilters` implementation which answers responses with HTML content or redirects in `clientToProxyRequest` like this: ```java +import java.nio.charset.StandardCharsets; + +import static java.nio.charset.StandardCharsets.UTF_8; + public class AnswerRequestFilter extends HttpFiltersAdapter { - private final String answer; - - public AnswerRequestFilter(HttpRequest originalRequest, String answer) { - super(originalRequest, null); - this.answer = answer; - } - - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - ByteBuf buffer = Unpooled.wrappedBuffer(answer.getBytes("UTF-8")); - HttpResponse response = new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK, buffer); - HttpHeaders.setContentLength(response, buffer.readableBytes()); - HttpHeaders.setHeader(response, HttpHeaders.Names.CONTENT_TYPE, "text/html"); - return response; - } + private final String answer; + + public AnswerRequestFilter(HttpRequest originalRequest, String answer) { + super(originalRequest, null); + this.answer = answer; + } + + @Override + public HttpResponse clientToProxyRequest(HttpObject httpObject) { + ByteBuf buffer = Unpooled.wrappedBuffer(answer.getBytes(UTF_8)); + HttpResponse response = new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK, buffer); + HttpHeaders.setContentLength(response, buffer.readableBytes()); + HttpHeaders.setHeader(response, HttpHeaders.Names.CONTENT_TYPE, "text/html"); + return response; + } } ``` -On answering a redirect, you should add a Connection: close header, to avoid +On answering a redirect, you should add a Connection: close header, to avoid blocking behavior: ```java HttpHeaders.setHeader(response, Names.CONNECTION, Values.CLOSE); ``` -With this trick, you can implement an UI to your application very easy. +With this trick, you can implement a UI to your application very easy. If you want to create additional proxy servers with similar configuration but listening on different ports, you can clone an existing server. The cloned @@ -149,8 +506,94 @@ stopped, all are stopped. existingServer.clone().withPort(8081).start() ``` + +### Logging Activity Tracker + +LittleProxy includes a `LoggingActivityTracker` that can log detailed information about each request and response handled by the proxy. It supports multiple standard log formats, which can be useful for integration with log analysis tools. + +To use it, wrap your functionality or simply add it to your server bootstrap: + +```java +import org.littleshoot.proxy.extras.ActivityLogger; +import org.littleshoot.proxy.extras.LoggingActivityTracker; +import org.littleshoot.proxy.extras.LogFormat; + +// ... + +DefaultHttpProxyServer.bootstrap() + . + +withPort(8080) + . + +plusActivityTracker(new ActivityLogger(LogFormat.CLF)) // Use Common Log Format + . + +start(); +``` + +#### Supported Log Formats + +The `LogFormat` enum provides several standard formats: + +* **`CLF` (Common Log Format)**: The standard NCSA Common log format. + * Example: `127.0.0.1 - - [24/Dec/2025:00:00:00 +0000] "GET /index.html HTTP/1.1" 200 1234` +* **`ELF` (Extended Log Format)**: Uses the NCSA Combined Log Format, which includes Referer and User-Agent. + * Example: `127.0.0.1 - - [date] "GET /..." 200 123 "http://referer" "Mozilla/5.0"` +* **`W3C`**: A standard W3C Extended Log File Format (space-separated fields). + * Example: `2025-12-24 00:00:00 127.0.0.1 GET /index.html 200 1234 "Mozilla/5.0"` +* **`JSON`**: Structured logging in JSON format, ideal for modern log aggregators (ELK, Splunk, etc.). Includes duration. + * Example: `{"timestamp":"...","client_ip":"127.0.0.1","method":"GET","duration":15,...}` +* **`LTSV` (Labeled Tab-Separated Values)**: Efficient, human-readable, and machine-parsable format. + * Example: `time:2025-...\thost:127.0.0.1\tmethod:GET\t...` +* **`CSV` (Comma-Separated Values)**: Standard CSV format for easy import into spreadsheets. + * Example: `"timestamp","127.0.0.1","GET",...` +* **`SQUID`**: Squid native access log format. Useful for tools expecting Squid logs. +* **`HAPROXY`**: A format mimicking HAProxy's HTTP logging, focusing on timing and status. + For examples of configuring logging, see [src/test/resources/log4j.xml](src/test/resources/log4j.xml). +#### Customizing Logging Configuration + +You can customize the `log4j.xml` configuration to control how logs are output. This is particularly useful for separating access logs from system logs. + +**1. Standard Output (Default)** + +To print access logs to the console without standard Log4j prefixes (timestamps, thread names, etc.), use a specific appender for the tracker: + +```xml + + + + + + + + + + + +``` + +**2. Dedicated Access Log File** + +To write access logs to a separate file (e.g., `access.log`) and exclude them from the main log, use a `FileAppender`: + +```xml + + + + + + + + + + + + +``` + If you have questions, please visit our Google Group here: https://groups.google.com/forum/#!forum/littleproxy2 @@ -159,7 +602,19 @@ https://groups.google.com/forum/#!forum/littleproxy2 accepting posts from new users. But it's still a great resource if you're searching for older answers.) -To subscribe, send an e-mail to [LittleProxy2+subscribe@googlegroups.com](mailto:LittleProxy2+subscribe@googlegroups.com). +To subscribe, send an e-mail to [LittleProxy2+subscribe@googlegroups.com](mailto:LittleProxy2+subscribe@googlegroups.com). + +## Performance and Logging Guide + +For comprehensive information on logging performance optimization, including: + +- **Synchronous vs Asynchronous Logging**: Performance comparison and use cases +- **Activity Log Formats**: All supported formats (CLF, ELF, JSON, SQUID, W3C, LTSV, CSV, HAPROXY) +- **Log Filtering**: BurstFilter and custom filter implementations +- **Groovy Script Filters**: Dynamic filtering with sampling and conditional logic +- **Best Practices**: Production-ready recommendations and troubleshooting + +Please see our **[Performance and Logging Guide](PERFORMANCE_AND_LOGGING.md)**. Acknowledgments --------------- diff --git a/RELEASE_NOTES.md b/RELEASE_NOTES.md index e286bc81..e2ec212a 100644 --- a/RELEASE_NOTES.md +++ b/RELEASE_NOTES.md @@ -1,5 +1,327 @@ # Release Notes +- 2.10.0 (Under construction, https://github.com/LittleProxy/LittleProxy/milestone/54?closed=1) + - TBD + +- 2.9.1 (28.08.2026, https://github.com/LittleProxy/LittleProxy/milestone/53?closed=1) + - Bump Netty from 4.2.16.Final to 4.2.17.Final (#782) + - Bump Selenium from 4.46.0 to 4.48.0 (#784) (#792) + - Bump Guava from 33.6.0-jre to 33.7.1-jre (#787) (#788) + +- 2.9.0 (06.08.2026, https://github.com/LittleProxy/LittleProxy/milestone/52?closed=1) + - #737 Fix order of handlers to support proxy protocol decoding for inbound requests (#729) by Krasimir Marinov + - fix "IllegalReferenceCountException: null" on ClientToProxyConnection (#625) by James Baldassari + - #71 Fix assumption about CONNECT in MITM mode should use SSL (#717) by Charles Lescot + - #77 Fix: CONNECT response not returned to HttpFilters (#721) by Charles Lescot + - #57 Fix: Upstream Socks Proxy Not Authenticating (#719) by Charles Lescot + - Bump netty.version from 4.2.15.Final to 4.2.16.Final (#770) + - Bump org.seleniumhq.selenium:selenium-java from 4.45.0 to 4.46.0 (#771) + +- 2.8.0 (27.06.2026, https://github.com/LittleProxy/LittleProxy/milestone/51?closed=1) + - Fix unsupported `com.lmax.disruptor` scope (#744) by Alexey Venderov + - Bump selenium from 4.41.0 to 4.45.0 + - Bump jackson from 2.21.2 to 2.22.0 + - Bump netty from 4.2.12.Final to 4.2.15.Final (#761) + - Bump dnsjava from 3.6.4 to 3.6.5 (#754) + +- 2.7.0 (06.04.2026, https://github.com/LittleProxy/LittleProxy/milestone/50?closed=1) + - fix WebSocket proxying by Andrei Solntsev + - add WebSocket frame observation hook by Andrei Solntsev + - #464 fix authentication bug when forwarding request with "Proxy-Authorization" header to a chained proxy (#691) by Charles Lescot + - quickly resolve localhost name (#698) by Andrei Solntsev + - #56 #439 fix timeout detection logic in ClientToProxyConnection (#699) by Charles Lescot + - #680 fix issue with generated jks and cert at root (#681) by Charles Lescot + - migration to log4j2 and async logging available (#684) by Charles Lescot + - enhance ActivityTracker interface with new lifecycle methods (#708) by Charles Lescot + - add unit test demonstrating how to do "internal redirect" (#68) (#718) by Charles Lescot + +- 2.6.0 (19.01.2026, https://github.com/LittleProxy/LittleProxy/milestone/49?closed=1) + - Feature/activity tracker logging (#668) by Charles Lescot + - add explicit logging support with "activity_log_format" option (CLI and config file) supporting CLF, ELF, JSON, LTSV, CSV, SQUID, HAPROXY formats (#668) by Charles Lescot + - move littleproxy.properties and log4j.xml in standard maven location (#667) by Charles Lescot + - fix: LittleProxy is not starting with jdk higher than jdk 11. (#669) by Charles Lescot + - migrate LittleProxy own tests from MockServer to WireMock (#670) by Charles Lescot + - bump Selenium from 4.39.0 to 4.40.0 (#676) + +- 2.5.0 (19.12.2025, https://github.com/LittleProxy/LittleProxy/milestone/48?closed=1) + - enhance documentation with run.bash script options (#659) by Charles Lescot + - add the --server option to Launch proxy as a server (#660) by Charles Lescot + - Add the --log-config option to launcher (#661) by Charles Lescot + - Add the --config option to launcher (#662) by Charles Lescot + - add options "--name", "--address", "--nic", "--allow_local_only", "--authenticate_ssl_clients", "--transparent", "--throttling", "--allow_requests_to_origin_server" (#666) by Charles Lescot + - add options "--proxy_alias", "--allow_proxy_protocol option", "--send_proxy_protocol", "--client_to_proxy_worker_threads", "--proxy_to_server_worker_threads" options (#666) by Charles Lescot + - remove unused and deprecated NetworkUtils.java (#663) by Charles Lescot + - remove "withListenOnAllAddresses" method from DefaultHttpProxyServer.java and HttpProxyServerBootstrap.java (#664) by Charles Lescot + +- 2.4.7 (15.12.2025, https://github.com/LittleProxy/LittleProxy/milestone/47?closed=1) + - Bump Netty from 4.2.7.Final to 4.2.9.Final (#655) (#658) + - bump Selenium from 4.38.0 to 4.39.0 (#653) + +- 2.4.6 (28.10.2025, https://github.com/LittleProxy/LittleProxy/milestone/46?closed=1) + - Bump Netty from 4.2.5.Final to 4.2.7.Final + - Bump Log4j from 2.25.1 to 2.25.2 + - Bump Selenium from 4.35.0 to 4.38.0 + - remove most Guava usages + +- 2.4.5 (08.09.2025, https://github.com/LittleProxy/LittleProxy/milestone/45?closed=1) + - fix compile warnings (#610) by leeyazhou + - bump Selenium from 4.34.0 to 4.35.0 (#604) + - bump Netty from 4.2.3.Final to 4.2.5.Final (#611) (#605) + - bump Jackson from 2.19.2 to 2.20.0 (#609) + +- 2.4.4 (10.08.2025, https://github.com/LittleProxy/LittleProxy/milestone/44?closed=1) + - Bump Netty from 4.2.2.Final to 4.2.3.Final (#596) + - Bump Jackson from 2.19.1 to 2.19.2 (#598) + - Bump Log4j from 2.25.0 to 2.25.1 (#595) + +- 2.4.3 (02.07.2025, https://github.com/LittleProxy/LittleProxy/milestone/43?closed=1) + - added "autoStop" property for ServerGroup (#590) -- thanks to Alex Panchenko + - fix problem when proxy is stopped quickly after starting (#590) -- thanks to Alex Panchenko + - Bump selenium from 4.32.0 to 4.34.0 (#588) + - Bump netty from 4.2.1.Final to 4.2.2.Final (#582) + - Bump jackson from 2.19.0 to 2.19.1 (#584) + - Bump log4j from 2.24.3 to 2.25.0 (#585) + +- 2.4.2 (08.05.2025, https://github.com/LittleProxy/LittleProxy/milestone/42?closed=1) + - fix memory leak in ProxyToServerConnection (#573) + - Bump selenium from 4.31.0 to 4.32.0 (#576) + - Bump netty from 4.2.0.Final to 4.2.1.Final (#577) + +- 2.4.1 (22.04.2025, https://github.com/LittleProxy/LittleProxy/milestone/41?closed=1) + - Bump netty from 4.1.116.Final to 4.2.0.Final (#550) (#564) + - Bump selenium from 4.27.0 to 4.31.0 (#546) (#560) (#566) + - Bump dnsjava from 3.6.2 to 3.6.3 + - Bump slf4j from 2.0.16 to 2.0.17 (#549) + - Bump jackson from 2.18.2 to 2.18.3 (#554) + +- 2.4.0 (02.01.2025, https://github.com/LittleProxy/LittleProxy/milestone/40?closed=1) + - Migrate nullability annotations from JSR 305 to JSpecify (#533) (#534) + - Bump Netty from 4.1.115.Final to 4.1.116.Final (#530) + - Bump Log4j from 2.24.2 to 2.24.3 (#527) + - Bump Guava from 33.3.1-jre to 33.4.0-jre (#528) + +- 2.3.3 (04.12.2024, https://github.com/LittleProxy/LittleProxy/milestone/39?closed=1) + - Bump Netty from 4.1.114.Final to 4.1.115.Final (#521) + - Bump Selenium from 4.26.0 to 4.27.0 (#524) + - Bump Jackson from 2.18.1 to 2.18.2 (#525) + +- 2.3.2 (06.11.2024, https://github.com/LittleProxy/LittleProxy/milestone/38?closed=1) + - Expose proxy to server ctx to access client address -- thanks to Teodora Kostova (#520) + - Bump Netty from 4.1.113.Final to 4.1.114.Final + - Bump Selenium from 4.25.0 to 4.26.0 + +- 2.3.1 (30.09.2024, https://github.com/LittleProxy/LittleProxy/milestone/37?closed=1) + - Bump org.seleniumhq.selenium:selenium-java from 4.24.0 to 4.25.0 (#492) + - Bump dnsjava:dnsjava from 3.6.1 to 3.6.2 (#491) + - Bump org.apache.logging.log4j:log4j-core from 2.23.1 to 2.24.1 (#489) (#497) + +- 2.3.0 (06.09.2024, https://github.com/LittleProxy/LittleProxy/milestone/36?closed=1) + - #487 remove UDP protocol support (#488) + - Bump Netty from 4.1.112.Final to 4.1.113.Final (#486) + +- 2.2.4 (04.09.2024, https://github.com/LittleProxy/LittleProxy/milestone/35?closed=1) + - Bump Selenium from 4.22.0 to 4.24.0 + - Bump Netty from 4.1.111.Final to 4.1.112.Final + - Bump dnsjava from 3.5.3 to 3.6.1 + +- 2.2.3 (21.06.2024, https://github.com/LittleProxy/LittleProxy/milestone/34?closed=1) + - Bump selenium from 4.21.0 to 4.22.0 (#433) + - #37 fix ClassCastException: "PooledUnsafeDirectByteBuf cannot be cast to HttpObject" (#434) + +- 2.2.2 (12.06.2024, https://github.com/LittleProxy/LittleProxy/milestone/33?closed=1) + - Bump selenium from 4.20.0 to 4.21.0 + - Bump jackson from 2.17.0 to 2.17.1 + - Bump netty from 4.1.109.Final to 4.1.111.Final + +- 2.2.1 (25.04.2024, https://github.com/LittleProxy/LittleProxy/milestone/32?closed=1) + - Bump Netty from 4.1.107.Final to 4.1.109.Final + - Bump Selenium from 4.18.1 to 4.20.0 + - Bump Jackson from 2.16.1 to 2.17.0 + - Bump Slf4j from 2.0.12 to 2.0.13 + - Bump Log4j from 2.23.0 to 2.23.1 + - Bump Guava from 33.0.0-jre to 33.1.0-jre + +- 2.2.0 (22.02.2024, https://github.com/LittleProxy/LittleProxy/milestone/31?closed=1) + - Move the project from groupId "xyz.rogfam" to "io.github.littleproxy" + - Migrate from JUnit4/Hamcrest to JUnit5/AssertJ (#373) + - Bump Netty from 4.1.106.Final to 4.1.107.Final (#376) + - Bump Selenium from 4.17.0 to 4.18.1 (#378) (#379) + - Bump log4j from 2.22.1 to 2.23.0 (#381) + +- 2.1.2 (07.02.2024, https://github.com/LittleProxy/LittleProxy/milestone/30?closed=1) + - Refactoring & code cleanup & setup IDEA inspections (#370) (#371) + - Bump Netty from 4.1.103.Final to 4.1.106.Final + - Bump slf4j from 2.0.9 to 2.0.12 + - Bump log4j from 2.22.0 to 2.22.1 (#357) + - Bump Selenium from 4.16.1 to 4.17.0 (#367) + - Bump Guava from 32.1.3-jre to 33.0.0-jre + +- 2.1.1 (15.12.2023, https://github.com/LittleProxy/LittleProxy/milestone/29?closed=1) + - Bump Netty from 4.1.101.Final to 4.1.103.Final #345 #346 + - Bump selenium from 4.15.0 to 4.16.1 #343 #344 + - Bump log4j from 2.21.1 to 2.22.0 #336 + - Bump commons-lang3 from 3.13.0 to 3.14.0 #338 + +- 2.1.0 (20.11.2023, https://github.com/LittleProxy/LittleProxy/milestone/28?closed=1) + - Upgrade from Java 8 to Java 11+ + - Bump Selenium from 4.13.0 to 4.15.0 + - Bump Netty from 4.1.99.Final to 4.1.101.Final + - Bump Jackson from 2.15.2 to 2.16.0 + - Bump Log4j from 2.20.0 to 2.21.1 + - Bump Guava from 32.1.2-jre to 32.1.3-jre + +- 2.0.22 (08.10.2023, https://github.com/LittleProxy/LittleProxy/milestone/27?closed=1) + - #35 Fix Websocket race condition while protocol switching -- thanks to Craig Andrews for PR #308 + - #307 bump Netty from 4.1.98.Final to 4.1.99.Final + +- 2.0.21 (27.09.2023, https://github.com/LittleProxy/LittleProxy/milestone/26?closed=1) + - #301 Always use a new connection for websockets -- thanks to Craig Andrews + - #299 fix problem with filtering proxy Authorization header -- thanks to Matthias Kraaz for PR #304 + - #297 #306 Bump org.seleniumhq.selenium:selenium-java from 4.12.0 to 4.13.0 + - #303 Bump `netty.version` from 4.1.97.Final to 4.1.98.Final. + +- 2.0.20 (04.09.2023, https://github.com/LittleProxy/LittleProxy/milestone/25?closed=1) + - #295 #131 fix memory leak "LEAK: ByteBuf.release() was not called..." -- thanks to Sujit Joshi for the fix + - #284 #291 Bump netty.version from 4.1.95.Final to 4.1.97.Final + - #286 #293 Bump org.seleniumhq.selenium:selenium-java from 4.10.0 to 4.12.0 + - #285 Bump org.apache.commons:commons-lang3 from 3.12.0 to 3.13.0 + - #287 Bump com.google.guava:guava from 32.1.1-jre to 32.1.2-jre + - #294 Bump slf4j.version from 2.0.7 to 2.0.9 + +- 2.0.19 (22.07.2023, https://github.com/LittleProxy/LittleProxy/milestone/24?closed=1) + - #283 fix memory leak: On proxy connection unregister, unregister downstream channels - thanks to Craig Andrews + - #274 Bump Selenium from 4.9.1 to 4.10.0 (see https://github.com/SeleniumHQ/selenium) + - #266 Bump Jackson from 2.15.1 to 2.15.2 + - #281 Bump guava from 32.0.0-jre to 32.1.1-jre + - #282 Bump Netty from 4.1.93.Final to 4.1.95.Final + - #264 Migrate Jetty 9 to Jetty 11 - thanks to Valery Yatsynovich + +- 2.0.18 (29.05.2023, https://github.com/LittleProxy/LittleProxy/milestone/23?closed=1) + - Bump Selenium from 4.8.3 to 4.9.1 (see https://github.com/SeleniumHQ/selenium) + - #242 Bump Netty from 4.1.90.Final to 4.1.93.Final + - Bump Jackson from 2.14.2 to 2.15.1 + - Bump guava from 31.1-jre to 32.0.0-jre + +- 2.0.17 (01.04.2023, https://github.com/LittleProxy/LittleProxy/milestone/22?closed=1) + - #235 Bump netty.version from 4.1.89.Final to 4.1.90.Final + - bump Jackson from 2.13.4 to latest 2.14.2 (fixes several CVEs) + - #236 Bump slf4j.version from 2.0.6 to 2.0.7 + - #241 Bump selenium-java from 4.8.1 to 4.8.3 + +- 2.0.16 (27.02.2023, https://github.com/LittleProxy/LittleProxy/milestone/21?closed=1) + - rename "master" branch to "main" + - #207 Remove redundant file generated by unit test -- thanks to Valery Yatsynovich + - #206 Export certificate to generated by SelfSignedMitmManager KeyStore directory -- thanks to Valery Yatsynovich + - Bump slf4j.version from 2.0.5 to 2.0.6 + - Bump log4j-core from 2.19.0 to 2.20.0 + - Bump selenium-java from 4.7.1 to 4.8.1 + - Bump netty.version from 4.1.86.Final to 4.1.89.Final + +- 2.0.15 (14.12.2022, https://github.com/LittleProxy/LittleProxy/milestone/20?closed=1) + - Bump netty-codec-haproxy from 4.1.85.Final to 4.1.86.Final + - Bump selenium-java from 4.6.0 to 4.7.1 + - Bump slf4j.version from 2.0.4 to 2.0.5 + - Bump httpclient from 4.5.13 to 4.5.14 + +- 2.0.14 (21.11.2022, https://github.com/LittleProxy/LittleProxy/milestone/19?closed=1) + - #184 Respectful KeyStore file path while generating certs by `SelfSignedMitmManager` -- thanks to Valery Yatsynovich + - #187 CI: run build on all major OS-s -- thanks to Valery Yatsynovich + - #183 Bump netty from 4.1.82.Final to 4.1.85.Final -- thanks to Valery Yatsynovich for fixing tests after upgrading Netty. + - #189 Bump slf4j.version from 2.0.3 to 2.0.4 + - Bump jackson-databind from 2.13.2.2 to 2.13.4 + - #191 Bump dnsjava from 3.5.1 to 3.5.2 + +- 2.0.13 (04.10.2022) + - #170 restore transitive dependencies in generated pom -- thanks to Mateusz Pietryga for PR #171 + - Bump slf4j from 2.0.1 to 2.0.3 + - Bump selenium-java from 4.4.0 to 4.5.0 + +- 2.0.12 (23.09.2022) + - #145 Restore Keep-Alive value when filtering short-circuit response -- thanks to krlvm for PR + - Bump netty from 4.1.79.Final to 4.1.82.Final + - Bump slf4j from 1.7.36 to 2.0.1 + - Bump log4j-core from 2.18.0 to 2.19.0 + +- 2.0.11 (13.08.2022) + - #131 fix memory leak: release byte buffer when closing request - see PR #141 + - #142 fix some "modify response" problem, see https://github.com/adamfisk/LittleProxy/issues/359 + - #144 HTTP CONNECT can't be Keep-Alive - thanks Michel Belleau for PR #144 + +- 2.0.10 (20.07.2022) + - #135 Bump netty.version from 4.1.77.Final to 4.1.79.Final + - #132 Bump selenium-java from 4.1.4 to 4.3.0 + - #118 Bump dnsjava from 3.5.0 to 3.5.1 + +- 2.0.9 (10.05.2022) + - #115 reverted to maven-shade-plugin 3.2.4 (because 3.3.0 generated artifact without compile/runtime dependencies) + +- 2.0.8 (06.05.2022) + - #26 fixed TLS 1.3 handshake bug -- thanks Dan Powell for PR https://github.com/LittleProxy/LittleProxy/pull/26 + - Bumped log4j-core from 2.17.0 to 2.17.2 + - Bumped netty from 4.1.71 to 4.1.76 + - Bumped slf4j from 1.7.30 to 1.7.36 + - Bumped jackson from 2.11.3 to 2.12.6.1 + - Bumped guava from 30.1-jre to 31.1-jre + - Bumped commons-cli from 1.4 to 1.5.0 + - Relocated slf4j-log4j to slf4j-reload4j + - moved the project to https://github.com/LittleProxy/LittleProxy + - moved CI from Travis to https://github.com/LittleProxy/LittleProxy/actions + +- 2.0.7 (21.12.2021) + - Bumped log4j-core from 2.16.0 to 2.17.0 + +- 2.0.6 + - Use single Hamcrest dependency in tests + - Improve logging performance + - Bumped netty-codec from 4.1.63.Final to 4.1.68.Final + - Bump netty-codec-http from 4.1.68.Final to 4.1.71.Final + - Bumped log4j-core from 2.14.0 to 2.16.0 + - Added public key file + +- 2.0.5 + - Bumped jetty-server from 9.4.34.v20201102 to 9.4.41.v20210516. + +- 2.0.4 + - Android compatibility fix (PR #76) + - Fix NoSuchElementException when switching protocols to WebSocket (PR #78) + - Prevent NullPointerException in ProxyUtils::isHEAD (PR #79) + - Fixes in ThrottlingTest, Upgrade to Netty 4.1.63.Final (PR #65) + - Fix NPEs in getReadThrottle and getWriteThrottle when globalTrafficShapingHandler is null (PR #80) + +- 2.0.3 + - Upgrade guava to 30.1 + - Threads are now set as daemon (not user, which is the default) threads so the JVM exits as expected when all other threads stop. + - Close thread pool if proxy fails to start + +- 2.0.2 + - Support for WebSockets with MITM in transparent mode + - Support for per request conditional MITM + +- 2.0.1 + - Removed beta tag from version + - Updated various dependency versions + - Re-ordered the release notes so the newest stuff is at the top + +- 2.0.0-beta-6 + - Cleaned up old code to conform with newer version of Netty + - Deprecated UDT support because it's deprecated in Netty + - Removed performance test code because it seems to be confusing GitHub into thinking that this is a PHP project + +- 2.0.0-beta-5 + - Treat an upstream SOCKS proxy as if it is the origin server + - Fixed memoryLeak in ClientToProxyConnection + +- 2.0.0-beta-4 + - Allow users to set their own server group within the bootstrap helper + - Added support for chained SOCKS proxies + +- 2.0.0-beta-3 + - Upgraded Netty, guava, Hamcrest, Jetty, Selenium, Apache commons cli and lang3 + - Upgrade Maven plugins to the latest versions + +- 2.0.0-beta-2 + - Added support for proxy protocol. See https://www.haproxy.com/blog/haproxy/proxy-protocol/ and https://www.haproxy.org/download/1.8/doc/proxy-protocol.txt for protocol details. + - 2.0.0-beta-1 - New Maven coordinates - Moved from Java 7 to 8 @@ -7,24 +329,3 @@ - **Breaking change:** Made client details available to ChainedProxyManager - Refactored MITM manager to accept engine with user-defined parameters - Added ability to load keystore from classpath - -- 2.0.0-beta-2 - - Added support for proxy protocol. See https://www.haproxy.com/blog/haproxy/proxy-protocol/ and https://www.haproxy.org/download/1.8/doc/proxy-protocol.txt for protocol details. - -- 2.0.0-beta-3 - - Upgraded Netty, guava, Hamcrest, Jetty, Selenium, Apache commons cli and lang3 - - Upgrade Maven plugins to the latest versions - -- 2.0.0-beta-4 - - Allow users to set their own sesrvergroup within the bootstrap helper - - Added support for chained SOCKS proxies - -- 2.0.0-beta-5 - - Treat an upstream SOCKS proxy as if it is the origin server - - Fixed memoryLeak in ClientToProxyConnection - -- 2.0.0-beta-6 - - Cleaned up old code to conform with newer version of Netty - - Deprecated UDT support because it's deprecated in Netty - - Removed performance test code because it seems to be confusing GitHub into thinking that this is a PHP project. - \ No newline at end of file diff --git a/config/littleproxy.properties b/config/littleproxy.properties new file mode 100644 index 00000000..29e089c4 --- /dev/null +++ b/config/littleproxy.properties @@ -0,0 +1,3 @@ +name=MyLittleProxy +idle_connection_timeout=40 +activity_log_format=ELF \ No newline at end of file diff --git a/docs/pooling-features.md b/docs/pooling-features.md new file mode 100644 index 00000000..81d5f778 --- /dev/null +++ b/docs/pooling-features.md @@ -0,0 +1,677 @@ +# Server Connection Pooling — Full Feature Set + +This document describes **all** features on this branch that are not present in the +main branch. The branch introduces a complete shared server connection pooling +infrastructure (based on the design from PR #724) and extends it with MITM / HTTPS +upstream support. + +--- + +## Background + +On the main branch, every client connection gets its own dedicated +`ProxyToServerConnection`. When client A and client B both connect to +`http://example.com`, two separate TCP sockets are opened to the same server. +When the traffic is MITM'd HTTPS, the upstream TLS connection is tied to the +client's session for its entire lifetime and discarded on disconnect. + +This branch addresses the connection explosion problem by introducing a **shared +server connection pool**, and extends it to HTTPS upstream connections intercepted +via MITM. + +--- + +## Base Connection Pooling (PR #724 Infrastructure) + +The entire pool infrastructure is new on this branch. On the main branch, all +connections are created with `ProxyToServerConnection.create(...)` and tracked in a +per-client `serverConnectionsByHostAndPort` map. + +### New Interfaces and Classes + +| Type | File | Purpose | +|---|---|---| +| Interface | `ServerConnectionPool` | Contract for pooling server connections | +| Config | `ServerConnectionPoolConfig` | Configuration bean (pool type, sizes, timeouts) | +| Config | `DefaultHttpProxyServerConfig` | Server-level config carrier that includes pool config | +| Metrics | `PoolMetrics` | Active/idle/borrow/return/eviction counters | +| Model | `PendingRequest` | Tracks a request awaiting a response (for HTTP pipelining) | +| Enum | `ServerConnectionPoolType` | `CONCURRENT_MAP` | + +### Pool Implementations + +Single backend implementation: + +| Pool type | Class | Approach | Dependencies | +|---|---|---|---| +| `CONCURRENT_MAP` | `ConcurrentMapServerConnectionPool` | `ConcurrentHashMap` + per-host `Queue` of available connections | None (pure Netty/Java) | + +The implementation: +- Implements `getOrCreateConnection(host, chainedProxyAddr, client, filters, request)` + + `releaseConnection(connection)` + `removeConnection(connection)` + `closeAll()` +- Tracks pending requests per channel for HTTP pipelining support +- Enforces per-host and global connection limits +- Supports idle timeout eviction and optional connection validation on borrow +- Exposes `getMetrics()` returning active/idle/total connections and cumulative operation counts + +### `ServerConnectionPool` Interface + +```java +ProxyToServerConnection getOrCreateConnection( + String serverHostAndPort, + @Nullable InetSocketAddress chainedProxyAddress, + ClientToProxyConnection clientConnection, + HttpFilters initialFilters, + HttpRequest initialHttpRequest); + +void releaseConnection(ProxyToServerConnection connection); +void removeConnection(ProxyToServerConnection connection); +void registerPendingRequest(Channel, ClientToProxyConnection, HttpRequest, HttpFilters); +PendingRequest removePendingRequest(Channel); +PendingRequest peekPendingRequest(Channel); +void drainPendingRequests(Channel); +void closeAll(); + +PoolMetrics getMetrics(); + +// Computes a compound pool key: "host:port:" +default String computePoolKey(String serverHostAndPort, + @Nullable InetSocketAddress chainedProxyAddress); +``` + +Pool keys incorporate the chained proxy address, so connections through different +upstream proxies are isolated even when the target host is the same. + +### How Plain HTTP Pooling Works + +In `ClientToProxyConnection.doReadHTTPInitial()`, the selection logic was changed +from a simple `serverConnectionsByHostAndPort.get(serverHostAndPort)` to a decision +tree: + +```java +boolean useSharedPool = + usePool + && !ProxyUtils.isCONNECT(httpRequest) + && !isTunneling() + && !isMitming() + && !ProxyUtils.isSwitchingToWebSocketProtocol(httpRequest); +``` + +When `useSharedPool` is true, `pool.getOrCreateConnection(...)` is called instead of +`ProxyToServerConnection.create(...)`. The pool either returns an idle connection or +creates a new one via `ProxyToServerConnection.createForPool(...)`, which attaches +the `connectionPool` reference so the connection knows it is pool-managed. + +**Excluded from pooling:** Tunneling connections (non-MITM CONNECT, raw TCP tunnels) +and WebSocket protocol upgrades are never pooled. Tunneling has no HTTP request/response +boundary to trigger pool release — the connection stays dedicated for the tunnel's +lifetime. WebSocket connections replace HTTP codecs with raw frame handlers after the +upgrade handshake, so they cannot be returned to the pool. Both always use +`serverConnectionsByHostAndPort` with dedicated `ProxyToServerConnection.create()`. + +On the response path, `markResponseComplete()` calls +`connectionPool.releaseConnection(this)`, returning the connection to the available +queue. HTTP pipelining is handled via `registerPendingRequest` / +`removePendingRequest` — when a pipelined response arrives, the next pending request +is dequeued and routed. + +### Connection Lifecycle for Pooled Connections + +| Event | Non-pooled (main branch) | Pooled (this branch) | +|---|---|---| +| New request arrives | `ProxyToServerConnection.create()` | `pool.getOrCreateConnection()` → creates or borrows | +| Request sent | Stored in `currentHttpRequest` | Registered as `PendingRequest` in pool | +| Response received | Forwarded to client | `removePendingRequest()` → forwarded to correct client | +| Response complete | Connection stays in `AWAITING_INITIAL` | `releaseConnection()` → returned to pool | +| Server disconnect | `clientConnection.serverDisconnected()` | `pool.removeConnection()` + drain pending | +| Client disconnect | `serverConnection.disconnect()` (close) | `serverConnection.disconnect()` → `pool.removeConnection()` via `channelInactive` | + +### New Builder Methods on `HttpProxyServerBootstrap` + +```java +.withSharedServerConnectionPool(boolean) // master switch +.withServerConnectionPoolType(ServerConnectionPoolType) // CONCURRENT_MAP default +.withMaxConnectionsPerHost(int) // default 10 +.withMaxConnections(int) // default 200 +.withPoolIdleTimeout(Duration) // null = no idle eviction +``` + +### How Pool Configuration Reaches the Server + +The configuration flows through three layers: + +1. `DefaultHttpProxyServerBootstrap` stores raw builder fields +2. `build()` creates a `ServerConnectionPoolConfig` and a `DefaultHttpProxyServerConfig` +3. `DefaultHttpProxyServer` reads the config on construction + +Properties file parsing in `DefaultHttpProxyServerBootstrap(Properties props)` maps +each key: + +```properties +use_shared_server_connection_pool=true +server_connection_pool_type=CONCURRENT_MAP +max_connections_per_host=10 +max_total_connections=200 +``` + +### Refactoring: `DefaultHttpProxyServerConfig` + +A new `DefaultHttpProxyServerConfig` class was extracted to carry all server +configuration as a single object. The old pattern of passing individual fields +through the bootstrap → server constructor was replaced with a config object, +enabling cleaner cloning and property-based construction. This class holds +~25 fields including the `ServerConnectionPoolConfig`. + +### Changes to `ClientToProxyConnection` + +- **Request routing**: `doReadHTTPInitial()` now branches on `useSharedPool`. The + pooled path calls `pool.getOrCreateConnection()` with the resolved chained proxy + address. +- **Connection tracking**: Pooled connections are not stored in + `serverConnectionsByHostAndPort` (they live in the pool instead). +- **Backpressure** (`becameSaturated` / `becameWritable`): Both methods now also + check `currentServerConnection` (the transient reference set per request) in + addition to the `serverConnectionsByHostAndPort` values. This is necessary + because pooled connections are not in that map. +- **`serverBecameWriteable`**: Now also checks `currentServerConnection` for + saturation before resuming client reads. +- **`disconnected()`**: Now handles pooled connections by releasing them to the + pool instead of disconnecting them. +- **`recordClientConnected()`**: Now called from `requestRead()` callback rather + than during CONNECT setup, fixing a timing issue with pooled connections. +- **`getClientAddress()`**: Fixed a `ClassCastException` when `remoteAddress()` + returns a non-`InetSocketAddress` type. + +### Changes to `ProxyToServerConnection` + +- **`connectionPool` field**: All pooled connections carry a reference back to + their pool. +- **`createForPool()` static factory**: Creates a connection with pool awareness, + resolving chained proxies and filters the same way as the non-pooled path. +- **`getClientConnection()` method**: Routes responses to the correct client. + For pooled connections, uses the per-request `currentClientConnectionForRequest` + field (set before each write) instead of the constructor-injected + `clientConnection` reference. +- **`write()` method**: When the connection is in pool-managed CONNECT reuse + state (not `DISCONNECTED` but needs a new flow), triggers `connectAndWrite()` + for the CONNECT request. +- **`connectionSucceeded()`**: Explicitly releases the initial request reference + to prevent memory leaks with pooled connections (where `initialRequest` is + retained). +- **`disconnected()`**: Pool-managed connections call + `connectionPool.removeConnection(this)` + `drainPendingRequests(channel)` before + notifying the client. +- **`readRaw()`**: Uses `getClientConnection()` instead of `clientConnection`. +- **All event recording methods** (`recordServerConnected`, `recordServerDisconnected`, + `recordConnectionSaturated`, etc.): Use `getClientConnection()` instead of + `clientConnection` to route activity tracker events to the correct client. +- **`SendProxyProtocolHeader`**: Fixed to handle IPv6 addresses by selecting + `HAProxyProxiedProtocol.TCP6` when either endpoint uses an `Inet6Address`. + +### New `PendingRequest` Class + +A simple holder for a client connection, HTTP request, and filters, stored in a +per-channel FIFO queue to support HTTP pipelining over pooled connections: + +```java +class PendingRequest { + ClientToProxyConnection getClientConnection(); + HttpRequest getRequest(); + HttpFilters getFilters(); +} +``` + +--- + +## HTTP-Only Pooling Scenario + +When the proxy handles only plain HTTP (no MITM manager, no CONNECT), pooling is fully +determined by `useSharedServerConnectionPool`: + +```java +HttpProxyServer server = DefaultHttpProxyServer.bootstrap() + .withPort(8080) + .withSharedServerConnectionPool(true) + .start(); +``` + +### What gets pooled + +All non-CONNECT, non-tunneling, non-WebSocket HTTP requests go through the shared +pool. Each request acquires a connection from the pool, the response flows back, +and `markResponseComplete()` releases the connection immediately. + +### Connection lifecycle + +``` +Client A ── GET /api ──→ pool.getOrCreateConnection() ──→ [idle conn] or [new TCP] + ↓ + response received → markResponseComplete() → releaseConnection() + ↓ +Client B ── GET /api ──→ pool.getOrCreateConnection() ──→ [same idle conn from A] +``` + +When the pool has an idle connection for the target `host:port`, it is reused +directly — no new TCP socket. When all connections are busy, the pool either waits +(blocking borrow) or creates a new one up to `maxConnectionsPerHost`. + +### Use case + +A high-traffic forward proxy serving many clients hitting the same REST APIs. +Without pooling, each request opens a new TCP socket, does a TCP handshake (and +potentially TLS), then tears it down. With pooling, sockets stay alive and are +reused across clients, dramatically reducing latency and server load. + +### Backpressure in HTTP-only mode + +When using the pool, `currentServerConnection` is set on each request and cleared +on response complete. The `becameSaturated()` / `becameWritable()` / `serverBecameWriteable()` +methods in `ClientToProxyConnection` check this transient reference in addition to the +`serverConnectionsByHostAndPort` map, because pooled connections are not stored in that map. +This ensures backpressure signals flow correctly through the pooled connection to pause/resume +client reads. + +### Pool sizing for HTTP-only + +For HTTP-only workloads, `maxConnectionsPerHost` (default 10) limits concurrent requests +to any single origin. `maxConnections` (default 200) limits the total across all origins. +If your clients make many concurrent requests to the same server, increase +`maxConnectionsPerHost`. If you proxy to many different origins, increase `maxConnections`. + +--- + +## Mixed HTTP + HTTPS (MITM) Pooling Scenario + +When the proxy handles both plain HTTP and MITM'd HTTPS traffic, the pool serves both, +but the MITM paths require additional flags. + +### Configuration matrix + +```java +// HTTP only: pool is used for all non-CONNECT requests +HttpProxyServer.bootstrap() + .withPort(8080) + .withSharedServerConnectionPool(true) + .start(); + +// HTTP + MITM cross-client reuse: HTTPS upstream connections survive client sessions +HttpProxyServer.bootstrap() + .withPort(8080) + .withManInTheMiddle(myMitmManager) + .withSharedServerConnectionPool(true) + .withPoolSharedMitmConnections(true) + .start(); + +// HTTP + MITM per-request: full pooling, no dedicated upstream per session +HttpProxyServer.bootstrap() + .withPort(8080) + .withManInTheMiddle(myMitmManager) + .withSharedServerConnectionPool(true) + .withPoolSharedMitmConnections(true) + .withPoolPerRequestInMitm(true) + .start(); +``` + +### How the `useSharedPool` decision works + +In `ClientToProxyConnection.doReadHTTPInitial()`, every request goes through a +single decision tree: + +```java +boolean usePool = proxyServer.getServerConnectionPool() != null; +boolean isConnect = ProxyUtils.isCONNECT(httpRequest); +boolean poolSharedMitm = usePool && proxyServer.isPoolSharedMitmConnections(); +boolean poolPerRequest = usePool && proxyServer.isPoolPerRequestInMitm(); + +boolean useSharedPool = + usePool + && !isTunneling() + && !ProxyUtils.isSwitchingToWebSocketProtocol(httpRequest) + && (poolPerRequest || !isMitming()) + && (poolSharedMitm || !isConnect); +``` + +A request reaches the pool when: +- The pool is enabled (`usePool`) +- It is not a WebSocket upgrade or tunneling request — these are **always excluded** + because tunneling has no request/response boundary and WebSocket replaces HTTP codecs + with raw frame handlers, making pool return impossible +- **For CONNECT requests**: only if `poolSharedMitmConnections=true` +- **For MITM requests (after CONNECT)**: always if `poolPerRequestInMitm=true`; + otherwise the dedicated `serverConnectionsByHostAndPort` path is used +- **For plain HTTP requests**: always (no MITM/CONNECT gating applies) + +### What happens during a CONNECT in mixed mode + +``` +CONNECT example.com:443 + → useSharedPool? (only if poolSharedMitmConnections=true) + → Yes: pool.getOrCreateConnection() + → Pool returns idle connection (TCP + TLS already up) OR creates new one + → initializeConnectionFlow() with isReused check + → Reused: skip ConnectChannel + EncryptChannel → RespondCONNECTSuccessful → MitmEncryptClientChannel + → New: ConnectChannel → EncryptChannel → RespondCONNECTSuccessful → MitmEncryptClientChannel + → After CONNECT flow completes: + poolPerRequest? → releaseToPool() immediately + !poolPerRequest? → pinned to client session (mitmPooled=true), released on disconnect +``` + +While the CONNECT is being established, plain HTTP requests to different hosts +continue to use the pool independently — the CONNECT flow does not block them. + +### What happens during an HTTP GET in mixed mode (after CONNECT) + +``` +GET /api (through existing MITM tunnel) + → useSharedPool? + → poolPerRequest=true: pool.getOrCreateConnection() → borrows a (possibly different) connection + → poolPerRequest=false: use dedicated serverConnectionsByHostAndPort entry + → Response complete: + → poolPerRequest=true: markResponseComplete() → releaseConnection() + → poolPerRequest=false: connection stays ready for next request +``` + +### Pool sizing for mixed workloads + +In mixed mode, the pool serves both plain HTTP requests and MITM upstream connections +from the same pool. Each MITM upstream TLS connection counts toward the per-host and +global limits. If you expect many concurrent MITM sessions to the same host, ensure +`maxConnectionsPerHost` is high enough to accommodate both HTTP and HTTPS demand. + +Example: with `maxConnectionsPerHost=10`, if 8 HTTP requests and 5 MITM sessions all +target `api.example.com:443`, the 11th request will fail to acquire a connection. +Raise the limit or monitor `PoolMetrics` for borrow failures. + +### Separate pool keys for MITM vs. plain HTTP + +The pool key is computed by `ServerConnectionPool.computePoolKey()`: +- Plain HTTP to `example.com:443`: key = `"example.com:443:direct"` +- MITM CONNECT to `example.com:443`: key = `"example.com:443:direct"` (same key) + +This means a plain HTTP request and a MITM upstream connection to the same +`host:port` compete for the same pool entries. This is intentional — they are +connections to the same server and should be limited together. + +--- + +## Feature 1: `poolSharedMitmConnections` — Cross-Client Upstream Reuse + +### What it does + +When enabled, the upstream TLS connection created during a MITM `CONNECT` handshake is +stored in the shared connection pool (keyed by `host:port`). When a second (or third, +etc.) client connects to the **same** target host, the pool returns the cached +connection instead of opening a fresh TCP socket and TLS handshake. + +### Configuration + +```java +HttpProxyServer server = DefaultHttpProxyServer.bootstrap() + .withPort(8080) + .withManInTheMiddle(myMitmManager) + .withSharedServerConnectionPool(true) // master switch + .withPoolSharedMitmConnections(true) // Phase 1 + .start(); +``` + +Or via properties file: + +```properties +use_shared_server_connection_pool=true +pool_shared_mitm_connections=true +``` + +### How it works + +- The pool guard `!isMitming()` is replaced with a combined condition that checks the + new flag (`poolSharedMitmConnections`). +- During the CONNECT flow in `ClientToProxyConnection.doReadHTTPInitial()`, when + the flag is set, the connection is obtained via `pool.getOrCreateConnection()` + instead of `ProxyToServerConnection.create()` + `serverConnectionsByHostAndPort.put()`. +- The existing `initializeConnectionFlow()` is extended: when the connection was + retrieved from the pool and its channel is already active (`channel != null && + channel.isActive()`), the flow **skips** `ConnectChannel` and `EncryptChannel` + — the TCP socket and TLS handshake are already done from a previous session. +- When the client disconnects, the server connection is released back to the pool + (via `releaseToPool()`) instead of being closed, making it available for the next + client. + +### Key implementation details + +- A new `isReused` boolean is computed at the top of `initializeConnectionFlow()`. + When true, `ConnectChannel` and all chained/send-proxy/encrypt steps are bypassed. +- The `mitmPooled` flag on `ProxyToServerConnection` prevents `markResponseComplete()` + from releasing the connection after each individual HTTP response — the connection + stays pinned to the MITM session until the client disconnects. +- `ProxyToServerConnection` gains a `connectionPool` field (null for non-pooled + connections), `createForPool()` static factory, `releaseToPool()`, + `isManagedByPool()`, `isConnected()`, `isAvailableForNewRequest()`, and + `getClientConnection()` — the latter routes responses to the correct client + even when the `clientConnection` field points to the original client that + created the connection. + +### Use case + +A browser opens 10 tabs to `https://api.example.com`. Each tab creates a separate +client TCP connection through the proxy. Without pooling, 10 upstream TLS sockets +are opened to `api.example.com:443`. With pooling, the first tab's upstream +connection is reused for tabs 2–10. + +### Value + +- Reduces upstream TLS handshake overhead (CPU, latency) +- Reduces server-side connection load +- Lowers the number of concurrent outbound sockets +- Most impactful for proxy deployments with many clients hitting the same origins + +## Feature 2: `poolPerRequestInMitm` — Per-Request (vs. Per-Session) Borrowing + +*Requires `poolSharedMitmConnections=true`.* + +### What it does + +Without this flag, the pooled connection is pinned to the client's MITM session for +its lifetime — it is released back to the pool only when the client disconnects. +With this flag, the connection is released back to the pool **after each individual +HTTP request** completes, even while the client TCP session remains open. Subsequent +HTTP requests through the same MITM tunnel acquire a (possibly different) connection +from the pool. + +### Configuration + +```java +HttpProxyServer server = DefaultHttpProxyServer.bootstrap() + .withPort(8080) + .withManInTheMiddle(myMitmManager) + .withSharedServerConnectionPool(true) // master switch + .withPoolSharedMitmConnections(true) // Phase 1 + .withPoolPerRequestInMitm(true) // Phase 2 + .start(); +``` + +Or via properties file: + +```properties +use_shared_server_connection_pool=true +pool_shared_mitm_connections=true +pool_per_request_in_mitm=true +``` + +### How it works + +- The `useSharedPool` condition in `ClientToProxyConnection.doReadHTTPInitial()` + is extended: when `poolPerRequestInMitm` is true, MITM requests (after the + CONNECT) go through the pool path. +- The CONNECT response flow sets `releaseToPoolOnConnectComplete = true` on the + server connection. After the CONNECT flow completes (in + `connectionSucceeded()`), the connection is immediately released to the pool. +- `serverConnectionsByHostAndPort` is NOT used for MITM connections when this + flag is set — every HTTP request goes through `pool.getOrCreateConnection()`. +- `markResponseComplete()` calls `connectionPool.releaseConnection(this)` for + per-request connections (ones where `mitmPooled` is false). This returns the + connection to the available queue for another client's request. +- HTTP pipelining is handled via `PendingRequest` tracking in the pool. +- The `disconnected()` method in `ClientToProxyConnection` avoids a double-release: + per-request connections are already in the pool by the time the client + disconnects, so they just need `removeConnection()` on disconnect, not + `releaseToPool()`. + +### Key implementation details + +- The `releaseToPoolOnConnectComplete` transient flag on + `ProxyToServerConnection` is set during `initializeConnectionFlow()` and + consumed once in `connectionSucceeded()`. After release, subsequent HTTP + requests from the same MITM tunnel go through `doReadHTTPInitial()` → pool + path. +- `setCurrentClientConnectionForRequest()` is called on the borrowed connection + to ensure responses are routed to the correct client. +- A critical fix was made during development: `MitmEncryptClientChannel.execute()` + used the constructor field `clientConnection` (the original client that + initiated the CONNECT) instead of `getClientConnection()` (the client that + is making the current request). This caused + `testMultipleRequestsOverHTTPS` to fail with an SSL handshake error because + the encrypt was applied to the wrong channel. The fix was to use + `getClientConnection()` consistently. + +### Use case + +A client opens a single HTTPS connection and sends a GET, then sits idle for +30 seconds, then sends a POST. Without per-request pooling, the upstream +connection is held idle the whole time. With per-request pooling, it is +returned to the pool after the GET, available for other clients, and +re-acquired for the POST. + +### Value + +- Better connection utilization — idle time during a client session is + reclaimed for other clients +- Enables a smaller connection pool to serve the same workload +- Connections are not held captive by idle client sessions + +## Combined Scenario + +With both flags enabled, upstream connections are pooled across clients **and** +released between requests within a single client session. This is the most +aggressive connection-sharing mode, providing the highest potential connection +reuse. + +## Pool Metrics + +The `ServerConnectionPool` interface exposes `getMetrics()`, returning a +`PoolMetrics` object with: + +| Metric | Description | +|------------------------|------------------------------------------| +| `getTotalConnections()` | Total connections in the pool | +| `getActiveConnections()`| Connections currently borrowed | +| `getIdleConnections()` | Connections available for reuse | +| `getBorrowCount()` | Cumulative borrow operations | +| `getReturnCount()` | Cumulative return operations | +| `getEvictionCount()` | Cumulative eviction operations | +| `getValidationFailureCount()` | Connections that failed validation | + +These metrics are accessible from tests or monitoring code by casting +`HttpProxyServer` to `DefaultHttpProxyServer` and calling +`getServerConnectionPool().getMetrics()`. + +## Feature Interaction Matrix + +| `useSharedPool` | `poolSharedMitm` | `poolPerRequest` | Behavior | +|---|---|---|---| +| `false` | — | — | Legacy: dedicated connection per client per host. All connections are tracked in `serverConnectionsByHostAndPort` and closed on disconnect. Main-branch behavior. | +| `true` | `false` | `false` | Plain HTTP is pooled. CONNECT and MITM use dedicated connections (main-branch pool behavior). | +| `true` | `true` | `false` | Plain HTTP **and** MITM upstream connections are pooled across client sessions. Each MITM session holds one pooled connection until the client disconnects. | +| `true` | `true` | `true` | Full pooling: plain HTTP and MITM connections are borrowed per request and released back to the pool between requests within the same client session. | + +(`poolPerRequestInMitm=true` without `poolSharedMitmConnections=true` has no +effect — per-request mode requires the shared pool for MITM.) + +## Configuration Reference + +### Builder methods on `HttpProxyServerBootstrap` + +```java +// Base pooling (PR #724 infrastructure) +.withSharedServerConnectionPool(boolean) +.withServerConnectionPoolType(ServerConnectionPoolType) +.withMaxConnectionsPerHost(int) +.withMaxConnections(int) +.withPoolIdleTimeout(Duration) + +// MITM-specific (this branch) +.withPoolSharedMitmConnections(boolean) // default false +.withPoolPerRequestInMitm(boolean) // default false +``` + +### Properties file keys (for `--config`) + +```properties +# Base pooling +use_shared_server_connection_pool=true +server_connection_pool_type=CONCURRENT_MAP +max_connections_per_host=10 +max_total_connections=200 + +# MITM-specific +pool_shared_mitm_connections=true +pool_per_request_in_mitm=true +``` + +## Test Coverage + +### Unit tests + +| Test class | Tests added | What it covers | +|---|---|---| +| `ServerConnectionPoolConfigTest` | 8 | Default values, fluent setters/getters, independence from `enabled` flag | +| `DefaultHttpProxyServerBootstrapTest` | 5 | Property parsing for both flags via `Properties` constructor | + +### Integration tests + +| Test class | Tag | Tests | What it covers | +|---|---|---|---| +| `MitmWithSharedPoolTest` | — | 4 | GET, POST, cross-client reuse, pool metrics (borrow count) | +| `MitmWithPerRequestPoolTest` | `slow-test` | 4 | GET, POST, sequential reuse, cross-client reuse | + +### Existing tests exercising the underlying pool infrastructure + +| Test class | Tests | What it covers | +|---|---|---| +| `ConcurrentMapServerConnectionPoolTest` | 24 | Pool implementation: borrow, release, eviction, pending requests | +| `SharedConnectionPoolTest` | 13 | Integrated shared pool for plain HTTP | +| `ServerConnectionPoolTypeTest` | 6 | Pool type selection | +| `ClientToProxyConnectionShortCircuitTest` | 5 | Short-circuit filter response with pooled connections | +| `ClientToProxyConnectionBackpressureTest` | 15 | Backpressure / saturation with pooled connections | + +## Excluded Protocols + +The following traffic types are never pooled and always use dedicated +`ProxyToServerConnection` instances tracked in `serverConnectionsByHostAndPort`: + +| Protocol | Reason | +|---|---| +| **WebSocket** (`Upgrade: websocket`) | After the HTTP upgrade handshake, the HTTP codecs are replaced with `WebSocketFramePipeHandler`. The connection becomes a raw frame pipe between client and server with no HTTP request/response lifecycle to trigger pool release. | +| **Tunneling** (non-MITM CONNECT, i.e. regular HTTPS) | Once the CONNECT response is sent, the connection enters raw TCP tunneling mode. There are no HTTP request/response boundaries — all data flows bidirectionally until the tunnel closes, so the connection cannot be returned to the pool. This is the default HTTPS proxy behavior (no intercept), and it is **never** pooled. Only MITM-intercepting HTTPS can be pooled (behind `poolSharedMitmConnections`). | +| **Switching Protocols** (other `Upgrade` headers) | Same as WebSocket — HTTP codecs are removed and the connection switches to a different protocol, making pool return impossible. | + +The `useSharedPool` condition explicitly checks `!isTunneling()` and +`!ProxyUtils.isSwitchingToWebSocketProtocol(httpRequest)` before every request. +No configuration flag can override these exclusions. + +## HTTP/2 Compatibility Note + +The underlying pooling architecture (pool key generation, connection tracking, +and connection lifecycle) is independent of HTTP version and could support HTTP/2 +multiplexed streams. Future HTTP/2 server connections would reuse the same +`PoolMetrics` infrastructure to track success/failure of multiplexed requests, +though ALPN and stream-tracking adjustments would be required beyond the current +implementation. + +## Known Limitations + +- **Metrics API not on `HttpProxyServer` interface:** To access pool metrics + from client code, cast to `DefaultHttpProxyServer`. This is a minor API gap. +- **No `encryptForMitm()` integration test:** The code path reached via + `disableSslForNonTls` (retry after failed TLS to a plain-text server) has + no integration test coverage. A bug in `encryptForMitm()` analogous to the + `MitmEncryptClientChannel` bug was fixed during development. diff --git a/littleproxy.properties b/littleproxy.properties deleted file mode 100644 index 452ac128..00000000 --- a/littleproxy.properties +++ /dev/null @@ -1,4 +0,0 @@ -# Exposes proxy connection properties via JMX. -jmx=false -# Idle connections are disconnected after X seconds of inactivity -idle_connection_timeout=70 \ No newline at end of file diff --git a/littleproxy_cert b/littleproxy_cert deleted file mode 100644 index d31898a2..00000000 Binary files a/littleproxy_cert and /dev/null differ diff --git a/pom.xml b/pom.xml index f46b2e48..e3e64b68 100644 --- a/pom.xml +++ b/pom.xml @@ -1,57 +1,97 @@ 4.0.0 - xyz.rogfam + io.github.littleproxy littleproxy jar - 2.0.0-beta-6-SNAPSHOT + 2.10.0-SNAPSHOT LittleProxy LittleProxy is a high performance HTTP proxy written in Java and using the Netty networking framework. - https://github.com/mrog/LittleProxy + https://github.com/LittleProxy/LittleProxy UTF-8 UTF-8 github - 4.1.41.Final - 1.7.28 - 1.8 + 11 + 17 + + + 3.27.7 + 1.11.0 + 1.6.0 + 2.22.0 + 3.20.0 + 4.0.0 + 3.6.5 + 0.1.6 + 2.42.0 + 33.7.1-jre + 4.5.14 + 2.22.2 + 11.0.24 + 6.1.3 + 2.26.1 + 5.23.0 + 4.2.18.Final + 4.49.0 + 2.0.19 + 3.13.2 The Apache Software License, Version 2.0 - http://www.apache.org/licenses/LICENSE-2.0 + https://www.apache.org/licenses/LICENSE-2.0 github - https://github.com/mrog/LittleProxy/issues + https://github.com/LittleProxy/LittleProxy/issues - scm:git:https://github.com/mrog/LittleProxy.git - scm:git:git@github.com:mrog/LittleProxy.git - scm:git:https://github.com/mrog/LittleProxy + scm:git:https://github.com/LittleProxy/LittleProxy.git + scm:git:git@github.com:LittleProxy/LittleProxy.git + scm:git:https://github.com/LittleProxy/LittleProxy HEAD - - - ossrh - https://oss.sonatype.org/content/repositories/snapshots - - - ossrh - https://oss.sonatype.org/service/local/staging/deploy/maven2/ - - - 2009 + + smoke-test + + + + org.apache.maven.plugins + maven-surefire-plugin + + slow-test + -ea -Xmx256m -XX:+HeapDumpOnOutOfMemoryError -XX:HeapDumpPath=target/smoke-tests.hprof + + + + + + + slow-tests + + + + org.apache.maven.plugins + maven-surefire-plugin + + slow-test + -ea -Xmx256m -XX:+HeapDumpOnOutOfMemoryError -XX:HeapDumpPath=target/slow-tests.hprof + + + + + release @@ -108,52 +148,47 @@ ${gpg.keyname} - ${gpg.keyname} gpg - - --pinentry-mode - loopback - - - org.sonatype.plugins - nexus-staging-maven-plugin - true - - ossrh - https://oss.sonatype.org/ - false - - - - org.apache.maven.plugins - maven-release-plugin - - true - false - release - deploy - - + com.google.guava guava - 27.1-jre + ${guava.version} + + + com.google.code.findbugs + jsr305 + + + + + + org.jspecify + jspecify + 1.0.1 + provided + + + com.google.errorprone + error_prone_annotations + ${error_prone.version} + provided commons-cli commons-cli - 1.4 + ${commons.cli.version} true @@ -161,183 +196,209 @@ org.apache.commons commons-lang3 - 3.8.1 + ${commons.lang3.version} + - junit - junit - 4.12 - test + io.netty + netty-buffer - - org.hamcrest - hamcrest-core - 2.1 - test + io.netty + netty-codec - - org.hamcrest - hamcrest-library - 2.1 - test + io.netty + netty-codec-http - - org.eclipse.jetty - jetty-server - 9.4.20.v20190813 - test + io.netty + netty-codec-haproxy + + + io.netty + netty-codec-socks + + + io.netty + netty-common + + + io.netty + netty-handler + + + io.netty + netty-handler-proxy + + + io.netty + netty-resolver + + + io.netty + netty-transport + - org.mockito - mockito-core - 2.25.1 - test + org.slf4j + slf4j-api + ${slf4j.version} + + + org.apache.logging.log4j + log4j-slf4j2-impl + ${log4j.version} + true + + + org.apache.logging.log4j + log4j-core + ${log4j.version} + true + + + + com.lmax + disruptor + ${disruptor.version} + true + - org.mock-server - mockserver-netty - 5.6.1 - test + org.littleshoot + dnssec4j + ${dnssec4j.version} + true - ch.qos.logback - logback-classic + org.littleshoot + dnsjava + + + log4j + log4j + + + org.slf4j + slf4j-log4j12 + + + org.apache.commons + commons-lang3 + + + org.slf4j + slf4j-api - org.seleniumhq.selenium - selenium-java - 3.141.59 - test + dnsjava + dnsjava + ${dnsjava.version} + true - io.netty - netty + org.slf4j + slf4j-api + - org.apache.logging.log4j - log4j-core - 2.11.2 - true + org.junit.jupiter + junit-jupiter + ${junit.version} + test + + + org.assertj + assertj-core + ${assertj.version} + test - org.apache.httpcomponents - httpclient - 4.5.8 + org.eclipse.jetty + jetty-server + ${jetty.version} test - io.netty - netty-all + org.mockito + mockito-core + ${mockito.version} + test + - io.netty - netty-example + org.wiremock + wiremock-standalone + ${wiremock.version} test - com.barchart.udt - barchart-udt-bundle - 2.3.0 + commons-io + commons-io + ${commons.io.version} + test - org.littleshoot - dnssec4j - 0.1.6 - true + org.seleniumhq.selenium + selenium-java + ${selenium.version} + test - org.littleshoot - dnsjava + io.netty + netty - dnsjava - dnsjava - 2.1.8 - true - - - - org.slf4j - slf4j-log4j12 - ${slf4j.version} - true + org.apache.httpcomponents + httpclient + ${httpclient.version} + test - org.slf4j - slf4j-api - ${slf4j.version} + io.netty + netty-example + ${netty.version} + test - org.apache.commons commons-exec - 1.3 + ${commons.exec.version} test - - - - io.netty - netty-all - ${netty.version} - - - io.netty - netty-buffer - ${netty.version} - - - io.netty - netty-codec - ${netty.version} - - - io.netty - netty-codec-haproxy - ${netty.version} - - - io.netty - netty-codec-http - ${netty.version} - - - io.netty - netty-codec-socks - ${netty.version} - + io.netty - netty-common + netty-bom ${netty.version} + import + pom + + io.netty netty-example @@ -362,43 +423,14 @@ - - io.netty - netty-handler - ${netty.version} - - - io.netty - netty-handler-proxy - ${netty.version} - - - io.netty - netty-transport - ${netty.version} - - - io.netty - netty-transport-rxtx - ${netty.version} - - - io.netty - netty-transport-sctp - ${netty.version} - - - io.netty - netty-transport-udp - ${netty.version} - + - - - com.fasterxml.jackson.core - jackson-databind - 2.9.9.3 + com.fasterxml.jackson + jackson-bom + ${jackson.bom.version} + import + pom @@ -409,80 +441,110 @@ org.apache.maven.plugins maven-enforcer-plugin - 3.0.0-M2 + 3.6.3 org.apache.maven.plugins maven-site-plugin - 3.8.2 + 3.22.0 org.apache.maven.plugins maven-release-plugin - 2.5.3 + 3.3.1 org.apache.maven.plugins maven-dependency-plugin - 3.1.1 + 3.11.0 org.apache.maven.plugins maven-clean-plugin - 3.1.0 + 3.5.0 org.apache.maven.plugins maven-deploy-plugin - 3.0.0-M1 + 3.2.0 org.apache.maven.plugins maven-compiler-plugin - 3.8.0 + 3.16.0 - ${java.version} - ${java.version} + ${java.version} UTF-8 + + + -XDcompilePolicy=simple + -Xplugin:ErrorProne -Xep:MissingSummary:OFF -Xep:JdkObsolete:OFF -Xep:ReferenceEquality:OFF -Xep:OperatorPrecedence:OFF + + --add-exports=jdk.compiler/com.sun.tools.javac.api=ALL-UNNAMED + --add-exports=jdk.compiler/com.sun.tools.javac.code=ALL-UNNAMED + --add-exports=jdk.compiler/com.sun.tools.javac.util=ALL-UNNAMED + --add-opens=jdk.compiler/com.sun.tools.javac.comp=ALL-UNNAMED + --add-opens=jdk.compiler/com.sun.tools.javac.tree=ALL-UNNAMED + --add-opens=jdk.compiler/com.sun.tools.javac.main=ALL-UNNAMED + + + + com.google.errorprone + error_prone_core + ${error_prone.version} + + + + + default-testCompile + test-compile + + testCompile + + + ${java.version.tests} + + + org.apache.maven.plugins maven-install-plugin - 3.0.0-M1 + 3.2.0 org.apache.maven.plugins maven-jar-plugin - 3.1.1 + 3.5.1 org.apache.maven.plugins maven-resources-plugin - 3.1.0 + 3.5.0 org.apache.maven.plugins maven-source-plugin - 3.0.1 + 3.4.0 org.apache.maven.plugins maven-javadoc-plugin - 3.1.1 + 3.12.0 - all,-missing + all,-missing,-reference private ${java.version} @@ -494,21 +556,14 @@ org.apache.maven.plugins maven-surefire-plugin - 2.22.1 + 3.6.0 org.apache.maven.plugins maven-gpg-plugin - 1.6 - - - - org.sonatype.plugins - nexus-staging-maven-plugin - 1.6.8 + 3.2.8 - @@ -517,7 +572,7 @@ org.apache.maven.plugins maven-surefire-plugin - -Xmx1g -XX:MaxPermSize=256m + -Xmx1g @@ -534,7 +589,7 @@ org.apache.maven.plugins maven-shade-plugin - 3.2.1 + 3.6.2 package @@ -542,6 +597,7 @@ shade + false true littleproxy-shade @@ -549,7 +605,11 @@ org.bouncycastle:* - + + org.apache.logging.log4j:log4j-core:${log4j.version} + org.apache.logging.log4j:log4j-slf4j2-impl:${log4j.version} + com.lmax:disruptor:${disruptor.version} + *:* @@ -565,11 +625,16 @@ org.littleshoot.proxy.Launcher + true - - log4j.xml - src/main/config/log4j.xml + + + + META-INF/log4j2-plugin.dat + + + META-INF/maven/org.apache.logging.log4j/log4j-core/pom.properties @@ -591,12 +656,55 @@ 3.0.4 + + + junit:junit + org.hamcrest:hamcrest-core + + - + + org.sonatype.central + central-publishing-maven-plugin + 0.11.0 + true + + central + true + published + + + + com.diffplug.spotless + spotless-maven-plugin + 3.10.2 + + + + src/main/java/**/*.java + src/test/java/**/*.java + + + 1.28.0 + + + + + + + + + + apply + + compile + + + @@ -621,7 +729,7 @@ org.apache.maven.plugins maven-project-info-reports-plugin - 3.0.0 + 3.9.0 @@ -632,7 +740,7 @@ org.apache.maven.plugins maven-surefire-report-plugin - 3.0.0-M3 + 3.6.0 false @@ -641,7 +749,7 @@ org.apache.maven.plugins maven-checkstyle-plugin - 3.0.0 + 3.6.0 @@ -675,13 +783,13 @@ org.apache.maven.plugins maven-jxr-plugin - 3.0.0 + 3.6.0 org.apache.maven.plugins maven-pmd-plugin - 3.11.0 + 3.28.0 true utf-8 @@ -693,7 +801,7 @@ org.codehaus.mojo versions-maven-plugin - 2.7 + 2.22.0 @@ -709,7 +817,7 @@ org.codehaus.mojo taglist-maven-plugin - 2.4 + 3.2.3 true @@ -724,7 +832,6 @@ - mrogers @@ -735,5 +842,14 @@ Developer -7 + + asolntsev + Andrei Solntsev + andrei.solntsev+littleproxy@gmail.com + LittleProxy + https://github.com/LittleProxy + Maintainer + +3 + diff --git a/public-gpg.key b/public-gpg.key new file mode 100644 index 00000000..ce63b313 --- /dev/null +++ b/public-gpg.key @@ -0,0 +1,31 @@ +-----BEGIN PGP PUBLIC KEY BLOCK----- + +mQENBFx5TsYBCAC5MFtvCeSvvt60CAY3P2TAoPxIdcOowB4yg/jMzw7bGOwBX5xm +e7cfvIeuSvgR3UuMLbvcJU3R4q6WRCm8OsISFJjGduKO89FNS3EvetJwIXjPgla+ +v9LjmEVl4wzaBRSH6vzYQl3EbIHMVaW138HQLCbaULf6CxPp1+vBeAVHEmXa9qit +Ly3sPgOlCm9g/x8o25k4wdj4xwpSFlUe/yl+yPqK9xjVNt0V9oR4NBm9bhsxOeo6 +LH/nQcaCKNRK5wK3ytVN1FgNEYv9cARuN6X7Qvi3Oi0P1+lVVpKJu0Ux0Gx0UB3O +imnN3V61Pcs/yzHnyyW2OD0PUbX3JvcC4xdhABEBAAG0I01hcmsgUm9nZXJzIDxt +YXJrLnJvZ2Vyc0BnbWFpbC5jb20+iQFUBBMBCAA+AhsDBQsJCAcCBhUKCQgLAgQW +AgMBAh4BAheAFiEEzqtrzoRv2mNdpZZkFNFpBtlIKpIFAmDKNJ4FCQn0gFgACgkQ +FNFpBtlIKpI2LwgAiYRIzB5ISp/Ie90TztYvFjByShxjYSIBE25pWhGOez9wgaEM +g2uDIhQflVkY7U062Aqecnk686NsMA/McfgP0d2mr9fUxpxQoDHqOqtlmcWkt1k3 +sVZEPLM61OCyXUbU6cITcyDk9v1CLWfba/LKi0TteH+HQRa+/xrctHtuwpvPB2E4 +HCeVZT54DeyW1/YaHyJ1npo4AnGo+x1JGFOIVfJ6X74YHxvi0axcfdjJZVIP4qIZ +yWWu5YcvhDFjddpDfd/Hh9be+Ym3msVA764anPKHJS00Eg6PQSoesqluh0c0ZUN3 +VxiNVk0+xBx68IgcSYP6kSP5JxuW8CP3detE2rkBDQRceU7GAQgAwaQT351vLvTU +5kq0/SJxN1S41o4MYShLlI60InjHOTybjE9CdVeEkCivhcwtJVxXmNrAIw8T1vK2 +UxbmEPaii5fe7VSBShndIfwC4KCtCxFg2T7VcV8dftpbHW3E22a0zLbhpkgLcDzO +EOFSRNl01k4j/Pjda+wHztVHcJdS5I1HzeEa20Sv9Agnsf4UGvRheLtj9h0aWb/o +XZEhaEYAUDwHSBMd8B9OoG/ZL7sdNbLRwNKZYgDZ+K2k/Hg59+s5UoM+EVo0/ppf +K1e5NFhiI7qsEUk2mh74+R8cjtdT5mARb4nZZKf1/cS7z1oS3st7Qv1wVkzLMnhA +Q+fpmo/TuQARAQABiQE8BBgBCAAmAhsMFiEEzqtrzoRv2mNdpZZkFNFpBtlIKpIF +AmDKNMIFCQn0gHwACgkQFNFpBtlIKpLoTAf/ZlPg4c4BfV5cZ6u3KanJsx8OpENn +raPEEnyOnJhZdQHmxKUokMgtMwZLheA50jOh80pTPdQjkKStIrcWygE26iBLODKc ++dzjdHZUazU/P6671VNnIZbSVk2mNuJgPUafrGGgyD+hWRxv7rJ12cV4xKvlntf9 +M4gu1aKiXKqaOYMnVXn3eU4lrcfJqVW8tqiHLX3xuWn7IS3JLzQ5PmN7zHKNZO8o +/beP/hjsW6ceO8AMkY0vMCeQMNRMovxdQ1VU5KcPJjoAHkpJQoA0rvSLrRYq6Vvk +Bm6EikIxjJMPg9o6hoeg1+lwd0rdJSh9mPC4qlaQLUq0hPOZ5gTVnPtWzQ== +=lK6I +-----END PGP PUBLIC KEY BLOCK----- + diff --git a/run.bash b/run.bash index a69d457f..9cd2e049 100755 --- a/run.bash +++ b/run.bash @@ -4,12 +4,74 @@ function die() { exit 1 } +# Show help if requested +if [[ "$1" == "--help" || "$1" == "-h" || "$1" == "help" ]]; then + echo "LittleProxy Run Script" + echo "Usage: $0 [options]" + echo "" + echo "Options:" + echo " --server Run as server" + echo " --config Configuration file path" + echo " --port Port to listen on" + echo " --log_config Log4j2 configuration file path" + echo " --activity_log_format Activity log format (CLF, JSON, etc.)" + echo " --async_logging_default Use asynchronous logging with default config" + echo " --help, -h, help Show this help message" + echo "" + echo "Example:" + echo " $0 --server --config ./config/littleproxy.properties --port 9092 --async_logging_default" + exit 0 +fi + mvn package -Dmaven.test.skip=true || die "Could not package" fullPath=`dirname $0` jar=`find $fullPath/target/littleproxy*-littleproxy-shade.jar` cp=`echo $jar | sed 's,./,'$fullPath'/,'` -javaArgs="-server -XX:+HeapDumpOnOutOfMemoryError -Xmx800m -jar "$cp" $*" + +# Initialize Java arguments +javaArgs="-server -XX:+HeapDumpOnOutOfMemoryError -Xmx800m" + +# Parse arguments to handle --async_logging_default flag and build proper argument list +async_logging_default=false +remaining_args=() +log_config_set=false +custom_log_config="" + +while [[ $# -gt 0 ]]; do + case "$1" in + --async_logging_default) + async_logging_default=true + shift + ;; + --log_config) + log_config_set=true + custom_log_config="$2" + remaining_args+=("--log_config" "$2") + shift 2 + ;; + --log_config=*) + log_config_set=true + custom_log_config="${1#*=}" + remaining_args+=("$1") + shift + ;; + *) + remaining_args+=("$1") + shift + ;; + esac +done + +# Add async logging if flag is set AND no custom log config is provided +if [ "$async_logging_default" = true ] && [ "$log_config_set" = false ]; then + echo "Async logging enabled (using default async configuration)" + remaining_args+=("--log_config" "./target/classes/littleproxy_async_log4j2.xml") +elif [ "$async_logging_default" = true ] && [ "$log_config_set" = true ]; then + echo "Warning: --async_logging_default flag ignored because custom --log_config is specified: $custom_log_config" +fi + +javaArgs="$javaArgs -jar "$cp" ${remaining_args[@]}" echo "Running using Java on path at `which java` with args $javaArgs" java $javaArgs || die "Java process exited abnormally" diff --git a/src/main/config/log4j.xml b/src/main/config/log4j.xml deleted file mode 100644 index f284b3fa..00000000 --- a/src/main/config/log4j.xml +++ /dev/null @@ -1,33 +0,0 @@ - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - \ No newline at end of file diff --git a/src/main/java/org/littleshoot/proxy/ActivityTracker.java b/src/main/java/org/littleshoot/proxy/ActivityTracker.java index 28b9cc5e..9b8bb352 100644 --- a/src/main/java/org/littleshoot/proxy/ActivityTracker.java +++ b/src/main/java/org/littleshoot/proxy/ActivityTracker.java @@ -2,146 +2,142 @@ import io.netty.handler.codec.http.HttpRequest; import io.netty.handler.codec.http.HttpResponse; - -import javax.net.ssl.SSLSession; import java.net.InetSocketAddress; +import javax.net.ssl.SSLSession; /** - *

* Interface for receiving information about activity in the proxy. - *

- * - *

- * Sub-classes may wish to extend {@link ActivityTrackerAdapter} for sensible - * defaults. - *

+ * + *

Sub-classes may wish to extend {@link ActivityTrackerAdapter} for sensible defaults. */ public interface ActivityTracker { - /** - * Record that a client connected. - */ - void clientConnected(InetSocketAddress clientAddress); - - /** - * Record that a client's SSL handshake completed. - */ - void clientSSLHandshakeSucceeded(InetSocketAddress clientAddress, - SSLSession sslSession); - - /** - * Record that a client disconnected. - */ - void clientDisconnected(InetSocketAddress clientAddress, - SSLSession sslSession); - - /** - * Record that the proxy received bytes from the client. - * - * @param flowContext - * if full information is available, this will be a - * {@link FullFlowContext}. - * @param numberOfBytes - */ - void bytesReceivedFromClient(FlowContext flowContext, - int numberOfBytes); - - /** - *

- * Record that proxy received an {@link HttpRequest} from the client. - *

- * - *

- * Note - on chunked transfers, this is only called once (for the initial - * HttpRequest object). - *

- * - * @param flowContext - * if full information is available, this will be a - * {@link FullFlowContext}. - * @param httpRequest - */ - void requestReceivedFromClient(FlowContext flowContext, - HttpRequest httpRequest); - - /** - * Record that the proxy attempted to send bytes to the server. - * - * @param flowContext - * provides contextual information about the flow - * @param numberOfBytes - */ - void bytesSentToServer(FullFlowContext flowContext, int numberOfBytes); - - /** - *

- * Record that proxy attempted to send a request to the server. - *

- * - *

- * Note - on chunked transfers, this is only called once (for the initial - * HttpRequest object). - *

- * - * @param flowContext - * provides contextual information about the flow - * @param httpRequest - */ - void requestSentToServer(FullFlowContext flowContext, - HttpRequest httpRequest); - - /** - * Record that the proxy received bytes from the server. - * - * @param flowContext - * provides contextual information about the flow - * @param numberOfBytes - */ - void bytesReceivedFromServer(FullFlowContext flowContext, int numberOfBytes); - - /** - *

- * Record that the proxy received an {@link HttpResponse} from the server. - *

- * - *

- * Note - on chunked transfers, this is only called once (for the initial - * HttpRequest object). - *

- * - * @param flowContext - * provides contextual information about the flow - * @param httpResponse - */ - void responseReceivedFromServer(FullFlowContext flowContext, - HttpResponse httpResponse); - - /** - * Record that the proxy sent bytes to the client. - * - * @param flowContext - * if full information is available, this will be a - * {@link FullFlowContext}. - * @param numberOfBytes - */ - void bytesSentToClient(FlowContext flowContext, int numberOfBytes); - - /** - *

- * Record that the proxy sent a response to the client. - *

- * - *

- * Note - on chunked transfers, this is only called once (for the initial - * HttpRequest object). - *

- * - * @param flowContext - * if full information is available, this will be a - * {@link FullFlowContext}. - * @param httpResponse - */ - void responseSentToClient(FlowContext flowContext, - HttpResponse httpResponse); + /** Record that a client connected. */ + void clientConnected(FlowContext flowContext); + + /** Record that a client's SSL handshake started. */ + void clientSSLHandshakeStarted(FlowContext flowContext); + + /** Record that a client's SSL handshake completed. */ + void clientSSLHandshakeSucceeded(FlowContext flowContext, SSLSession sslSession); + + /** Record that a client disconnected. */ + void clientDisconnected(FlowContext flowContext, SSLSession sslSession); + + /** + * Record that the proxy received bytes from the client. + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + * @param numberOfBytes + */ + void bytesReceivedFromClient(FlowContext flowContext, int numberOfBytes); + + /** + * Record that proxy received an {@link HttpRequest} from the client. + * + *

Note - on chunked transfers, this is only called once (for the initial HttpRequest object). + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + * @param httpRequest + */ + void requestReceivedFromClient(FlowContext flowContext, HttpRequest httpRequest); + + /** + * Record that the proxy attempted to send bytes to the server. + * + * @param flowContext provides contextual information about the flow + * @param numberOfBytes + */ + void bytesSentToServer(FullFlowContext flowContext, int numberOfBytes); + + /** + * Record that proxy attempted to send a request to the server. + * + *

Note - on chunked transfers, this is only called once (for the initial HttpRequest object). + * + * @param flowContext provides contextual information about the flow + * @param httpRequest + */ + void requestSentToServer(FullFlowContext flowContext, HttpRequest httpRequest); + + /** + * Record that the proxy received bytes from the server. + * + * @param flowContext provides contextual information about the flow + * @param numberOfBytes + */ + void bytesReceivedFromServer(FullFlowContext flowContext, int numberOfBytes); + + /** + * Record that the proxy received an {@link HttpResponse} from the server. + * + *

Note - on chunked transfers, this is only called once (for the initial HttpRequest object). + * + * @param flowContext provides contextual information about the flow + * @param httpResponse + */ + void responseReceivedFromServer(FullFlowContext flowContext, HttpResponse httpResponse); + + /** + * Record that the proxy sent bytes to the client. + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + * @param numberOfBytes + */ + void bytesSentToClient(FlowContext flowContext, int numberOfBytes); + + /** + * Record that the proxy sent a response to the client. + * + *

Note - on chunked transfers, this is only called once (for the initial HttpRequest object). + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + * @param httpResponse + */ + void responseSentToClient(FlowContext flowContext, HttpResponse httpResponse); + + /** + * Record that the proxy connected to the server. + * + * @param flowContext provides contextual information about the flow + * @param serverAddress the address of the server that was connected + */ + void serverConnected(FullFlowContext flowContext, InetSocketAddress serverAddress); + + /** + * Record that the proxy disconnected from the server. + * + * @param flowContext provides contextual information about the flow + * @param serverAddress the address of the server that was disconnected + */ + void serverDisconnected(FullFlowContext flowContext, InetSocketAddress serverAddress); + + /** + * Record that a connection became saturated (not writable). + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + */ + void connectionSaturated(FlowContext flowContext); + + /** + * Record that a connection became writable again after being saturated. + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + */ + void connectionWritable(FlowContext flowContext); + + /** + * Record that a connection timed out due to idle timeout. + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + */ + void connectionTimedOut(FlowContext flowContext); + /** + * Record that an exception was caught on a connection. + * + * @param flowContext if full information is available, this will be a {@link FullFlowContext}. + * @param cause the exception that was caught + */ + void connectionExceptionCaught(FlowContext flowContext, Throwable cause); } diff --git a/src/main/java/org/littleshoot/proxy/ActivityTrackerAdapter.java b/src/main/java/org/littleshoot/proxy/ActivityTrackerAdapter.java index 9a04838a..f0f052ec 100644 --- a/src/main/java/org/littleshoot/proxy/ActivityTrackerAdapter.java +++ b/src/main/java/org/littleshoot/proxy/ActivityTrackerAdapter.java @@ -2,67 +2,66 @@ import io.netty.handler.codec.http.HttpRequest; import io.netty.handler.codec.http.HttpResponse; - -import javax.net.ssl.SSLSession; import java.net.InetSocketAddress; +import javax.net.ssl.SSLSession; /** - * Adapter of {@link ActivityTracker} interface that provides default no-op - * implementations of all methods. + * Adapter of {@link ActivityTracker} interface that provides default no-op implementations of all + * methods. */ public class ActivityTrackerAdapter implements ActivityTracker { - @Override - public void bytesReceivedFromClient(FlowContext flowContext, - int numberOfBytes) { - } - - @Override - public void requestReceivedFromClient(FlowContext flowContext, - HttpRequest httpRequest) { - } - - @Override - public void bytesSentToServer(FullFlowContext flowContext, int numberOfBytes) { - } - - @Override - public void requestSentToServer(FullFlowContext flowContext, - HttpRequest httpRequest) { - } - - @Override - public void bytesReceivedFromServer(FullFlowContext flowContext, - int numberOfBytes) { - } - - @Override - public void responseReceivedFromServer(FullFlowContext flowContext, - HttpResponse httpResponse) { - } - - @Override - public void bytesSentToClient(FlowContext flowContext, - int numberOfBytes) { - } - - @Override - public void responseSentToClient(FlowContext flowContext, - HttpResponse httpResponse) { - } - - @Override - public void clientConnected(InetSocketAddress clientAddress) { - } - - @Override - public void clientSSLHandshakeSucceeded(InetSocketAddress clientAddress, - SSLSession sslSession) { - } - - @Override - public void clientDisconnected(InetSocketAddress clientAddress, - SSLSession sslSession) { - } + @Override + public void bytesReceivedFromClient(FlowContext flowContext, int numberOfBytes) {} + + @Override + public void requestReceivedFromClient(FlowContext flowContext, HttpRequest httpRequest) {} + + @Override + public void bytesSentToServer(FullFlowContext flowContext, int numberOfBytes) {} + + @Override + public void requestSentToServer(FullFlowContext flowContext, HttpRequest httpRequest) {} + + @Override + public void bytesReceivedFromServer(FullFlowContext flowContext, int numberOfBytes) {} + + @Override + public void responseReceivedFromServer(FullFlowContext flowContext, HttpResponse httpResponse) {} + + @Override + public void bytesSentToClient(FlowContext flowContext, int numberOfBytes) {} + + @Override + public void responseSentToClient(FlowContext flowContext, HttpResponse httpResponse) {} + + @Override + public void clientConnected(FlowContext flowContext) {} + + @Override + public void clientSSLHandshakeStarted(FlowContext flowContext) {} + + @Override + public void clientSSLHandshakeSucceeded(FlowContext flowContext, SSLSession sslSession) {} + + @Override + public void clientDisconnected(FlowContext flowContext, SSLSession sslSession) {} + + @Override + public void serverConnected(FullFlowContext flowContext, InetSocketAddress serverAddress) {} + + @Override + public void serverDisconnected(FullFlowContext flowContext, InetSocketAddress serverAddress) {} + + @Override + public void connectionSaturated(FlowContext flowContext) {} + + @Override + public void connectionWritable(FlowContext flowContext) {} + + @Override + public void connectionTimedOut(FlowContext flowContext) {} + @Override + public void connectionExceptionCaught(FlowContext flowContext, Throwable cause) {} } diff --git a/src/main/java/org/littleshoot/proxy/ChainedProxy.java b/src/main/java/org/littleshoot/proxy/ChainedProxy.java index d39d0f09..2fc960e7 100644 --- a/src/main/java/org/littleshoot/proxy/ChainedProxy.java +++ b/src/main/java/org/littleshoot/proxy/ChainedProxy.java @@ -1,92 +1,77 @@ package org.littleshoot.proxy; import io.netty.handler.codec.http.HttpObject; - import java.net.InetSocketAddress; /** - *

* Encapsulates information needed to connect to a chained proxy. - *

- * - *

- * Sub-classes may wish to extend {@link ChainedProxyAdapter} for sensible - * defaults. - *

+ * + *

Sub-classes may wish to extend {@link ChainedProxyAdapter} for sensible defaults. */ public interface ChainedProxy extends SslEngineSource { - /** - * Return the {@link InetSocketAddress} for connecting to the chained proxy. - * Returning null indicates that we won't chain. - * - * @return The Chain Proxy with Host and Port. - */ - InetSocketAddress getChainedProxyAddress(); + /** + * Return the {@link InetSocketAddress} for connecting to the chained proxy. Returning null + * indicates that we won't chain. + * + * @return The Chain Proxy with Host and Port. + */ + InetSocketAddress getChainedProxyAddress(); - /** - * (Optional) ensure that the connection is opened from a specific local - * address (useful when doing NAT traversal). - */ - InetSocketAddress getLocalAddress(); + /** + * (Optional) ensure that the connection is opened from a specific local address (useful when + * doing NAT traversal). + */ + InetSocketAddress getLocalAddress(); - /** - * Tell LittleProxy what kind of TransportProtocol to use to communicate - * with the chained proxy. - */ - TransportProtocol getTransportProtocol(); + /** + * Tell LittleProxy what kind of TransportProtocol to use to communicate with the chained proxy. + */ + TransportProtocol getTransportProtocol(); - /** - * Tell LittleProxy the type of chained proxy that it will be - * connecting to. This setting determines what type of requests - * LittleProxy will use to communicate with the chained proxy. - * @return the chained proxy type. - */ - ChainedProxyType getChainedProxyType(); + /** + * Tell LittleProxy the type of chained proxy that it will be connecting to. This setting + * determines what type of requests LittleProxy will use to communicate with the chained proxy. + * + * @return the chained proxy type. + */ + ChainedProxyType getChainedProxyType(); - /** - * (Optional) implement this method if the chained proxy requires - * a username. - * @return the username to send to the chained proxy. - */ - String getUsername(); + /** + * (Optional) implement this method if the chained proxy requires a username. + * + * @return the username to send to the chained proxy. + */ + String getUsername(); - /** - * (Optional) implement this method if the chained proxy requires - * a password. - * @return the password to send to the chained proxy. - */ - String getPassword(); + /** + * (Optional) implement this method if the chained proxy requires a password. + * + * @return the password to send to the chained proxy. + */ + String getPassword(); - /** - * Implement this method to tell LittleProxy whether or not to encrypt - * connections to the chained proxy for the given request. If true, - * LittleProxy will call {@link SslEngineSource#newSslEngine()} to obtain an - * SSLContext used by the downstream proxy. - * - * @return true of the connection to the chained proxy should be encrypted - */ - boolean requiresEncryption(); + /** + * Implement this method to tell LittleProxy whether to encrypt connections to the chained proxy + * for the given request. If true, LittleProxy will call {@link SslEngineSource#newSslEngine()} to + * obtain an SSLContext used by the downstream proxy. + * + * @return true of the connection to the chained proxy should be encrypted + */ + boolean requiresEncryption(); - /** - * Filters requests on their way to the chained proxy. - */ - void filterRequest(HttpObject httpObject); + /** Filters requests on their way to the chained proxy. */ + void filterRequest(HttpObject httpObject); - /** - * Called to let us know that connecting to this proxy succeeded. - */ - void connectionSucceeded(); + /** Called to let us know that connecting to this proxy succeeded. */ + void connectionSucceeded(); - /** - * Called to let us know that connecting to this proxy failed. - * - * @param cause - * exception that caused this failure (may be null) - */ - void connectionFailed(Throwable cause); + /** + * Called to let us know that connecting to this proxy failed. + * + * @param cause exception that caused this failure (maybe null) + */ + void connectionFailed(Throwable cause); - /** - * Called to let us know that we were disconnected. - */ - void disconnected(); + /** Called to let us know that we were disconnected. */ + void disconnected(); } diff --git a/src/main/java/org/littleshoot/proxy/ChainedProxyAdapter.java b/src/main/java/org/littleshoot/proxy/ChainedProxyAdapter.java index 947c7486..21fd8c80 100644 --- a/src/main/java/org/littleshoot/proxy/ChainedProxyAdapter.java +++ b/src/main/java/org/littleshoot/proxy/ChainedProxyAdapter.java @@ -1,78 +1,71 @@ package org.littleshoot.proxy; import io.netty.handler.codec.http.HttpObject; - -import javax.net.ssl.SSLEngine; import java.net.InetSocketAddress; +import javax.net.ssl.SSLEngine; -/** - * Convenience base class for implementations of {@link ChainedProxy}. - */ +/** Convenience base class for implementations of {@link ChainedProxy}. */ public class ChainedProxyAdapter implements ChainedProxy { - /** - * {@link ChainedProxy} that simply has the downstream proxy make a direct - * connection to the upstream server. - */ - public static ChainedProxy FALLBACK_TO_DIRECT_CONNECTION = new ChainedProxyAdapter(); + /** + * {@link ChainedProxy} that simply has the downstream proxy make a direct connection to the + * upstream server. + */ + public static final ChainedProxy FALLBACK_TO_DIRECT_CONNECTION = new ChainedProxyAdapter(); + + @Override + public InetSocketAddress getChainedProxyAddress() { + return null; + } + + @Override + public InetSocketAddress getLocalAddress() { + return null; + } + + @Override + public TransportProtocol getTransportProtocol() { + return TransportProtocol.TCP; + } + + @Override + public ChainedProxyType getChainedProxyType() { + return ChainedProxyType.HTTP; + } + + @Override + public String getUsername() { + return null; + } - @Override - public InetSocketAddress getChainedProxyAddress() { - return null; - } + @Override + public String getPassword() { + return null; + } - @Override - public InetSocketAddress getLocalAddress() { - return null; - } + @Override + public boolean requiresEncryption() { + return false; + } - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public ChainedProxyType getChainedProxyType() { - return ChainedProxyType.HTTP; - } - - @Override - public String getUsername() { - return null; - } - - @Override - public String getPassword() { - return null; - } + @Override + public SSLEngine newSslEngine() { + return null; + } - @Override - public boolean requiresEncryption() { - return false; - } + @Override + public void filterRequest(HttpObject httpObject) {} - @Override - public SSLEngine newSslEngine() { - return null; - } - - @Override - public void filterRequest(HttpObject httpObject) { - } - - @Override - public void connectionSucceeded() { - } + @Override + public void connectionSucceeded() {} - @Override - public void connectionFailed(Throwable cause) { - } + @Override + public void connectionFailed(Throwable cause) {} - @Override - public void disconnected() { - } + @Override + public void disconnected() {} - @Override - public SSLEngine newSslEngine(String peerHost, int peerPort) { - return null; - } + @Override + public SSLEngine newSslEngine(String peerHost, int peerPort) { + return null; + } } diff --git a/src/main/java/org/littleshoot/proxy/ChainedProxyManager.java b/src/main/java/org/littleshoot/proxy/ChainedProxyManager.java index 4ebf3f3f..ee2e92ea 100644 --- a/src/main/java/org/littleshoot/proxy/ChainedProxyManager.java +++ b/src/main/java/org/littleshoot/proxy/ChainedProxyManager.java @@ -1,37 +1,23 @@ package org.littleshoot.proxy; import io.netty.handler.codec.http.HttpRequest; -import org.littleshoot.proxy.impl.ClientDetails; - import java.util.Queue; +import org.littleshoot.proxy.impl.ClientDetails; -/** - *

- * Interface for classes that manage chained proxies. - *

- */ +/** Interface for classes that manage chained proxies. */ public interface ChainedProxyManager { - /** - *

- * Based on the given httpRequest, add any {@link ChainedProxy}s to the list - * that should be used to process the request. The downstream proxy will - * attempt to connect to each of these in the order that they appear until - * it successfully connects to one. - *

- * - *

- * To allow the proxy to fall back to a direct connection, you can add - * {@link ChainedProxyAdapter#FALLBACK_TO_DIRECT_CONNECTION} to the end of - * the list. - *

- * - *

- * To keep the proxy from attempting any connection, leave the list blank. - * This will cause the proxy to return a 502 response. - *

- */ - void lookupChainedProxies(HttpRequest httpRequest, - Queue chainedProxies, - ClientDetails clientDetails); -} \ No newline at end of file + /** + * Based on the given httpRequest, add any {@link ChainedProxy}s to the list that should be used + * to process the request. The downstream proxy will attempt to connect to each of these in the + * order that they appear until it successfully connects to one. + * + *

To allow the proxy to fall back to a direct connection, you can add {@link + * ChainedProxyAdapter#FALLBACK_TO_DIRECT_CONNECTION} to the end of the list. + * + *

To keep the proxy from attempting any connection, leave the list blank. This will cause the + * proxy to return a 502 response. + */ + void lookupChainedProxies( + HttpRequest httpRequest, Queue chainedProxies, ClientDetails clientDetails); +} diff --git a/src/main/java/org/littleshoot/proxy/ChainedProxyType.java b/src/main/java/org/littleshoot/proxy/ChainedProxyType.java index 1e9a4272..e75c1291 100644 --- a/src/main/java/org/littleshoot/proxy/ChainedProxyType.java +++ b/src/main/java/org/littleshoot/proxy/ChainedProxyType.java @@ -1,8 +1,8 @@ package org.littleshoot.proxy; -/** - * Enumeration of chained proxy types supported by LittleProxy. - */ +/** Enumeration of chained proxy types supported by LittleProxy. */ public enum ChainedProxyType { - HTTP, SOCKS4, SOCKS5 + HTTP, + SOCKS4, + SOCKS5 } diff --git a/src/main/java/org/littleshoot/proxy/DefaultHostResolver.java b/src/main/java/org/littleshoot/proxy/DefaultHostResolver.java index 3519dcd9..ae274cb0 100644 --- a/src/main/java/org/littleshoot/proxy/DefaultHostResolver.java +++ b/src/main/java/org/littleshoot/proxy/DefaultHostResolver.java @@ -5,14 +5,13 @@ import java.net.UnknownHostException; /** - * Default implementation of {@link HostResolver} that just uses - * {@link InetAddress#getByName(String)}. + * Default implementation of {@link HostResolver} that just uses {@link + * InetAddress#getByName(String)}. */ public class DefaultHostResolver implements HostResolver { - @Override - public InetSocketAddress resolve(String host, int port) - throws UnknownHostException { - InetAddress addr = InetAddress.getByName(host); - return new InetSocketAddress(addr, port); - } + @Override + public InetSocketAddress resolve(String host, int port) throws UnknownHostException { + InetAddress address = InetAddress.getByName(host); + return new InetSocketAddress(address, port); + } } diff --git a/src/main/java/org/littleshoot/proxy/DnsSecServerResolver.java b/src/main/java/org/littleshoot/proxy/DnsSecServerResolver.java index 90ee9d0d..2c35fa5d 100644 --- a/src/main/java/org/littleshoot/proxy/DnsSecServerResolver.java +++ b/src/main/java/org/littleshoot/proxy/DnsSecServerResolver.java @@ -1,14 +1,12 @@ package org.littleshoot.proxy; -import org.littleshoot.dnssec4j.VerifiedAddressFactory; - import java.net.InetSocketAddress; import java.net.UnknownHostException; +import org.littleshoot.dnssec4j.VerifiedAddressFactory; public class DnsSecServerResolver implements HostResolver { - @Override - public InetSocketAddress resolve(String host, int port) - throws UnknownHostException { - return VerifiedAddressFactory.newInetSocketAddress(host, port, true); - } + @Override + public InetSocketAddress resolve(String host, int port) throws UnknownHostException { + return VerifiedAddressFactory.newInetSocketAddress(host, port, true); + } } diff --git a/src/main/java/org/littleshoot/proxy/FlowContext.java b/src/main/java/org/littleshoot/proxy/FlowContext.java index ae4230fa..6eaa0f2b 100644 --- a/src/main/java/org/littleshoot/proxy/FlowContext.java +++ b/src/main/java/org/littleshoot/proxy/FlowContext.java @@ -1,42 +1,100 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.impl.ClientToProxyConnection; - +import io.netty.handler.codec.haproxy.HAProxyMessage; +import java.net.InetSocketAddress; +import java.util.Map; +import java.util.Objects; +import java.util.concurrent.ConcurrentHashMap; import javax.net.ssl.SSLEngine; import javax.net.ssl.SSLSession; -import java.net.InetSocketAddress; +import org.littleshoot.proxy.impl.ClientToProxyConnection; /** - *

- * Encapsulates contextual information for flow information that's being - * reported to a {@link ActivityTracker}. - *

+ * Encapsulates contextual information for flow information that's being reported to a {@link + * ActivityTracker}. */ public class FlowContext { - private final InetSocketAddress clientAddress; - private final SSLSession clientSslSession; - - public FlowContext(ClientToProxyConnection clientConnection) { - super(); - this.clientAddress = clientConnection.getClientAddress(); - SSLEngine sslEngine = clientConnection.getSslEngine(); - this.clientSslSession = sslEngine != null ? sslEngine.getSession() - : null; - } + private final ClientToProxyConnection clientConnection; + private final long connectionId; + private final Map timingData = new ConcurrentHashMap<>(); - /** - * The address of the client. - */ - public InetSocketAddress getClientAddress() { - return clientAddress; - } + /** + * Creates a new FlowContext for the given client connection. + * + * @param clientConnection the client-side connection that owns this flow + */ + public FlowContext(ClientToProxyConnection clientConnection) { + this.clientConnection = clientConnection; + this.connectionId = clientConnection.getId(); + } - /** - * If using SSL, this returns the {@link SSLSession} on the client - * connection. - */ - public SSLSession getClientSslSession() { - return clientSslSession; + /** + * The client's address: the PROXY header's source address when PROXY protocol is in use, + * otherwise the TCP peer. Resolved lazily so a header received after construction is reflected. + * + * @return the client's socket address + */ + public InetSocketAddress getClientAddress() { + HAProxyMessage haProxyMessage = clientConnection.getHaProxyMessage(); + if (haProxyMessage != null + && haProxyMessage.sourceAddress() != null + && !haProxyMessage.sourceAddress().isBlank()) { + return new InetSocketAddress(haProxyMessage.sourceAddress(), haProxyMessage.sourcePort()); } + return clientConnection.getClientAddress(); + } + + /** + * If using SSL, this returns the {@link SSLSession} on the client connection. + * + * @return the SSL session, or null if the client connection is not using SSL + */ + public SSLSession getClientSslSession() { + SSLEngine sslEngine = clientConnection.getSslEngine(); + return sslEngine != null ? sslEngine.getSession() : null; + } + + /** + * Stores timing data for this flow. + * + * @param key the timing metric key + * @param value the timing value in milliseconds + */ + public void setTimingData(String key, Long value) { + Objects.requireNonNull(key, "timing key must not be null"); + Objects.requireNonNull(value, "timing value must not be null"); + timingData.put(key, value); + } + + /** + * Retrieves timing data for this flow. + * + * @param key the timing metric key + * @return the timing value in milliseconds, or null if not available + */ + public Long getTimingData(String key) { + return timingData.get(key); + } + + /** + * Gets all timing data for this flow. + * + * @return map of all timing data + */ + public Map getTimings() { + return Map.copyOf(timingData); + } + + @Override + public boolean equals(Object o) { + if (this == o) return true; + if (!(o instanceof FlowContext)) return false; + FlowContext that = (FlowContext) o; + return connectionId == that.connectionId; + } + @Override + public int hashCode() { + return Long.hashCode(connectionId); + } } diff --git a/src/main/java/org/littleshoot/proxy/FullFlowContext.java b/src/main/java/org/littleshoot/proxy/FullFlowContext.java index cb8ac1d7..94ac86a7 100644 --- a/src/main/java/org/littleshoot/proxy/FullFlowContext.java +++ b/src/main/java/org/littleshoot/proxy/FullFlowContext.java @@ -1,35 +1,38 @@ package org.littleshoot.proxy; +import io.netty.channel.ChannelHandlerContext; import org.littleshoot.proxy.impl.ClientToProxyConnection; import org.littleshoot.proxy.impl.ProxyToServerConnection; /** - * Extension of {@link FlowContext} that provides additional information (which - * we know after actually processing the request from the client). + * Extension of {@link FlowContext} that provides additional information (which we know after + * actually processing the request from the client). */ public class FullFlowContext extends FlowContext { - private final String serverHostAndPort; - private final ChainedProxy chainedProxy; + private final String serverHostAndPort; + private final ChainedProxy chainedProxy; + private final ChannelHandlerContext ctx; - public FullFlowContext(ClientToProxyConnection clientConnection, - ProxyToServerConnection serverConnection) { - super(clientConnection); - this.serverHostAndPort = serverConnection.getServerHostAndPort(); - this.chainedProxy = serverConnection.getChainedProxy(); - } + public FullFlowContext( + ClientToProxyConnection clientConnection, ProxyToServerConnection serverConnection) { + super(clientConnection); + serverHostAndPort = serverConnection.getServerHostAndPort(); + chainedProxy = serverConnection.getChainedProxy(); + this.ctx = serverConnection.getContext(); + } - /** - * The host and port for the server (i.e. the ultimate endpoint). - */ - public String getServerHostAndPort() { - return serverHostAndPort; - } + /** The host and port for the server (i.e. the ultimate endpoint). */ + public String getServerHostAndPort() { + return serverHostAndPort; + } - /** - * The chained proxy (if proxy chaining). - */ - public ChainedProxy getChainedProxy() { - return chainedProxy; - } + /** The chained proxy (if proxy chaining). */ + public ChainedProxy getChainedProxy() { + return chainedProxy; + } + /** The proxy to server channel context. */ + public ChannelHandlerContext getProxyToServerContext() { + return ctx; + } } diff --git a/src/main/java/org/littleshoot/proxy/HostResolver.java b/src/main/java/org/littleshoot/proxy/HostResolver.java index 17197c9f..e3fb7eef 100644 --- a/src/main/java/org/littleshoot/proxy/HostResolver.java +++ b/src/main/java/org/littleshoot/proxy/HostResolver.java @@ -3,9 +3,7 @@ import java.net.InetSocketAddress; import java.net.UnknownHostException; -/** - * Resolves host and port into an InetSocketAddress. - */ +/** Resolves host and port into an InetSocketAddress. */ public interface HostResolver { - InetSocketAddress resolve(String host, int port) throws UnknownHostException; + InetSocketAddress resolve(String host, int port) throws UnknownHostException; } diff --git a/src/main/java/org/littleshoot/proxy/HttpFilters.java b/src/main/java/org/littleshoot/proxy/HttpFilters.java index b102d1d4..909919a3 100644 --- a/src/main/java/org/littleshoot/proxy/HttpFilters.java +++ b/src/main/java/org/littleshoot/proxy/HttpFilters.java @@ -1,211 +1,210 @@ package org.littleshoot.proxy; import io.netty.channel.ChannelHandlerContext; -import io.netty.handler.codec.http.*; -import org.littleshoot.proxy.impl.ProxyUtils; - +import io.netty.handler.codec.http.HttpContent; +import io.netty.handler.codec.http.HttpObject; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; +import io.netty.handler.codec.http.LastHttpContent; import java.net.InetSocketAddress; +import java.util.function.Supplier; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.impl.ProxyUtils; /** - *

- * Interface for objects that filter {@link HttpObject}s, including both - * requests and responses, and informs of different steps in request/response. - *

- * - *

- * Multiple methods are defined, corresponding to different steps in the request - * processing lifecycle. Some of these methods is given the current object - * (request, response or chunk) and is allowed to modify it in place. Others - * provide a notification of when specific operations happen (i.e. connection in - * queue, DNS resolution, SSL handshaking and so forth). - *

- * - *

- * Because HTTP transfers can be chunked, for any given request or response, the - * filter methods that can modify request/response in place may be called - * multiple times, once for the initial {@link HttpRequest} or - * {@link HttpResponse}, and once for each subsequent {@link HttpContent}. The - * last chunk will always be a {@link LastHttpContent} and can be checked for - * being last using {@link ProxyUtils#isLastChunk(HttpObject)}. - *

- * - *

- * {@link HttpFiltersSource#getMaximumRequestBufferSizeInBytes()} and - * {@link HttpFiltersSource#getMaximumResponseBufferSizeInBytes()} can be used - * to instruct the proxy to buffer the {@link HttpObject}s sent to all of its - * request/response filters, in which case it will buffer up to the specified - * limit and then send either complete {@link HttpRequest}s or - * {@link HttpResponse}s to the filter methods. When buffering, if the proxy - * receives more data than fits in the specified maximum bytes to buffer, the - * proxy will stop processing the request and respond with a 502 Bad Gateway - * error. - *

- * - *

- * A new instance of {@link HttpFilters} is created for each request, so these - * objects can be stateful. - *

- * - *

- * To monitor (and time measure?) the different steps the request/response goes - * through, many informative methods are provided. Those steps are reported in - * the following order: + * Interface for objects that filter {@link HttpObject}s, including both requests and responses, and + * informs of different steps in request/response. + * + *

Multiple methods are defined, corresponding to different steps in the request processing + * lifecycle. Some of these methods is given the current object (request, response or chunk) and is + * allowed to modify it in place. Others provide a notification of when specific operations happen + * (i.e. connection in queue, DNS resolution, SSL handshaking and so forth). + * + *

Because HTTP transfers can be chunked, for any given request or response, the filter methods + * that can modify request/response in place may be called multiple times, once for the initial + * {@link HttpRequest} or {@link HttpResponse}, and once for each subsequent {@link HttpContent}. + * The last chunk will always be a {@link LastHttpContent} and can be checked for being last using + * {@link ProxyUtils#isLastChunk(HttpObject)}. + * + *

{@link HttpFiltersSource#getMaximumRequestBufferSizeInBytes()} and {@link + * HttpFiltersSource#getMaximumResponseBufferSizeInBytes()} can be used to instruct the proxy to + * buffer the {@link HttpObject}s sent to all of its request/response filters, in which case it will + * buffer up to the specified limit and then send either complete {@link HttpRequest}s or {@link + * HttpResponse}s to the filter methods. When buffering, if the proxy receives more data than fits + * in the specified maximum bytes to buffer, the proxy will stop processing the request and respond + * with a 502 Bad Gateway error. + * + *

A new instance of {@link HttpFilters} is created for each request, so these objects can be + * stateful. + * + *

To monitor (and time measure?) the different steps the request/response goes through, many + * informative methods are provided. Those steps are reported in the following order: + * *

    - *
  1. clientToProxyRequest
  2. - *
  3. proxyToServerConnectionQueued
  4. - *
  5. proxyToServerResolutionStarted
  6. - *
  7. proxyToServerResolutionSucceeded
  8. - *
  9. proxyToServerRequest (can be multiple if chunked)
  10. - *
  11. proxyToServerConnectionStarted
  12. - *
  13. proxyToServerConnectionFailed (if connection couldn't be established)
  14. - *
  15. proxyToServerConnectionSSLHandshakeStarted (only if HTTPS required)
  16. - *
  17. proxyToServerConnectionSucceeded
  18. - *
  19. proxyToServerRequestSending
  20. - *
  21. proxyToServerRequestSent
  22. - *
  23. serverToProxyResponseReceiving
  24. - *
  25. serverToProxyResponse (can be multiple if chuncked)
  26. - *
  27. serverToProxyResponseReceived
  28. - *
  29. proxyToClientResponse
  30. + *
  31. clientToProxyRequest + *
  32. proxyToServerConnectionQueued + *
  33. proxyToServerResolutionStarted + *
  34. proxyToServerResolutionSucceeded + *
  35. proxyToServerRequest (can be multiple if chunked) + *
  36. proxyToServerConnectionStarted + *
  37. proxyToServerConnectionFailed (if connection couldn't be established) + *
  38. proxyToServerConnectionSSLHandshakeStarted (only if HTTPS required) + *
  39. proxyToServerConnectionSucceeded + *
  40. proxyToServerRequestSending + *
  41. proxyToServerRequestSent + *
  42. serverToProxyResponseReceiving + *
  43. serverToProxyResponse (can be multiple if chunked) + *
  44. serverToProxyResponseReceived + *
  45. proxyToClientResponse *
*/ public interface HttpFilters { - /** - * Filters requests on their way from the client to the proxy. To interrupt processing of this request and return a - * response to the client immediately, return an HttpResponse here. Otherwise, return null to continue processing as - * usual. - *

- * Important: When returning a response, you must include a mechanism to allow the client to determine the length - * of the message (see RFC 7230, section 3.3.3: https://tools.ietf.org/html/rfc7230#section-3.3.3 ). For messages that - * may contain a body, you may do this by setting the Transfer-Encoding to chunked, setting an appropriate - * Content-Length, or by adding a "Connection: close" header to the response (which will instruct LittleProxy to close - * the connection). If the short-circuit response contains body content, it is recommended that you return a - * FullHttpResponse. - * - * @param httpObject Client to Proxy HttpRequest (and HttpContent, if chunked) - * @return a short-circuit response, or null to continue processing as usual - */ - HttpResponse clientToProxyRequest(HttpObject httpObject); - - /** - * Filters requests on their way from the proxy to the server. To interrupt processing of this request and return a - * response to the client immediately, return an HttpResponse here. Otherwise, return null to continue processing as - * usual. - *

- * Important: When returning a response, you must include a mechanism to allow the client to determine the length - * of the message (see RFC 7230, section 3.3.3: https://tools.ietf.org/html/rfc7230#section-3.3.3 ). For messages that - * may contain a body, you may do this by setting the Transfer-Encoding to chunked, setting an appropriate - * Content-Length, or by adding a "Connection: close" header to the response. (which will instruct LittleProxy to close - * the connection). If the short-circuit response contains body content, it is recommended that you return a - * FullHttpResponse. - * - * @param httpObject Proxy to Server HttpRequest (and HttpContent, if chunked) - * @return a short-circuit response, or null to continue processing as usual - */ - HttpResponse proxyToServerRequest(HttpObject httpObject); - - /** - * Informs filter that proxy to server request is being sent. - */ - void proxyToServerRequestSending(); - - /** - * Informs filter that the HTTP request, including any content, has been sent. - */ - void proxyToServerRequestSent(); - - /** - * Filters responses on their way from the server to the proxy. - * - * @param httpObject - * Server to Proxy HttpResponse (and HttpContent, if chunked) - * @return the modified (or unmodified) HttpObject. Returning null will - * force a disconnect. - */ - HttpObject serverToProxyResponse(HttpObject httpObject); - - /** - * Informs filter that a timeout occurred before the server response was received by the client. The timeout may have - * occurred while the client was sending the request, waiting for a response, or after the client started receiving - * a response (i.e. if the response from the server "stalls"). - * - * See {@link HttpProxyServerBootstrap#withIdleConnectionTimeout(int)} for information on setting the timeout. - */ - void serverToProxyResponseTimedOut(); - - /** - * Informs filter that server to proxy response is being received. - */ - void serverToProxyResponseReceiving(); - - /** - * Informs filter that server to proxy response has been received. - */ - void serverToProxyResponseReceived(); - - /** - * Filters responses on their way from the proxy to the client. - * - * @param httpObject - * Proxy to Client HttpResponse (and HttpContent, if chunked) - * @return the modified (or unmodified) HttpObject. Returning null will - * force a disconnect. - */ - HttpObject proxyToClientResponse(HttpObject httpObject); - - /** - * Informs filter that proxy to server connection is in queue. - */ - void proxyToServerConnectionQueued(); - - /** - * Filter DNS resolution from proxy to server. - * - * @param resolvingServerHostAndPort - * Server "HOST:PORT" - * @return alternative address resolution. Returning null will let normal - * DNS resolution continue. - */ - InetSocketAddress proxyToServerResolutionStarted( - String resolvingServerHostAndPort); - - /** - * Informs filter that proxy to server DNS resolution failed for the specified host and port. - * - * @param hostAndPort hostname and port the proxy failed to resolve - */ - void proxyToServerResolutionFailed(String hostAndPort); - - /** - * Informs filter that proxy to server DNS resolution has happened. - * - * @param serverHostAndPort - * Server "HOST:PORT" - * @param resolvedRemoteAddress - * Address it was proxyToServerResolutionSucceeded to - */ - void proxyToServerResolutionSucceeded(String serverHostAndPort, - InetSocketAddress resolvedRemoteAddress); - - /** - * Informs filter that proxy to server connection is initiating. - */ - void proxyToServerConnectionStarted(); - - /** - * Informs filter that proxy to server ssl handshake is initiating. - */ - void proxyToServerConnectionSSLHandshakeStarted(); - - /** - * Informs filter that proxy to server connection has failed. - */ - void proxyToServerConnectionFailed(); - - /** - * Informs filter that proxy to server connection has succeeded. - * - * @param serverCtx the {@link io.netty.channel.ChannelHandlerContext} used to connect to the server - */ - void proxyToServerConnectionSucceeded(ChannelHandlerContext serverCtx); - + /** + * Filters requests on their way from the client to the proxy. To interrupt processing of this + * request and return a response to the client immediately, return an HttpResponse here. + * Otherwise, return null to continue processing as usual. + * + *

Important: When returning a response, you must include a mechanism to allow the + * client to determine the length of the message (see RFC 7230, section 3.3.3 ). + * + *

For messages that may contain a body, you may do this by setting the Transfer-Encoding to + * chunked, setting an appropriate Content-Length, or by adding a "Connection: close" header to + * the response (which will instruct LittleProxy to close the connection). If the short-circuit + * response contains body content, it is recommended that you return a FullHttpResponse. + * + * @param httpObject Client to Proxy HttpRequest (and HttpContent, if chunked) + * @return a short-circuit response, or null to continue processing as usual + */ + @Nullable HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject); + + /** + * Filters requests on their way from the proxy to the server. To interrupt processing of this + * request and return a response to the client immediately, return an HttpResponse here. + * Otherwise, return null to continue processing as usual. + * + *

Important: When returning a response, you must include a mechanism to allow the + * client to determine the length of the message (see RFC 7230, section 3.3.3 ). For + * messages that may contain a body, you may do this by setting the Transfer-Encoding to chunked, + * setting an appropriate Content-Length, or by adding a "Connection: close" header to the + * response. (which will instruct LittleProxy to close the connection). If the short-circuit + * response contains body content, it is recommended that you return a FullHttpResponse. + * + * @param httpObject Proxy to Server HttpRequest (and HttpContent, if chunked) + * @return a short-circuit response, or null to continue processing as usual + */ + @Nullable HttpResponse proxyToServerRequest(@NonNull HttpObject httpObject); + + /** Informs filter that proxy to server request is being sent. */ + void proxyToServerRequestSending(); + + /** Informs filter that the HTTP request, including any content, has been sent. */ + void proxyToServerRequestSent(); + + /** + * Filters responses on their way from the server to the proxy. + * + * @param httpObject Server to Proxy HttpResponse (and HttpContent, if chunked) + * @return the modified (or unmodified) HttpObject. Returning null will force a disconnect. + */ + @Nullable HttpObject serverToProxyResponse(@NonNull HttpObject httpObject); + + /** + * Informs filter that a timeout occurred before the server response was received by the client. + * The timeout may have occurred while the client was sending the request, waiting for a response, + * or after the client started receiving a response (i.e. if the response from the server + * "stalls"). + * + *

See {@link HttpProxyServerBootstrap#withIdleConnectionTimeout(int)} for information on + * setting the timeout. + */ + void serverToProxyResponseTimedOut(); + + /** Informs filter that server to proxy response is being received. */ + void serverToProxyResponseReceiving(); + + /** Informs filter that server to proxy response has been received. */ + void serverToProxyResponseReceived(); + + /** + * Filters responses on their way from the proxy to the client. + * + * @param httpObject Proxy to Client HttpResponse (and HttpContent, if chunked) + * @return the modified (or unmodified) HttpObject. Returning null will force a disconnect. + */ + @Nullable HttpObject proxyToClientResponse(@NonNull HttpObject httpObject); + + /** Informs filter that proxy to server connection is in queue. */ + void proxyToServerConnectionQueued(); + + /** + * Filter DNS resolution from proxy to server. + * + * @param resolvingServerHostAndPort Server "HOST:PORT" + * @return alternative address resolution. Returning null will let normal DNS resolution continue. + */ + @Nullable InetSocketAddress proxyToServerResolutionStarted( + @NonNull String resolvingServerHostAndPort); + + /** + * Informs filter that proxy to server DNS resolution failed for the specified host and port. + * + * @param hostAndPort hostname and port the proxy failed to resolve + */ + void proxyToServerResolutionFailed(@NonNull String hostAndPort); + + /** + * Informs filter that proxy to server DNS resolution has happened. + * + * @param serverHostAndPort Server "HOST:PORT" + * @param resolvedRemoteAddress Address it was proxyToServerResolutionSucceeded to + */ + void proxyToServerResolutionSucceeded( + @NonNull String serverHostAndPort, @NonNull InetSocketAddress resolvedRemoteAddress); + + /** Informs filter that proxy to server connection is initiating. */ + void proxyToServerConnectionStarted(); + + /** Informs filter that proxy to server ssl handshake is initiating. */ + void proxyToServerConnectionSSLHandshakeStarted(); + + /** Informs filter that proxy to server connection has failed. */ + void proxyToServerConnectionFailed(); + + /** + * Informs filter that proxy to server connection has succeeded. + * + * @param serverCtx the {@link io.netty.channel.ChannelHandlerContext} used to connect to the + * server + */ + void proxyToServerConnectionSucceeded(@NonNull ChannelHandlerContext serverCtx); + + /** + * Allow this proxy to act as an SSL man in the middle. + * + *

Has no impact if man in the middle is not enabled. + * + * @return true to allow mitm, false to not mitm the proxy to server connection. + */ + boolean proxyToServerAllowMitm(); + + /** + * Notifies the filter that a WebSocket frame has been received and is about to be forwarded. + * Called after the HTTP connection has been upgraded to WebSocket. + * + *

The {@code frameBytes} contain the raw, unmodified WebSocket frame as received from the + * network. Client-to-server frames are masked per RFC 6455; server-to-client frames are not. + * + *

This method is informational — the frame cannot be modified or suppressed here. + * + *

Important: The {@code frameBytes} supplier must be called synchronously within this + * method. Storing the supplier for later invocation will result in undefined behavior as the + * underlying buffer is released after this method returns. + * + * @param frameBytes the raw bytes of the WebSocket frame + * @param fromClient true if the frame was sent by the client, false if sent by the server + */ + default void webSocketFrameReceived(Supplier frameBytes, boolean fromClient) {} } diff --git a/src/main/java/org/littleshoot/proxy/HttpFiltersAdapter.java b/src/main/java/org/littleshoot/proxy/HttpFiltersAdapter.java index 2871364a..8d901cf6 100644 --- a/src/main/java/org/littleshoot/proxy/HttpFiltersAdapter.java +++ b/src/main/java/org/littleshoot/proxy/HttpFiltersAdapter.java @@ -4,103 +4,98 @@ import io.netty.handler.codec.http.HttpObject; import io.netty.handler.codec.http.HttpRequest; import io.netty.handler.codec.http.HttpResponse; - import java.net.InetSocketAddress; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; -/** - * Convenience base class for implementations of {@link HttpFilters}. - */ +/** Convenience base class for implementations of {@link HttpFilters}. */ +@NullMarked public class HttpFiltersAdapter implements HttpFilters { - /** - * A default, stateless, no-op {@link HttpFilters} instance. - */ - public static final HttpFiltersAdapter NOOP_FILTER = new HttpFiltersAdapter(null); - - protected final HttpRequest originalRequest; - protected final ChannelHandlerContext ctx; - - public HttpFiltersAdapter(HttpRequest originalRequest, - ChannelHandlerContext ctx) { - this.originalRequest = originalRequest; - this.ctx = ctx; - } - - public HttpFiltersAdapter(HttpRequest originalRequest) { - this(originalRequest, null); - } - - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - return null; - } - - @Override - public HttpResponse proxyToServerRequest(HttpObject httpObject) { - return null; - } - - @Override - public void proxyToServerRequestSending() { - } - - @Override - public void proxyToServerRequestSent() { - } - - @Override - public HttpObject serverToProxyResponse(HttpObject httpObject) { - return httpObject; - } - - @Override - public void serverToProxyResponseTimedOut() { - } - - @Override - public void serverToProxyResponseReceiving() { - } - - @Override - public void serverToProxyResponseReceived() { - } - - @Override - public HttpObject proxyToClientResponse(HttpObject httpObject) { - return httpObject; - } - - @Override - public void proxyToServerConnectionQueued() { - } - - @Override - public InetSocketAddress proxyToServerResolutionStarted( - String resolvingServerHostAndPort) { - return null; - } - - @Override - public void proxyToServerResolutionFailed(String hostAndPort) { - } - - @Override - public void proxyToServerResolutionSucceeded(String serverHostAndPort, - InetSocketAddress resolvedRemoteAddress) { - } - - @Override - public void proxyToServerConnectionStarted() { - } - - @Override - public void proxyToServerConnectionSSLHandshakeStarted() { - } - - @Override - public void proxyToServerConnectionFailed() { - } - - @Override - public void proxyToServerConnectionSucceeded(ChannelHandlerContext serverCtx) { - } + /** A default, stateless, no-op {@link HttpFilters} instance. */ + public static final HttpFiltersAdapter NOOP_FILTER = new HttpFiltersAdapter(null); + + @Nullable protected final HttpRequest originalRequest; + @Nullable protected final ChannelHandlerContext ctx; + + public HttpFiltersAdapter( + @Nullable HttpRequest originalRequest, @Nullable ChannelHandlerContext ctx) { + this.originalRequest = originalRequest; + this.ctx = ctx; + } + + public HttpFiltersAdapter(@Nullable HttpRequest originalRequest) { + this(originalRequest, null); + } + + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + return null; + } + + @Nullable + @Override + public HttpResponse proxyToServerRequest(@NonNull HttpObject httpObject) { + return null; + } + + @Override + public void proxyToServerRequestSending() {} + + @Override + public void proxyToServerRequestSent() {} + + @Override + public HttpObject serverToProxyResponse(HttpObject httpObject) { + return httpObject; + } + + @Override + public void serverToProxyResponseTimedOut() {} + + @Override + public void serverToProxyResponseReceiving() {} + + @Override + public void serverToProxyResponseReceived() {} + + @Override + public HttpObject proxyToClientResponse(HttpObject httpObject) { + return httpObject; + } + + @Override + public void proxyToServerConnectionQueued() {} + + @Nullable + @Override + public InetSocketAddress proxyToServerResolutionStarted( + @NonNull String resolvingServerHostAndPort) { + return null; + } + + @Override + public void proxyToServerResolutionFailed(@NonNull String hostAndPort) {} + + @Override + public void proxyToServerResolutionSucceeded( + @NonNull String serverHostAndPort, @NonNull InetSocketAddress resolvedRemoteAddress) {} + + @Override + public void proxyToServerConnectionStarted() {} + + @Override + public void proxyToServerConnectionSSLHandshakeStarted() {} + + @Override + public void proxyToServerConnectionFailed() {} + + @Override + public void proxyToServerConnectionSucceeded(@NonNull ChannelHandlerContext serverCtx) {} + + @Override + public boolean proxyToServerAllowMitm() { + return true; + } } diff --git a/src/main/java/org/littleshoot/proxy/HttpFiltersSource.java b/src/main/java/org/littleshoot/proxy/HttpFiltersSource.java index e554e61a..04e01d00 100644 --- a/src/main/java/org/littleshoot/proxy/HttpFiltersSource.java +++ b/src/main/java/org/littleshoot/proxy/HttpFiltersSource.java @@ -5,39 +5,37 @@ import io.netty.handler.codec.http.FullHttpResponse; import io.netty.handler.codec.http.HttpRequest; import io.netty.handler.codec.http.HttpResponse; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; -/** - * Factory for {@link HttpFilters}. - */ +/** Factory for {@link HttpFilters}. */ +@NullMarked public interface HttpFiltersSource { - /** - * Return an {@link HttpFilters} object for this request if and only if we - * want to filter the request and/or its responses. - */ - HttpFilters filterRequest(HttpRequest originalRequest, - ChannelHandlerContext ctx); + /** + * Return an {@link HttpFilters} object for this request if and only if we want to filter the + * request and/or its responses. + */ + @Nullable HttpFilters filterRequest( + @NonNull HttpRequest originalRequest, @Nullable ChannelHandlerContext ctx); - /** - * Indicate how many (if any) bytes to buffer for incoming - * {@link HttpRequest}s. A value of 0 or less indicates that no buffering - * should happen and that messages will be passed to the {@link HttpFilters} - * request filtering methods chunk by chunk. A positive value will cause - * LittleProxy to try an create a {@link FullHttpRequest} using the data - * received from the client, with its content already decompressed (in case - * the client was compressing it). If the request size exceeds the maximum - * buffer size, the request will fail. - */ - int getMaximumRequestBufferSizeInBytes(); + /** + * Indicate how many (if any) bytes to buffer for incoming {@link HttpRequest}s. A value of 0 or + * less indicates that no buffering should happen and that messages will be passed to the {@link + * HttpFilters} request filtering methods chunk by chunk. A positive value will cause LittleProxy + * to try to create a {@link FullHttpRequest} using the data received from the client, with its + * content already decompressed (in case the client was compressing it). If the request size + * exceeds the maximum buffer size, the request will fail. + */ + int getMaximumRequestBufferSizeInBytes(); - /** - * Indicate how many (if any) bytes to buffer for incoming - * {@link HttpResponse}s. A value of 0 or less indicates that no buffering - * should happen and that messages will be passed to the {@link HttpFilters} - * response filtering methods chunk by chunk. A positive value will cause - * LittleProxy to try an create a {@link FullHttpResponse} using the data - * received from the server, with its content already decompressed (in case - * the server was compressing it). If the response size exceeds the maximum - * buffer size, the response will fail. - */ - int getMaximumResponseBufferSizeInBytes(); + /** + * Indicate how many (if any) bytes to buffer for incoming {@link HttpResponse}s. A value of 0 or + * less indicates that no buffering should happen and that messages will be passed to the {@link + * HttpFilters} response filtering methods chunk by chunk. A positive value will cause LittleProxy + * to try to create a {@link FullHttpResponse} using the data received from the server, with its + * content already decompressed (in case the server was compressing it). If the response size + * exceeds the maximum buffer size, the response will fail. + */ + int getMaximumResponseBufferSizeInBytes(); } diff --git a/src/main/java/org/littleshoot/proxy/HttpFiltersSourceAdapter.java b/src/main/java/org/littleshoot/proxy/HttpFiltersSourceAdapter.java index 1b2cecfd..1858d317 100644 --- a/src/main/java/org/littleshoot/proxy/HttpFiltersSourceAdapter.java +++ b/src/main/java/org/littleshoot/proxy/HttpFiltersSourceAdapter.java @@ -2,30 +2,33 @@ import io.netty.channel.ChannelHandlerContext; import io.netty.handler.codec.http.HttpRequest; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; -/** - * Convenience base class for implementations of {@link HttpFiltersSource}. - */ +/** Convenience base class for implementations of {@link HttpFiltersSource}. */ +@NullMarked public class HttpFiltersSourceAdapter implements HttpFiltersSource { - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest, null); - } - - @Override - public HttpFilters filterRequest(HttpRequest originalRequest, - ChannelHandlerContext ctx) { - return filterRequest(originalRequest); - } + @Nullable + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest, null); + } - @Override - public int getMaximumRequestBufferSizeInBytes() { - return 0; - } + @Override + @Nullable + public HttpFilters filterRequest( + @NonNull HttpRequest originalRequest, @NonNull ChannelHandlerContext ctx) { + return filterRequest(originalRequest); + } - @Override - public int getMaximumResponseBufferSizeInBytes() { - return 0; - } + @Override + public int getMaximumRequestBufferSizeInBytes() { + return 0; + } + @Override + public int getMaximumResponseBufferSizeInBytes() { + return 0; + } } diff --git a/src/main/java/org/littleshoot/proxy/HttpProxyServer.java b/src/main/java/org/littleshoot/proxy/HttpProxyServer.java index 47e40d16..a2a21a1f 100644 --- a/src/main/java/org/littleshoot/proxy/HttpProxyServer.java +++ b/src/main/java/org/littleshoot/proxy/HttpProxyServer.java @@ -1,63 +1,44 @@ package org.littleshoot.proxy; import java.net.InetSocketAddress; +import java.time.Duration; -/** - * Interface for the top-level proxy server class. - */ +/** Interface for the top-level proxy server class. */ public interface HttpProxyServer { - int getIdleConnectionTimeout(); - - void setIdleConnectionTimeout(int idleConnectionTimeout); - - /** - * Returns the maximum time to wait, in milliseconds, to connect to a server. - */ - int getConnectTimeout(); - - /** - * Sets the maximum time to wait, in milliseconds, to connect to a server. - */ - void setConnectTimeout(int connectTimeoutMs); - - /** - *

- * Clone the existing server, with a port 1 higher and everything else the - * same. If the proxy was started with port 0 (JVM-assigned port), the cloned proxy will also use a JVM-assigned - * port. - *

- * - *

- * The new server will share event loops with the original server. The event - * loops will use whatever name was given to the first server in the clone - * group. The server group will not terminate until the original server and all clones terminate. - *

- * - * @return a bootstrap that allows customizing and starting the cloned - * server - */ - HttpProxyServerBootstrap clone(); - - /** - * Stops the server and all related clones. Waits for traffic to stop before shutting down. - */ - void stop(); - - /** - * Stops the server and all related clones immediately, without waiting for traffic to stop. - */ - void abort(); - - /** - * Return the address on which this proxy is listening. - */ - InetSocketAddress getListenAddress(); - - /** - *

- * Set the read/write throttle bandwidths (in bytes/second) for this proxy. - *

- */ - void setThrottle(long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond); + int getIdleConnectionTimeout(); + + void setIdleConnectionTimeout(int idleConnectionTimeoutInSeconds); + + void setIdleConnectionTimeout(Duration idleConnectionTimeout); + + /** Returns the maximum time to wait, in milliseconds, to connect to a server. */ + int getConnectTimeout(); + + /** Sets the maximum time to wait, in milliseconds, to connect to a server. */ + void setConnectTimeout(int connectTimeoutMs); + + /** + * Clone the existing server, with a port 1 higher and everything else the same. If the proxy was + * started with port 0 (JVM-assigned port), the cloned proxy will also use a JVM-assigned port. + * + *

The new server will share event loops with the original server. The event loops will use + * whatever name was given to the first server in the clone group. The server group will not + * terminate until the original server and all clones terminate. + * + * @return a bootstrap that allows customizing and starting the cloned server + */ + HttpProxyServerBootstrap clone(); + + /** Stops the server and all related clones. Waits for traffic to stop before shutting down. */ + void stop(); + + /** Stops the server and all related clones immediately, without waiting for traffic to stop. */ + void abort(); + + /** Return the address on which this proxy is listening. */ + InetSocketAddress getListenAddress(); + + /** Set the read/write throttle bandwidths (in bytes/second) for this proxy. */ + void setThrottle(long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond); } diff --git a/src/main/java/org/littleshoot/proxy/HttpProxyServerBootstrap.java b/src/main/java/org/littleshoot/proxy/HttpProxyServerBootstrap.java index 5ef9075c..35ce454e 100644 --- a/src/main/java/org/littleshoot/proxy/HttpProxyServerBootstrap.java +++ b/src/main/java/org/littleshoot/proxy/HttpProxyServerBootstrap.java @@ -1,306 +1,298 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.impl.ThreadPoolConfiguration; -import org.littleshoot.proxy.impl.ServerGroup; +import com.google.errorprone.annotations.CanIgnoreReturnValue; import java.net.InetSocketAddress; +import java.time.Duration; +import org.jspecify.annotations.NullMarked; +import org.littleshoot.proxy.impl.ServerGroup; +import org.littleshoot.proxy.impl.ThreadPoolConfiguration; /** - * Configures and starts an {@link HttpProxyServer}. The HttpProxyServer is - * built using {@link #start()}. Sensible defaults are available for all - * parameters such that {@link #start()} could be called immediately if you - * wish. + * Configures and starts an {@link HttpProxyServer}. The HttpProxyServer is built using {@link + * #start()}. Sensible defaults are available for all parameters such that {@link #start()} could be + * called immediately if you wish. */ +@NullMarked public interface HttpProxyServerBootstrap { - /** - *

- * Give the server a name (used for naming threads, useful for logging). - *

- * - *

- * Default = LittleProxy - *

- */ - HttpProxyServerBootstrap withName(String name); - - /** - *

- * Specify the {@link TransportProtocol} to use for incoming connections. - *

- * - *

- * Default = TCP - *

- */ - HttpProxyServerBootstrap withTransportProtocol( - TransportProtocol transportProtocol); - - /** - *

- * Listen for incoming connections on the given address. - *

- * - *

- * Default = [bound ip]:8080 - *

- */ - HttpProxyServerBootstrap withAddress(InetSocketAddress address); - - /** - *

- * Listen for incoming connections on the given port. - *

- * - *

- * Default = 8080 - *

- */ - HttpProxyServerBootstrap withPort(int port); - - /** - *

- * Specify whether or not to only allow local connections. - *

- * - *

- * Default = true - *

- */ - HttpProxyServerBootstrap withAllowLocalOnly(boolean allowLocalOnly); - - /** - * This method has no effect and will be removed in a future release. - * @deprecated use {@link #withNetworkInterface(InetSocketAddress)} to avoid listening on all local addresses - */ - @Deprecated - HttpProxyServerBootstrap withListenOnAllAddresses(boolean listenOnAllAddresses); - - /** - *

- * Specify an {@link SslEngineSource} to use for encrypting inbound - * connections. Enabling this will enable SSL client authentication - * by default (see {@link #withAuthenticateSslClients(boolean)}) - *

- * - *

- * Default = null - *

- * - *

- * Note - This and {@link #withManInTheMiddle(MitmManager)} are - * mutually exclusive. - *

- */ - HttpProxyServerBootstrap withSslEngineSource( - SslEngineSource sslEngineSource); - - /** - *

- * Specify whether or not to authenticate inbound SSL clients (only applies - * if {@link #withSslEngineSource(SslEngineSource)} has been set). - *

- * - *

- * Default = true - *

- */ - HttpProxyServerBootstrap withAuthenticateSslClients( - boolean authenticateSslClients); - - /** - *

- * Specify a {@link ProxyAuthenticator} to use for doing basic HTTP - * authentication of clients. - *

- * - *

- * Default = null - *

- */ - HttpProxyServerBootstrap withProxyAuthenticator( - ProxyAuthenticator proxyAuthenticator); - - /** - *

- * Specify a {@link ChainedProxyManager} to use for chaining requests to - * another proxy. - *

- * - *

- * Default = null - *

- */ - HttpProxyServerBootstrap withChainProxyManager( - ChainedProxyManager chainProxyManager); - - /** - *

- * Specify an {@link MitmManager} to use for making this proxy act as an SSL - * man in the middle - *

- * - *

- * Default = null - *

- * - *

- * Note - This and {@link #withSslEngineSource(SslEngineSource)} are - * mutually exclusive. - *

- */ - HttpProxyServerBootstrap withManInTheMiddle( - MitmManager mitmManager); - - /** - *

- * Specify a {@link HttpFiltersSource} to use for filtering requests and/or - * responses through this proxy. - *

- * - *

- * Default = null - *

- */ - HttpProxyServerBootstrap withFiltersSource( - HttpFiltersSource filtersSource); - - /** - *

- * Specify whether or not to use secure DNS lookups for outbound - * connections. - *

- * - *

- * Default = false - *

- */ - HttpProxyServerBootstrap withUseDnsSec( - boolean useDnsSec); - - /** - *

- * Specify whether or not to run this proxy as a transparent proxy. - *

- * - *

- * Default = false - *

- */ - HttpProxyServerBootstrap withTransparent( - boolean transparent); - - /** - *

- * Specify the timeout after which to disconnect idle connections, in - * seconds. - *

- * - *

- * Default = 70 - *

- */ - HttpProxyServerBootstrap withIdleConnectionTimeout( - int idleConnectionTimeout); - - /** - *

- * Specify the timeout for connecting to the upstream server on a new - * connection, in milliseconds. - *

- * - *

- * Default = 40000 - *

- */ - HttpProxyServerBootstrap withConnectTimeout( - int connectTimeout); - - /** - * Specify a custom {@link HostResolver} for resolving server addresses. - */ - HttpProxyServerBootstrap withServerResolver(HostResolver serverResolver); - - /** - * Specify a custom {@link ServerGroup} to use for managing this server's resources and such. - * If one isn't provided, a default one will be created using the {@link ThreadPoolConfiguration} provided - * - * @param group A custom server group - */ - HttpProxyServerBootstrap withServerGroup(ServerGroup group); - - /** - *

- * Add an {@link ActivityTracker} for tracking activity in this proxy. - *

- */ - HttpProxyServerBootstrap plusActivityTracker(ActivityTracker activityTracker); - - /** - *

- * Specify the read and/or write bandwidth throttles for this proxy server. 0 indicates not throttling. - *

- */ - HttpProxyServerBootstrap withThrottling(long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond); - - /** - * All outgoing-communication of the proxy-instance is goin' to be routed via the given network-interface - * - * @param inetSocketAddress to be used for outgoing communication - */ - HttpProxyServerBootstrap withNetworkInterface(InetSocketAddress inetSocketAddress); - - HttpProxyServerBootstrap withMaxInitialLineLength(int maxInitialLineLength); - - HttpProxyServerBootstrap withMaxHeaderSize(int maxHeaderSize); - - HttpProxyServerBootstrap withMaxChunkSize(int maxChunkSize); - - /** - * When true, the proxy will accept requests that appear to be directed at an origin server (i.e. the URI in the HTTP - * request will contain an origin-form, rather than an absolute-form, as specified in RFC 7230, section 5.3). - * This is useful when the proxy is acting as a gateway/reverse proxy. Note: This feature should not be - * enabled when running as a forward proxy; doing so may cause an infinite loop if the client requests the URI of the proxy. - * - * @param allowRequestToOriginServer when true, the proxy will accept origin-form HTTP requests - */ - HttpProxyServerBootstrap withAllowRequestToOriginServer(boolean allowRequestToOriginServer); - - /** - * Sets the alias to use when adding Via headers to incoming and outgoing HTTP messages. The alias may be any - * pseudonym, or if not specified, defaults to the hostname of the local machine. See RFC 7230, section 5.7.1. - * - * @param alias the pseudonym to add to Via headers - */ - HttpProxyServerBootstrap withProxyAlias(String alias); - - /** - *

- * Build and starts the server. - *

- * - * @return the newly built and started server - */ - HttpProxyServer start(); - - /** - * Set the configuration parameters for the proxy's thread pools. - * - * @param configuration thread pool configuration - * @return proxy server bootstrap for chaining - */ - HttpProxyServerBootstrap withThreadPoolConfiguration(ThreadPoolConfiguration configuration); - - /** - * Specifies if the proxy server should accept a proxy protocol header. Once set it works with request that - * include a proxy protocol header. The proxy server reads an incoming proxy protocol header from the - * client. - * @param allowProxyProtocol when true, the proxy will accept a proxy protocol header - */ - HttpProxyServerBootstrap withAcceptProxyProtocol(boolean allowProxyProtocol); - - /** - * Specifies if the proxy server should send a proxy protocol header. - * @param sendProxyProtocol when true, the proxy will send a proxy protocol header - */ - HttpProxyServerBootstrap withSendProxyProtocol(boolean sendProxyProtocol); + /** + * Give the server a name (used for naming threads, useful for logging). + * + *

Default = LittleProxy + */ + HttpProxyServerBootstrap withName(String name); + + /** + * Listen for incoming connections on the given address. + * + *

Default = [bound ip]:8080 + */ + HttpProxyServerBootstrap withAddress(InetSocketAddress address); + + /** + * Listen for incoming connections on the given port. + * + *

Default = 8080 + */ + HttpProxyServerBootstrap withPort(int port); + + /** + * Specify whether or not to only allow local connections. + * + *

Default = true + */ + HttpProxyServerBootstrap withAllowLocalOnly(boolean allowLocalOnly); + + /** + * Specify an {@link SslEngineSource} to use for encrypting inbound connections. Enabling this + * will enable SSL client authentication by default (see {@link + * #withAuthenticateSslClients(boolean)}) + * + *

Default = null + * + *

Note - This and {@link #withManInTheMiddle(MitmManager)} are mutually exclusive. + */ + HttpProxyServerBootstrap withSslEngineSource(SslEngineSource sslEngineSource); + + /** + * Specify whether or not to authenticate inbound SSL clients (only applies if {@link + * #withSslEngineSource(SslEngineSource)} has been set). + * + *

Default = true + */ + HttpProxyServerBootstrap withAuthenticateSslClients(boolean authenticateSslClients); + + /** + * Specify a {@link ProxyAuthenticator} to use for doing basic HTTP authentication of clients. + * + *

Default = null + */ + HttpProxyServerBootstrap withProxyAuthenticator(ProxyAuthenticator proxyAuthenticator); + + /** + * Specify a {@link ChainedProxyManager} to use for chaining requests to another proxy. + * + *

Default = null + */ + HttpProxyServerBootstrap withChainProxyManager(ChainedProxyManager chainProxyManager); + + /** + * Specify an {@link MitmManager} to use for making this proxy act as an SSL man in the middle + * + *

Default = null + * + *

Note - This and {@link #withSslEngineSource(SslEngineSource)} are mutually exclusive. + */ + HttpProxyServerBootstrap withManInTheMiddle(MitmManager mitmManager); + + /** + * Specify a {@link HttpFiltersSource} to use for filtering requests and/or responses through this + * proxy. + * + *

Default = null + */ + HttpProxyServerBootstrap withFiltersSource(HttpFiltersSource filtersSource); + + /** + * Specify whether or not to use secure DNS lookups for outbound connections. + * + *

Default = false + */ + @CanIgnoreReturnValue + HttpProxyServerBootstrap withUseDnsSec(boolean useDnsSec); + + /** + * Specify whether or not to run this proxy as a transparent proxy. + * + *

Default = false + */ + HttpProxyServerBootstrap withTransparent(boolean transparent); + + /** + * Specify the timeout after which to disconnect idle connections, in seconds. + * + *

Default = 70 + */ + HttpProxyServerBootstrap withIdleConnectionTimeout(int idleConnectionTimeoutInSeconds); + + /** + * Specify the timeout after which to disconnect idle connections + * + *

Default = 70 seconds + */ + HttpProxyServerBootstrap withIdleConnectionTimeout(Duration idleConnectionTimeout); + + /** + * Specify the timeout for connecting to the upstream server on a new connection, in milliseconds. + * + *

Default = 40000 + */ + HttpProxyServerBootstrap withConnectTimeout(int connectTimeout); + + /** Specify a custom {@link HostResolver} for resolving server addresses. */ + HttpProxyServerBootstrap withServerResolver(HostResolver serverResolver); + + /** + * Specify a custom {@link ServerGroup} to use for managing this server's resources and such. If + * one isn't provided, a default one will be created using the {@link ThreadPoolConfiguration} + * provided + * + * @param group A custom server group + */ + HttpProxyServerBootstrap withServerGroup(ServerGroup group); + + /** Add an {@link ActivityTracker} for tracking activity in this proxy. */ + HttpProxyServerBootstrap plusActivityTracker(ActivityTracker activityTracker); + + /** + * Specify the read and/or write bandwidth throttles for this proxy server. 0 indicates not + * throttling. + */ + HttpProxyServerBootstrap withThrottling( + long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond); + + /** + * All outgoing-communication of the proxy-instance is going to be routed via the given + * network-interface + * + * @param inetSocketAddress to be used for outgoing communication + */ + @CanIgnoreReturnValue + @NullMarked + HttpProxyServerBootstrap withNetworkInterface(InetSocketAddress inetSocketAddress); + + HttpProxyServerBootstrap withMaxInitialLineLength(int maxInitialLineLength); + + HttpProxyServerBootstrap withMaxHeaderSize(int maxHeaderSize); + + HttpProxyServerBootstrap withMaxChunkSize(int maxChunkSize); + + /** + * When true, the proxy will accept requests that appear to be directed at an origin server (i.e. + * the URI in the HTTP request will contain an origin-form, rather than an absolute-form, as + * specified in RFC 7230, section 5.3). This is useful when the proxy is acting as a + * gateway/reverse proxy. Note: This feature should not be enabled when running as a + * forward proxy; doing so may cause an infinite loop if the client requests the URI of the proxy. + * + * @param allowRequestToOriginServer when true, the proxy will accept origin-form HTTP requests + */ + HttpProxyServerBootstrap withAllowRequestToOriginServer(boolean allowRequestToOriginServer); + + /** + * Sets the alias to use when adding Via headers to incoming and outgoing HTTP messages. The alias + * may be any pseudonym, or if not specified, defaults to the hostname of the local machine. See + * RFC 7230, section 5.7.1. + * + * @param alias the pseudonym to add to Via headers + */ + HttpProxyServerBootstrap withProxyAlias(String alias); + + /** + * Build and starts the server. + * + * @return the newly built and started server + */ + HttpProxyServer start(); + + /** + * Set the configuration parameters for the proxy's thread pools. + * + * @param configuration thread pool configuration + * @return proxy server bootstrap for chaining + */ + HttpProxyServerBootstrap withThreadPoolConfiguration(ThreadPoolConfiguration configuration); + + /** + * Specifies if the proxy server should accept a proxy protocol header. Once set it works with + * request that include a proxy protocol header. The proxy server reads an incoming proxy protocol + * header from the client. + * + * @param allowProxyProtocol when true, the proxy will accept a proxy protocol header + */ + HttpProxyServerBootstrap withAcceptProxyProtocol(boolean allowProxyProtocol); + + /** + * Specifies if the proxy server should send a proxy protocol header. + * + * @param sendProxyProtocol when true, the proxy will send a proxy protocol header + */ + HttpProxyServerBootstrap withSendProxyProtocol(boolean sendProxyProtocol); + + /** + * Enable or disable the shared server connection pool. + * + *

When enabled, all client connections share a common pool of server connections, allowing + * connections to the same server to be reused across different client connections. This addresses + * connection explosion when many clients connect to the same servers. + * + *

Disabled by default for backwards compatibility. + * + * @param useSharedServerConnectionPool true to enable the shared pool + */ + HttpProxyServerBootstrap withSharedServerConnectionPool(boolean useSharedServerConnectionPool); + + /** + * Selects the server connection pool implementation to use when the shared pool is enabled. + * + *

Default is {@link ServerConnectionPoolType#CONCURRENT_MAP}. + * + * @param poolType the pool implementation to use + */ + HttpProxyServerBootstrap withServerConnectionPoolType(ServerConnectionPoolType poolType); + + /** + * Sets the maximum number of connections per host:port when using the shared connection pool. + * + *

Default is 10. This allows multiple connections to the same server for high concurrency + * scenarios. + * + * @param maxConnectionsPerHost the maximum number of connections per host:port + */ + HttpProxyServerBootstrap withMaxConnectionsPerHost(int maxConnectionsPerHost); + + /** + * Sets the maximum total number of connections in the shared connection pool. + * + *

Default is 200. + * + * @param maxConnections the maximum total number of pooled connections + */ + HttpProxyServerBootstrap withMaxConnections(int maxConnections); + + /** + * Sets the idle timeout for pooled connections. Connections that remain idle (available in the + * pool) for longer than this duration will be evicted. + * + *

Default is null (no eviction based on idle time). When set, connections will be evicted + * after being idle for the specified duration. + * + * @param idleTimeout the idle timeout duration, or null to disable idle eviction + */ + HttpProxyServerBootstrap withPoolIdleTimeout(Duration idleTimeout); + + /** + * When enabled, MITM connections will use the shared server connection pool. This allows upstream + * connections for MITM'd HTTPS traffic to be reused across different client connections. + * + *

Enabled only takes effect when {@link #withSharedServerConnectionPool(boolean)} is also + * enabled. + * + *

Disabled by default for backwards compatibility. + * + * @param poolSharedMitmConnections true to allow MITM connections to use the shared pool + */ + HttpProxyServerBootstrap withPoolSharedMitmConnections(boolean poolSharedMitmConnections); + + /** + * When enabled, each HTTP request through an MITM tunnel independently acquires and releases a + * server connection from the shared pool. This allows multiple clients' MITM requests to share + * upstream connections concurrently (one request at a time per connection). + * + *

Only takes effect when {@link #withPoolSharedMitmConnections(boolean)} is also enabled. + * + *

Disabled by default for backwards compatibility. + * + * @param poolPerRequestInMitm true to enable per-request pooling inside MITM tunnels + */ + HttpProxyServerBootstrap withPoolPerRequestInMitm(boolean poolPerRequestInMitm); } diff --git a/src/main/java/org/littleshoot/proxy/Launcher.java b/src/main/java/org/littleshoot/proxy/Launcher.java index b3d7e666..db15db38 100644 --- a/src/main/java/org/littleshoot/proxy/Launcher.java +++ b/src/main/java/org/littleshoot/proxy/Launcher.java @@ -1,137 +1,474 @@ package org.littleshoot.proxy; -import org.apache.commons.cli.*; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.ACCEPTOR_THREADS; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.ALLOW_PROXY_PROTOCOL; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.ALLOW_REQUESTS_TO_ORIGIN_SERVER; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.CLIENT_TO_PROXY_WORKER_THREADS; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.PROXY_TO_SERVER_WORKER_THREADS; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.SEND_PROXY_PROTOCOL; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_ALIAS; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_PASSWORD; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_PATH; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.SSL_CLIENTS_SEND_CERTS; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.SSL_CLIENTS_TRUST_ALL_SERVERS; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.THROTTLE_READ_BYTES_PER_SECOND; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.THROTTLE_WRITE_BYTES_PER_SECOND; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.TRANSPARENT; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.bootstrap; +import static org.littleshoot.proxy.impl.DefaultHttpProxyServer.bootstrapFromFile; + +import java.io.File; +import java.net.InetSocketAddress; +import java.net.URL; +import java.util.Arrays; +import org.apache.commons.cli.CommandLine; +import org.apache.commons.cli.CommandLineParser; +import org.apache.commons.cli.DefaultParser; +import org.apache.commons.cli.HelpFormatter; +import org.apache.commons.cli.Options; +import org.apache.commons.cli.ParseException; +import org.apache.commons.cli.UnrecognizedOptionException; import org.apache.commons.lang3.StringUtils; -import org.apache.log4j.xml.DOMConfigurator; +import org.apache.logging.log4j.core.config.Configurator; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.extras.ActivityLogger; +import org.littleshoot.proxy.extras.LogFormat; import org.littleshoot.proxy.extras.SelfSignedMitmManager; -import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; import org.littleshoot.proxy.impl.ProxyUtils; +import org.littleshoot.proxy.impl.ThreadPoolConfiguration; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import java.io.File; -import java.net.InetSocketAddress; -import java.util.Arrays; - -/** - * Launches a new HTTP proxy. - */ +/** Launches a new HTTP proxy. */ public class Launcher { - private static final Logger LOG = LoggerFactory.getLogger(Launcher.class); - - private static final String OPTION_DNSSEC = "dnssec"; - - private static final String OPTION_PORT = "port"; - - private static final String OPTION_HELP = "help"; - - private static final String OPTION_MITM = "mitm"; - - private static final String OPTION_NIC = "nic"; - - /** - * Starts the proxy from the command line. - * - * @param args - * Any command line arguments. - */ - public static void main(final String... args) { - pollLog4JConfigurationFileIfAvailable(); - LOG.info("Running LittleProxy with args: {}", Arrays.asList(args)); - final Options options = new Options(); - options.addOption(null, OPTION_DNSSEC, true, - "Request and verify DNSSEC signatures."); - options.addOption(null, OPTION_PORT, true, "Run on the specified port."); - options.addOption(null, OPTION_NIC, true, "Run on a specified Nic"); - options.addOption(null, OPTION_HELP, false, - "Display command line help."); - options.addOption(null, OPTION_MITM, false, "Run as man in the middle."); - - final CommandLineParser parser = new DefaultParser(); - final CommandLine cmd; - try { - cmd = parser.parse(options, args); - if (cmd.getArgs().length > 0) { - throw new UnrecognizedOptionException( - "Extra arguments were provided in " - + Arrays.asList(args)); - } - } catch (final ParseException e) { - printHelp(options, - "Could not parse command line: " + Arrays.asList(args)); - return; - } - if (cmd.hasOption(OPTION_HELP)) { - printHelp(options, null); - return; - } - final int defaultPort = 8080; - int port; - if (cmd.hasOption(OPTION_PORT)) { - final String val = cmd.getOptionValue(OPTION_PORT); - try { - port = Integer.parseInt(val); - } catch (final NumberFormatException e) { - printHelp(options, "Unexpected port " + val); - return; - } + public static final int DEFAULT_PORT = 8080; + private static final Logger LOG = LoggerFactory.getLogger(Launcher.class); + + private static final String OPTION_DNSSEC = "dnssec"; + + private static final String OPTION_PORT = "port"; + + private static final String OPTION_HELP = "help"; + + private static final String OPTION_MITM = "mitm"; + + private static final String OPTION_NIC = "nic"; + + private static final String OPTION_CONFIG = "config"; + + private static final String OPTION_LOG_CONFIG = "log_config"; + private static final String OPTION_SERVER = "server"; + private static final String OPTION_NAME = "name"; + private static final String OPTION_ADDRESS = "address"; + private static final String OPTION_PROXY_ALIAS = "proxy_alias"; + private static final String OPTION_ALLOW_LOCAL_ONLY = "allow_local_only"; + private static final String OPTION_AUTHENTICATE_SSL_CLIENTS = "authenticate_ssl_clients"; + private static final String OPTION_SSL_CLIENTS_TRUST_ALL_SERVERS = SSL_CLIENTS_TRUST_ALL_SERVERS; + private static final String OPTION_SSL_CLIENTS_SEND_CERTS = SSL_CLIENTS_SEND_CERTS; + private static final String OPTION_SSL_CLIENTS_KEYSTORE_PATH = SSL_CLIENTS_KEYSTORE_PATH; + private static final String OPTION_SSL_CLIENTS_KEYSTORE_ALIAS = SSL_CLIENTS_KEYSTORE_ALIAS; + private static final String OPTION_SSL_CLIENTS_KEYSTORE_PASSWORD = SSL_CLIENTS_KEYSTORE_PASSWORD; + private static final String OPTION_TRANSPARENT = TRANSPARENT; + private static final String OPTION_THROTTLE_READ_BYTES_PER_SECOND = + THROTTLE_READ_BYTES_PER_SECOND; + private static final String OPTION_THROTTLE_WRITE_BYTES_PER_SECOND = + THROTTLE_WRITE_BYTES_PER_SECOND; + private static final String OPTION_ALLOW_REQUEST_TO_ORIGIN_SERVER = + ALLOW_REQUESTS_TO_ORIGIN_SERVER; + private static final String OPTION_ALLOW_PROXY_PROTOCOL = ALLOW_PROXY_PROTOCOL; + private static final String OPTION_SEND_PROXY_PROTOCOL = SEND_PROXY_PROTOCOL; + private static final String OPTION_CLIENT_TO_PROXY_WORKER_THREADS = + CLIENT_TO_PROXY_WORKER_THREADS; + private static final String OPTION_PROXY_TO_SERVER_WORKER_THREADS = + PROXY_TO_SERVER_WORKER_THREADS; + private static final String OPTION_ACCEPTOR_THREADS = ACCEPTOR_THREADS; + private static final String OPTION_ACTIVITY_LOG_FORMAT = "activity_log_format"; + public static final int DELAY_IN_SECONDS_BETWEEN_RELOAD = 15; + private static final String DEFAULT_JKS_KEYSTORE_PATH = "littleproxy_keystore.jks"; + + @Nullable private volatile HttpProxyServer httpProxyServer; + + /** + * Starts the proxy from the command line. + * + * @param args Any command line arguments. + */ + public static void main(final String... args) { + Launcher launcher = new Launcher(); + launcher.start(args); + } + + protected void start(String[] args) { + final Options options = buildOptions(); + + CommandLine cmd = parseCommandLine(args, options); + + configureLogging(cmd); + + LOG.info("Running LittleProxy with args: {}", Arrays.asList(args)); + + if (cmd.hasOption(OPTION_HELP)) { + printHelp(options, null); + return; + } + + HttpProxyServerBootstrap bootstrap; + if (cmd.hasOption(OPTION_CONFIG)) { + String proxyConfigurationPath = cmd.getOptionValue(OPTION_CONFIG); + LOG.info("Using configuration file: {}", proxyConfigurationPath); + cmd.getOptionValue(OPTION_CONFIG); + bootstrap = bootstrapFromFile(proxyConfigurationPath); + } else { + bootstrap = bootstrap(); + } + + int port; + if (cmd.hasOption(OPTION_PORT)) { + final String val = cmd.getOptionValue(OPTION_PORT); + try { + port = Integer.parseInt(val); + } catch (final NumberFormatException e) { + printHelp(options, "Unexpected port " + val); + return; + } + } else { + port = DEFAULT_PORT; + } + bootstrap.withPort(port); + LOG.info("About to start server on port: '{}'", port); + + if (cmd.hasOption(OPTION_NIC)) { + final String val = cmd.getOptionValue(OPTION_NIC); + bootstrap.withNetworkInterface(new InetSocketAddress(val, 0)); + } + + if (cmd.hasOption(OPTION_MITM)) { + LOG.info("Running as Man in the Middle"); + String keyStorePath = DEFAULT_JKS_KEYSTORE_PATH; + if (cmd.hasOption(OPTION_SSL_CLIENTS_KEYSTORE_PATH)) { + keyStorePath = cmd.getOptionValue(OPTION_SSL_CLIENTS_KEYSTORE_PATH); + } + bootstrap.withManInTheMiddle(new SelfSignedMitmManager(keyStorePath, true, true)); + } + + if (cmd.hasOption(OPTION_DNSSEC)) { + final String val = cmd.getOptionValue(OPTION_DNSSEC); + if (ProxyUtils.isTrue(val)) { + LOG.info("Using DNSSEC"); + bootstrap.withUseDnsSec(true); + } else if (ProxyUtils.isFalse(val)) { + LOG.info("Not using DNSSEC"); + bootstrap.withUseDnsSec(false); + } else { + printHelp(options, "Unexpected value for " + OPTION_DNSSEC + "=:" + val); + return; + } + } + + if (cmd.hasOption(OPTION_NAME)) { + final String val = cmd.getOptionValue(OPTION_NAME); + LOG.info("Running with name: '{}'", val); + bootstrap.withName(val); + } + + if (cmd.hasOption(OPTION_ADDRESS)) { + final String val = cmd.getOptionValue(OPTION_ADDRESS); + LOG.info("Binding to address: '{}'", val); + InetSocketAddress address = ProxyUtils.resolveSocketAddress(val); + if (address != null) { + bootstrap.withAddress(address); + } + } + + if (cmd.hasOption(OPTION_PROXY_ALIAS)) { + final String val = cmd.getOptionValue(OPTION_PROXY_ALIAS); + LOG.info("Using proxy alias: '{}'", val); + if (val != null) { + bootstrap.withProxyAlias(val); + } + } + + if (cmd.hasOption(OPTION_ALLOW_LOCAL_ONLY)) { + final String val = cmd.getOptionValue(OPTION_ALLOW_LOCAL_ONLY); + LOG.info("Setting allow local only to: '{}'", val); + if (val != null) { + bootstrap.withAllowLocalOnly(Boolean.parseBoolean(val)); + } + } + + if (cmd.hasOption(OPTION_AUTHENTICATE_SSL_CLIENTS)) { + final String val = cmd.getOptionValue(OPTION_AUTHENTICATE_SSL_CLIENTS); + LOG.info("Setting authenticate SSL clients with a selfSigned cert : '{}'", val); + if (val != null) { + boolean trustAllServers = + Boolean.parseBoolean(cmd.getOptionValue(OPTION_SSL_CLIENTS_TRUST_ALL_SERVERS, "false")); + boolean sendCerts = + Boolean.parseBoolean(cmd.getOptionValue(OPTION_SSL_CLIENTS_SEND_CERTS, "false")); + SelfSignedSslEngineSource sslEngineSource; + if (cmd.hasOption(OPTION_SSL_CLIENTS_KEYSTORE_PATH)) { + String keyStorePath = cmd.getOptionValue(OPTION_SSL_CLIENTS_KEYSTORE_PATH); + if (cmd.hasOption(OPTION_SSL_CLIENTS_KEYSTORE_PASSWORD)) { + String keyStoreAlias = cmd.getOptionValue(OPTION_SSL_CLIENTS_KEYSTORE_ALIAS, ""); + String keyStorePassword = cmd.getOptionValue(OPTION_SSL_CLIENTS_KEYSTORE_PASSWORD); + sslEngineSource = + new SelfSignedSslEngineSource( + keyStorePath, trustAllServers, sendCerts, keyStoreAlias, keyStorePassword); + } else { + sslEngineSource = + new SelfSignedSslEngineSource(keyStorePath, trustAllServers, sendCerts); + } } else { - port = defaultPort; + sslEngineSource = + new SelfSignedSslEngineSource(DEFAULT_JKS_KEYSTORE_PATH, trustAllServers, sendCerts); } + bootstrap.withSslEngineSource(sslEngineSource); + bootstrap.withAuthenticateSslClients(Boolean.parseBoolean(val)); + } + } + if (cmd.hasOption(OPTION_TRANSPARENT)) { + String optionValue = cmd.getOptionValue(OPTION_TRANSPARENT); + LOG.info("Transparent proxy enabled :'{}'", optionValue); + if (optionValue != null) { + bootstrap.withTransparent(Boolean.parseBoolean(optionValue)); + } + } + long throttlingReadBytesPerSecond = 0; + long throttlingWriteBytesPerSecond = 0; + if (cmd.hasOption(OPTION_THROTTLE_READ_BYTES_PER_SECOND)) { + throttlingReadBytesPerSecond = + Long.parseLong(cmd.getOptionValue(OPTION_THROTTLE_READ_BYTES_PER_SECOND)); + } + if (cmd.hasOption(OPTION_THROTTLE_WRITE_BYTES_PER_SECOND)) { + throttlingWriteBytesPerSecond = + Long.parseLong(cmd.getOptionValue(OPTION_THROTTLE_WRITE_BYTES_PER_SECOND)); + } + if (throttlingReadBytesPerSecond > 0 || throttlingWriteBytesPerSecond > 0) { + LOG.info( + "Throttling enabled : read {} bytes/s, write {} bytes/s", + throttlingReadBytesPerSecond, + throttlingWriteBytesPerSecond); + bootstrap.withThrottling(throttlingReadBytesPerSecond, throttlingWriteBytesPerSecond); + } - System.out.println("About to start server on port: " + port); - HttpProxyServerBootstrap bootstrap = DefaultHttpProxyServer - .bootstrapFromFile("./littleproxy.properties") - .withPort(port) - .withAllowLocalOnly(false); + if (cmd.hasOption(OPTION_ALLOW_REQUEST_TO_ORIGIN_SERVER)) { + String optionValue = cmd.getOptionValue(OPTION_ALLOW_REQUEST_TO_ORIGIN_SERVER); + LOG.info("Allow request to origin server :'{}'", optionValue); + if (optionValue != null) { + bootstrap.withAllowRequestToOriginServer(Boolean.parseBoolean(optionValue)); + } + } - if (cmd.hasOption(OPTION_NIC)) { - final String val = cmd.getOptionValue(OPTION_NIC); - bootstrap.withNetworkInterface(new InetSocketAddress(val, 0)); - } + if (cmd.hasOption(OPTION_ALLOW_PROXY_PROTOCOL)) { + String optionValue = cmd.getOptionValue(OPTION_ALLOW_PROXY_PROTOCOL); + LOG.info("Allow proxy protocol :'{}'", optionValue); + if (optionValue != null) { + bootstrap.withAcceptProxyProtocol(Boolean.parseBoolean(optionValue)); + } + } - if (cmd.hasOption(OPTION_MITM)) { - LOG.info("Running as Man in the Middle"); - bootstrap.withManInTheMiddle(new SelfSignedMitmManager()); - } - - if (cmd.hasOption(OPTION_DNSSEC)) { - final String val = cmd.getOptionValue(OPTION_DNSSEC); - if (ProxyUtils.isTrue(val)) { - LOG.info("Using DNSSEC"); - bootstrap.withUseDnsSec(true); - } else if (ProxyUtils.isFalse(val)) { - LOG.info("Not using DNSSEC"); - bootstrap.withUseDnsSec(false); - } else { - printHelp(options, "Unexpected value for " + OPTION_DNSSEC - + "=:" + val); - return; - } - } + if (cmd.hasOption(OPTION_SEND_PROXY_PROTOCOL)) { + String optionValue = cmd.getOptionValue(OPTION_SEND_PROXY_PROTOCOL); + LOG.info("Send proxy protocol header:'{}'", optionValue); + if (optionValue != null) { + bootstrap.withSendProxyProtocol(Boolean.parseBoolean(optionValue)); + } + } - System.out.println("About to start..."); - bootstrap.start(); + ThreadPoolConfiguration threadPoolConfiguration = new ThreadPoolConfiguration(); + boolean threadPoolConfigSet = + false; // Flag to track if thread pool configuration is set through command line + // options + if (cmd.hasOption(OPTION_CLIENT_TO_PROXY_WORKER_THREADS)) { + String optionValue = cmd.getOptionValue(OPTION_CLIENT_TO_PROXY_WORKER_THREADS); + LOG.info("Setting client to proxy worker threads to :'{}'", optionValue); + if (optionValue != null) { + threadPoolConfiguration.withClientToProxyWorkerThreads(Integer.parseInt(optionValue)); + threadPoolConfigSet = true; + } + } + if (cmd.hasOption(OPTION_PROXY_TO_SERVER_WORKER_THREADS)) { + String optionValue = cmd.getOptionValue(OPTION_PROXY_TO_SERVER_WORKER_THREADS); + LOG.info("Setting proxy to server worker threads to :'{}'", optionValue); + if (optionValue != null) { + threadPoolConfiguration.withProxyToServerWorkerThreads(Integer.parseInt(optionValue)); + threadPoolConfigSet = true; + } + } + if (cmd.hasOption(OPTION_ACCEPTOR_THREADS)) { + String optionValue = cmd.getOptionValue(OPTION_ACCEPTOR_THREADS); + LOG.info("Setting acceptor threads to :'{}'", optionValue); + if (optionValue != null) { + threadPoolConfiguration.withAcceptorThreads(Integer.parseInt(optionValue)); + threadPoolConfigSet = true; + } + } + if (threadPoolConfigSet) { + bootstrap.withThreadPoolConfiguration(threadPoolConfiguration); } - private static void printHelp(final Options options, - final String errorMessage) { - if (!StringUtils.isBlank(errorMessage)) { - LOG.error(errorMessage); - System.err.println(errorMessage); - } + if (cmd.hasOption(OPTION_ACTIVITY_LOG_FORMAT)) { + String format = cmd.getOptionValue(OPTION_ACTIVITY_LOG_FORMAT); + try { + LogFormat logFormat = LogFormat.valueOf(format.toUpperCase()); + bootstrap.plusActivityTracker(new ActivityLogger(logFormat)); + LOG.info("Using activity log format: {}", logFormat); + } catch (IllegalArgumentException e) { + printHelp(options, "Unknown activity log format: " + format); + return; + } + } - final HelpFormatter formatter = new HelpFormatter(); - formatter.printHelp("littleproxy", options); + LOG.info("About to start..."); + httpProxyServer = bootstrap.start(); + if (cmd.hasOption(OPTION_SERVER)) { + Runtime.getRuntime().addShutdownHook(new Thread(() -> stop())); + try { + Thread.currentThread().join(); + } catch (InterruptedException e) { + stop(); + Thread.currentThread().interrupt(); + } } + } - private static void pollLog4JConfigurationFileIfAvailable() { - File log4jConfigurationFile = new File("src/test/resources/log4j.xml"); - if (log4jConfigurationFile.exists()) { - DOMConfigurator.configureAndWatch( - log4jConfigurationFile.getAbsolutePath(), 15); - } + public boolean isRunning() { + return httpProxyServer != null; + } + + public void stop() { + HttpProxyServer server = httpProxyServer; + if (server != null) { + LOG.info("Shutting down..."); + server.stop(); + httpProxyServer = null; + LOG.info("Shut down."); } + } + + @SuppressWarnings("java:S106") + private void configureLogging(CommandLine cmd) { + if (cmd.hasOption(OPTION_LOG_CONFIG)) { + String optionValue = cmd.getOptionValue(OPTION_LOG_CONFIG); + File logConfigPath = new File(optionValue); + if (logConfigPath.exists()) { + Configurator.initialize(null, logConfigPath.getAbsolutePath()); + } + } else { + // default log4j.xml file shipped with the jar + ClassLoader classLoader = Launcher.class.getClassLoader(); + URL defaultLogConfigUrl = classLoader.getResource("littleproxy_default_log4j2.xml"); + Configurator.initialize(null, defaultLogConfigUrl.toString()); + System.out.println("using 'littleproxy_default_log4j2.xml'"); + } + } + + private @NonNull CommandLine parseCommandLine(String[] args, Options options) { + final CommandLineParser parser = new DefaultParser(); + CommandLine cmd; + try { + cmd = parser.parse(options, args); + if (cmd.getArgs().length > 0) { + throw new UnrecognizedOptionException( + "Extra arguments were provided in " + Arrays.asList(args)); + } + } catch (final ParseException e) { + printHelp(options, "Could not parse command line: " + Arrays.asList(args)); + throw new IllegalArgumentException("Could not parse command line: " + Arrays.asList(args), e); + } + return cmd; + } + + protected @NonNull Options buildOptions() { + final Options options = new Options(); + options.addOption(null, OPTION_DNSSEC, true, "Request and verify DNSSEC signatures."); + options.addOption( + null, OPTION_CONFIG, true, "Path to proxy configuration file (relative or absolute)."); + options.addOption( + null, + OPTION_LOG_CONFIG, + true, + "Path to log4j configuration file (relative to current directory or absolute)."); + options.addOption(null, OPTION_PORT, true, "Run on the specified port."); + options.addOption(null, OPTION_NIC, true, "Run on a specified Nic"); + options.addOption(null, OPTION_HELP, false, "Display command line help."); + options.addOption(null, OPTION_MITM, false, "Run as man in the middle."); + options.addOption(null, OPTION_SERVER, false, "Run proxy as a server."); + options.addOption(null, OPTION_NAME, true, "name of the proxy."); + options.addOption(null, OPTION_ADDRESS, true, "address to bind the proxy."); + options.addOption(null, OPTION_PROXY_ALIAS, true, "alias for the proxy."); + options.addOption( + null, + OPTION_ALLOW_LOCAL_ONLY, + true, + "Allow only local connections to the proxy (true|false)."); + options.addOption( + null, + OPTION_AUTHENTICATE_SSL_CLIENTS, + true, + "Whether to authenticate SSL clients (true|false)."); + options.addOption( + null, + OPTION_SSL_CLIENTS_TRUST_ALL_SERVERS, + true, + "Whether SSL clients should trust all servers (true|false)."); + options.addOption( + null, + OPTION_SSL_CLIENTS_SEND_CERTS, + true, + "Whether SSL clients should send certificates (true|false)."); + options.addOption( + null, OPTION_SSL_CLIENTS_KEYSTORE_PATH, true, "Path to keystore for SSL clients."); + options.addOption( + null, OPTION_SSL_CLIENTS_KEYSTORE_ALIAS, true, "Alias for the keystore for SSL clients."); + options.addOption( + null, + OPTION_SSL_CLIENTS_KEYSTORE_PASSWORD, + true, + "Password for the keystore for SSL clients."); + options.addOption( + null, OPTION_TRANSPARENT, true, "Whether to run in transparent mode (true|false)."); + options.addOption( + null, OPTION_THROTTLE_READ_BYTES_PER_SECOND, true, "Throttling read bytes per second."); + options.addOption( + null, OPTION_THROTTLE_WRITE_BYTES_PER_SECOND, true, "Throttling write bytes per second."); + options.addOption( + null, + OPTION_ALLOW_REQUEST_TO_ORIGIN_SERVER, + true, + "Allow requests to origin server (true|false)."); + options.addOption( + null, OPTION_ALLOW_PROXY_PROTOCOL, true, "Allow Proxy Protocol (true|false)."); + options.addOption( + null, OPTION_SEND_PROXY_PROTOCOL, true, "send Proxy Protocol header (true|false)."); + options.addOption( + null, + OPTION_CLIENT_TO_PROXY_WORKER_THREADS, + true, + "Number of client-to-proxy worker threads."); + options.addOption( + null, + OPTION_PROXY_TO_SERVER_WORKER_THREADS, + true, + "Number of proxy-to-server worker threads."); + options.addOption(null, OPTION_ACCEPTOR_THREADS, true, "Number of acceptor threads."); + options.addOption( + null, OPTION_ACTIVITY_LOG_FORMAT, true, "Activity log format: CLF, ELF, JSON, SQUID, W3C"); + return options; + } + + @SuppressWarnings("java:S106") + private void printHelp(final Options options, final String errorMessage) { + if (!StringUtils.isBlank(errorMessage)) { + LOG.error(errorMessage); + // log4j is not yet loaded at this point in some cases + System.err.println(errorMessage); + } + + final HelpFormatter formatter = new HelpFormatter(); + formatter.printHelp("littleproxy", options); + } } diff --git a/src/main/java/org/littleshoot/proxy/MitmManager.java b/src/main/java/org/littleshoot/proxy/MitmManager.java index 7ba17eb3..bb209b95 100644 --- a/src/main/java/org/littleshoot/proxy/MitmManager.java +++ b/src/main/java/org/littleshoot/proxy/MitmManager.java @@ -1,54 +1,45 @@ package org.littleshoot.proxy; import io.netty.handler.codec.http.HttpRequest; - import javax.net.ssl.SSLEngine; import javax.net.ssl.SSLSession; /** - * MITMManagers encapsulate the logic required for letting LittleProxy act as a - * man in the middle for HTTPS requests. + * MITMManagers encapsulate the logic required for letting LittleProxy act as a man in the middle + * for HTTPS requests. */ public interface MitmManager { - /** - * Creates an {@link SSLEngine} for encrypting the server connection. The SSLEngine created by this method - * may use the given peer information to send SNI information when connecting to the upstream host. - * - * @param peerHost to start a client connection to the server. - * @param peerPort to start a client connection to the server. - * - * @return an SSLEngine used to connect to an upstream server - */ - SSLEngine serverSslEngine(String peerHost, int peerPort); + /** + * Creates an {@link SSLEngine} for encrypting the server connection. The SSLEngine created by + * this method may use the given peer information to send SNI information when connecting to the + * upstream host. + * + * @param peerHost to start a client connection to the server. + * @param peerPort to start a client connection to the server. + * @return an SSLEngine used to connect to an upstream server + */ + SSLEngine serverSslEngine(String peerHost, int peerPort); - /** - * Creates an {@link SSLEngine} for encrypting the server connection. - * - * @return an SSLEngine used to connect to an upstream server - */ - SSLEngine serverSslEngine(); + /** + * Creates an {@link SSLEngine} for encrypting the server connection. + * + * @return an SSLEngine used to connect to an upstream server + */ + SSLEngine serverSslEngine(); - /** - *

- * Creates an {@link SSLEngine} for encrypting the client connection based - * on the given serverSslSession. - *

- * - *

- * The serverSslSession is provided in case this method needs to inspect the - * server's certificates or something else about the encryption on the way - * to the server. - *

- * - *

- * This is the place where one would implement impersonation of the server - * by issuing replacement certificates signed by the proxy's own - * certificate. - *

- * - * @param httpRequest the HTTP CONNECT request that is being man-in-the-middled - * @param serverSslSession the {@link SSLSession} that's been established with the server - * @return the SSLEngine used to connect to the client - */ - SSLEngine clientSslEngineFor(HttpRequest httpRequest, SSLSession serverSslSession); + /** + * Creates an {@link SSLEngine} for encrypting the client connection based on the given + * serverSslSession. + * + *

The serverSslSession is provided in case this method needs to inspect the server's + * certificates or something else about the encryption on the way to the server. + * + *

This is the place where one would implement impersonation of the server by issuing + * replacement certificates signed by the proxy's own certificate. + * + * @param httpRequest the HTTP CONNECT request that is being man-in-the-middled + * @param serverSslSession the {@link SSLSession} that's been established with the server + * @return the SSLEngine used to connect to the client + */ + SSLEngine clientSslEngineFor(HttpRequest httpRequest, SSLSession serverSslSession); } diff --git a/src/main/java/org/littleshoot/proxy/ProxyAuthenticator.java b/src/main/java/org/littleshoot/proxy/ProxyAuthenticator.java index 9f96b689..3b1bb956 100644 --- a/src/main/java/org/littleshoot/proxy/ProxyAuthenticator.java +++ b/src/main/java/org/littleshoot/proxy/ProxyAuthenticator.java @@ -1,26 +1,22 @@ package org.littleshoot.proxy; /** - * Interface for objects that can authenticate someone for using our Proxy on - * the basis of a username and password. + * Interface for objects that can authenticate someone for using our Proxy on the basis of a + * username and password. */ public interface ProxyAuthenticator { - /** - * Authenticates the user using the specified userName and password. - * - * @param userName - * The user name. - * @param password - * The password. - * @return true if the credentials are acceptable, otherwise - * false. - */ - boolean authenticate(String userName, String password); - - /** - * The realm value to be used in the request for proxy authentication - * ("Proxy-Authenticate" header). Returning null will cause the string - * "Restricted Files" to be used by default. - */ - String getRealm(); + /** + * Authenticates the user using the specified userName and password. + * + * @param userName The username. + * @param password The password. + * @return true if the credentials are acceptable, otherwise false. + */ + boolean authenticate(String userName, String password); + + /** + * The realm value to be used in the request for proxy authentication ("Proxy-Authenticate" + * header). Returning null will cause the string "Restricted Files" to be used by default. + */ + String getRealm(); } diff --git a/src/main/java/org/littleshoot/proxy/ServerConnectionPoolType.java b/src/main/java/org/littleshoot/proxy/ServerConnectionPoolType.java new file mode 100644 index 00000000..60562995 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/ServerConnectionPoolType.java @@ -0,0 +1,7 @@ +package org.littleshoot.proxy; + +/** Defines which implementation backs the shared server connection pool. */ +public enum ServerConnectionPoolType { + /** Simple ConcurrentHashMap-based pool. */ + CONCURRENT_MAP +} diff --git a/src/main/java/org/littleshoot/proxy/SslEngineSource.java b/src/main/java/org/littleshoot/proxy/SslEngineSource.java index c88ba7f0..85cd4007 100644 --- a/src/main/java/org/littleshoot/proxy/SslEngineSource.java +++ b/src/main/java/org/littleshoot/proxy/SslEngineSource.java @@ -2,29 +2,21 @@ import javax.net.ssl.SSLEngine; -/** - * Source for {@link SSLEngine}s. - */ +/** Source for {@link SSLEngine}s. */ public interface SslEngineSource { - /** - * Returns an {@link SSLEngine} to use for a server connection from - * LittleProxy to the client. - */ - SSLEngine newSslEngine(); - - /** - * Returns an {@link SSLEngine} to use for a client connection from - * LittleProxy to the upstream server. * - * - * Note: Peer information is needed to send the server_name extension in - * handshake with Server Name Indication (SNI). - * - * @param peerHost - * to start a client connection to the server. - * @param peerPort - * to start a client connection to the server. - */ - SSLEngine newSslEngine(String peerHost, int peerPort); + /** Returns an {@link SSLEngine} to use for a server connection from LittleProxy to the client. */ + SSLEngine newSslEngine(); + /** + * Returns an {@link SSLEngine} to use for a client connection from LittleProxy to the upstream + * server. * + * + *

Note: Peer information is needed to send the server_name extension in handshake with Server + * Name Indication (SNI). + * + * @param peerHost to start a client connection to the server. + * @param peerPort to start a client connection to the server. + */ + SSLEngine newSslEngine(String peerHost, int peerPort); } diff --git a/src/main/java/org/littleshoot/proxy/TransportProtocol.java b/src/main/java/org/littleshoot/proxy/TransportProtocol.java index 49d4fd05..4e09e499 100644 --- a/src/main/java/org/littleshoot/proxy/TransportProtocol.java +++ b/src/main/java/org/littleshoot/proxy/TransportProtocol.java @@ -1,10 +1,6 @@ package org.littleshoot.proxy; -/** - * Enumeration of transport protocols supported by LittleProxy. - * - * UDT support is deprecated in Netty, so it's being deprecated here, too. We'll remove it when Netty removes it. - */ +/** Enumeration of transport protocols supported by LittleProxy */ public enum TransportProtocol { - TCP, @Deprecated UDT -} \ No newline at end of file + TCP +} diff --git a/src/main/java/org/littleshoot/proxy/UnknownChainedProxyTypeException.java b/src/main/java/org/littleshoot/proxy/UnknownChainedProxyTypeException.java index d6176db0..72fb1bcc 100644 --- a/src/main/java/org/littleshoot/proxy/UnknownChainedProxyTypeException.java +++ b/src/main/java/org/littleshoot/proxy/UnknownChainedProxyTypeException.java @@ -1,13 +1,14 @@ package org.littleshoot.proxy; /** - * This exception indicates that the system was asked to use an - * {@link ChainedProxyType} that it didn't know how to handle. + * This exception indicates that the system was asked to use an {@link ChainedProxyType} that it + * didn't know how to handle. */ public class UnknownChainedProxyTypeException extends RuntimeException { - private static final long serialVersionUID = 1L; - - public UnknownChainedProxyTypeException(ChainedProxyType chainedProxyType) { - super(String.format("Unknown %s: %s", ChainedProxyType.class.getSimpleName(), chainedProxyType)); - } + private static final long serialVersionUID = 1L; + + public UnknownChainedProxyTypeException(ChainedProxyType chainedProxyType) { + super( + String.format("Unknown %s: %s", ChainedProxyType.class.getSimpleName(), chainedProxyType)); + } } diff --git a/src/main/java/org/littleshoot/proxy/UnknownTransportProtocolException.java b/src/main/java/org/littleshoot/proxy/UnknownTransportProtocolException.java index b264b818..40477fce 100644 --- a/src/main/java/org/littleshoot/proxy/UnknownTransportProtocolException.java +++ b/src/main/java/org/littleshoot/proxy/UnknownTransportProtocolException.java @@ -1,12 +1,13 @@ package org.littleshoot.proxy; /** - * This exception indicates that the system was asked to use a TransportProtocol that it didn't know how to handle. + * This exception indicates that the system was asked to use a TransportProtocol that it didn't know + * how to handle. */ public class UnknownTransportProtocolException extends RuntimeException { - private static final long serialVersionUID = 1L; + private static final long serialVersionUID = 1L; - public UnknownTransportProtocolException(TransportProtocol transportProtocol) { - super(String.format("Unknown TransportProtocol: %1$s", transportProtocol)); - } + public UnknownTransportProtocolException(TransportProtocol transportProtocol) { + super(String.format("Unknown TransportProtocol: %1$s", transportProtocol)); + } } diff --git a/src/main/java/org/littleshoot/proxy/extras/ActivityLogger.java b/src/main/java/org/littleshoot/proxy/extras/ActivityLogger.java new file mode 100644 index 00000000..f4cc933c --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/extras/ActivityLogger.java @@ -0,0 +1,300 @@ +package org.littleshoot.proxy.extras; + +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; +import java.net.InetSocketAddress; +import java.time.ZoneId; +import java.time.ZonedDateTime; +import java.time.format.DateTimeFormatter; +import java.util.Locale; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import org.littleshoot.proxy.ActivityTrackerAdapter; +import org.littleshoot.proxy.FlowContext; +import org.littleshoot.proxy.FullFlowContext; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** An {@link org.littleshoot.proxy.ActivityTracker} that logs HTTP activity. */ +public class ActivityLogger extends ActivityTrackerAdapter { + + private static final Logger LOG = LoggerFactory.getLogger(ActivityLogger.class); + private static final String DATE_FORMAT_CLF = "dd/MMM/yyyy:HH:mm:ss Z"; + public static final String UTC = "UTC"; + public static final String USER_AGENT = "User-Agent"; + public static final String ISO_8601_PATTERN = "yyyy-MM-dd'T'HH:mm:ss.SSSZ"; + + private final LogFormat logFormat; + + private static class TimedRequest { + final HttpRequest request; + final long startTime; + + TimedRequest(HttpRequest request, long startTime) { + this.request = request; + this.startTime = startTime; + } + } + + private final Map requestMap = new ConcurrentHashMap<>(); + + public ActivityLogger(LogFormat logFormat) { + this.logFormat = logFormat; + } + + @Override + public void requestReceivedFromClient(FlowContext flowContext, HttpRequest httpRequest) { + requestMap.put(flowContext, new TimedRequest(httpRequest, System.currentTimeMillis())); + } + + @Override + public void responseSentToClient(FlowContext flowContext, HttpResponse httpResponse) { + TimedRequest timedRequest = requestMap.remove(flowContext); + if (timedRequest == null) { + return; + } + + String logMessage = formatLogEntry(flowContext, timedRequest, httpResponse); + if (logMessage != null) { + log(logMessage); + } + } + + protected void log(String message) { + LOG.info(message); + } + + @Override + public void clientDisconnected(FlowContext flowContext, javax.net.ssl.SSLSession sslSession) { + requestMap.remove(flowContext); + } + + @Override + public void connectionTimedOut(FlowContext flowContext) { + requestMap.remove(flowContext); + } + + @Override + public void connectionExceptionCaught(FlowContext flowContext, Throwable cause) { + requestMap.remove(flowContext); + } + + private String formatLogEntry( + FlowContext flowContext, TimedRequest timedInfo, HttpResponse response) { + HttpRequest request = timedInfo.request; + long duration = System.currentTimeMillis() - timedInfo.startTime; + + StringBuilder sb = new StringBuilder(); + InetSocketAddress clientAddress = flowContext.getClientAddress(); + String clientIp = clientAddress != null ? clientAddress.getAddress().getHostAddress() : "-"; + ZonedDateTime now = ZonedDateTime.now(ZoneId.of(UTC)); + + switch (logFormat) { + case CLF: + // host ident authuser [date] "request" status bytes + sb.append(clientIp).append(" "); + sb.append("- "); // ident + sb.append("- "); // authuser + sb.append("[").append(format(now, DATE_FORMAT_CLF)).append("] "); + sb.append("\"") + .append(request.method()) + .append(" ") + .append(getFullUrl(request)) + .append(" ") + .append(request.protocolVersion()) + .append("\" "); + sb.append(response.status().code()).append(" "); + sb.append(getContentLength(response)); + break; + + case ELF: + // Extended Log Format (ELF) - actually NCSA Combined Log Format + // host ident authuser [date] "request" status bytes "referer" "user-agent" + sb.append(clientIp).append(" "); + sb.append("- "); // ident + sb.append("- "); // authuser + sb.append("[").append(format(now, DATE_FORMAT_CLF)).append("] "); + sb.append("\"") + .append(request.method()) + .append(" ") + .append(getFullUrl(request)) + .append(" ") + .append(request.protocolVersion()) + .append("\" "); + sb.append(response.status().code()).append(" "); + sb.append(getContentLength(response)).append(" "); + sb.append("\"").append(getHeader(request, "Referer")).append("\" "); + sb.append("\"").append(getHeader(request, USER_AGENT)).append("\""); + break; + + case W3C: + // W3C Extended Log Format (simplified default) + // date time c-ip cs-method cs-uri-stem sc-status sc-bytes + // time-taken(optional/unavailable) cs(User-Agent) + DateTimeFormatter w3cDateTimeFormatter = + DateTimeFormatter.ofPattern("yyyy-MM-dd HH:mm:ss", Locale.US); + sb.append(now.format(w3cDateTimeFormatter)).append(" "); + sb.append(clientIp).append(" "); + sb.append(request.method()).append(" "); + sb.append(getFullUrl(request)).append(" "); + sb.append(response.status().code()).append(" "); + sb.append(getContentLength(response)).append(" "); + sb.append("\"").append(getHeader(request, USER_AGENT)).append("\""); + break; + + case JSON: + sb.append("{"); + sb.append("\"timestamp\":\"").append(format(now, ISO_8601_PATTERN)).append("\","); + sb.append("\"client_ip\":\"").append(clientIp).append("\","); + sb.append("\"method\":\"").append(request.method()).append("\","); + sb.append("\"uri\":\"").append(escapeJson(getFullUrl(request))).append("\","); + sb.append("\"protocol\":\"").append(request.protocolVersion()).append("\","); + sb.append("\"status\":").append(response.status().code()).append(","); + sb.append("\"bytes\":").append(getContentLength(response)).append(","); + sb.append("\"duration\":").append(duration).append(","); + sb.append("\"user_agent\":\"") + .append(escapeJson(getHeader(request, USER_AGENT))) + .append("\""); + sb.append("}"); + break; + + case LTSV: + // Labeled Tab-Separated Values + sb.append("time:").append(format(now, ISO_8601_PATTERN)).append("\t"); + sb.append("host:").append(clientIp).append("\t"); + sb.append("method:").append(request.method()).append("\t"); + sb.append("uri:").append(getFullUrl(request)).append("\t"); + sb.append("status:").append(response.status().code()).append("\t"); + sb.append("size:").append(getContentLength(response)).append("\t"); + sb.append("duration:").append(duration).append("\t"); + sb.append("ua:").append(getHeader(request, USER_AGENT)); + break; + + case CSV: + // Comma-Separated Values: timestamp,host,method,uri,status,bytes,duration,ua + sb.append("\"").append(format(now, ISO_8601_PATTERN)).append("\","); + sb.append("\"").append(clientIp).append("\","); + sb.append("\"").append(request.method()).append("\","); + sb.append("\"").append(escapeJson(getFullUrl(request))).append("\","); + sb.append(response.status().code()).append(","); + sb.append(getContentLength(response)).append(","); + sb.append(duration).append(","); + sb.append("\"").append(escapeJson(getHeader(request, USER_AGENT))).append("\""); + break; + + case SQUID: + // time elapsed remotehost code/status bytes method URL rfc931 + // peerstatus/peerhost type + long timestamp = now.toEpochSecond(); + sb.append(timestamp / 1000).append(".").append(timestamp % 1000).append(" "); + sb.append(duration).append(" "); // elapsed + sb.append(clientIp).append(" "); + sb.append("TCP_MISS/").append(response.status().code()).append(" "); + sb.append(getContentLength(response)).append(" "); + sb.append(request.method()).append(" "); + sb.append(getFullUrl(request)).append(" "); + sb.append("- "); // rfc931 + sb.append("DIRECT/").append(getServerIp(flowContext)).append(" "); + sb.append(getContentType(response)); + break; + + case HAPROXY: + // HAProxy HTTP format approximation + // client_ip:port [date] frontend backend/server Tq Tw Tc Tr Tr_tot status bytes + // ... + // simplified: client_ip [date] method uri status bytes duration + sb.append(clientIp).append(" "); + sb.append("[").append(format(now, "dd/MMM/yyyy:HH:mm:ss.SSS")).append("] "); + sb.append("\"") + .append(request.method()) + .append(" ") + .append(getFullUrl(request)) + .append(" ") + .append(request.protocolVersion()) + .append("\" "); + sb.append(response.status().code()).append(" "); + sb.append(getContentLength(response)).append(" "); + sb.append(duration); // duration in ms + break; + } + + return sb.toString(); + } + + /** + * Reconstructs the full URL from the request. If the URI is already absolute (starts with http:// + * or https://), returns it as-is. Otherwise, prepends the Host header to create a complete URL. + * + * @param request the HTTP request + * @return the full URL + */ + protected String getFullUrl(HttpRequest request) { + String uri = request.uri(); + + // Check if URI is already absolute (contains scheme) + if (uri.startsWith("http://") || uri.startsWith("https://")) { + return uri; + } + + // For CONNECT requests, the URI is just host:port + if (request.method().name().equals("CONNECT")) { + return uri; + } + + // Get host from Host header + String host = request.headers().get("Host"); + if (host == null || host.isEmpty()) { + // Fallback: return URI as-is if no Host header + return uri; + } + + // Determine scheme (default to http) + String scheme = "http"; + + // Reconstruct full URL + if (uri.startsWith("/")) { + return scheme + "://" + host + uri; + } else { + return scheme + "://" + host + "/" + uri; + } + } + + private String format(ZonedDateTime zonedDateTime, String pattern) { + DateTimeFormatter dtf = DateTimeFormatter.ofPattern(pattern, Locale.US); + return zonedDateTime.format(dtf); + } + + private String getContentLength(HttpResponse response) { + String len = response.headers().get("Content-Length"); + return len != null ? len : "-"; + } + + private String getHeader(HttpRequest request, String headerName) { + String val = request.headers().get(headerName); + return val != null ? val : "-"; + } + + private String getContentType(HttpResponse response) { + String val = response.headers().get("Content-Type"); + return val != null ? val : "-"; + } + + private String getServerIp(FlowContext context) { + if (context instanceof FullFlowContext) { + String hostAndPort = ((FullFlowContext) context).getServerHostAndPort(); + if (hostAndPort != null) { + // Returns "host:port", we want just the host/ip usually, or stick with + // host:port? + // Squid format usually asks for remotehost or peerhost. + // We will return request host. + return hostAndPort.split(":")[0]; + } + } + return "-"; + } + + private String escapeJson(String s) { + if (s == null) return ""; + return s.replace("\"", "\\\"").replace("\\", "\\\\"); + } +} diff --git a/src/main/java/org/littleshoot/proxy/extras/HAProxyMessageEncoder.java b/src/main/java/org/littleshoot/proxy/extras/HAProxyMessageEncoder.java index a58fb928..5fe5779a 100644 --- a/src/main/java/org/littleshoot/proxy/extras/HAProxyMessageEncoder.java +++ b/src/main/java/org/littleshoot/proxy/extras/HAProxyMessageEncoder.java @@ -7,18 +7,25 @@ /** * Encodes an HAProxy proxy protocol header * - * @see Proxy Protocol Specification + * @see Proxy Protocol + * Specification */ public class HAProxyMessageEncoder extends MessageToByteEncoder { - @Override - protected void encode(ChannelHandlerContext ctx, ProxyProtocolMessage msg, ByteBuf out) { - out.writeBytes(getHaProxyMessage(msg)); - } - - private byte [] getHaProxyMessage(ProxyProtocolMessage msg) { - return String.format("%s %s %s %s %s %s\r\n", msg.getCommand(), msg.getProxiedProtocol(), msg.getSourceAddress(), msg.getDestinationAddress(), msg.getSourcePort(), - msg.getDestinationPort()).getBytes(); - } + @Override + protected void encode(ChannelHandlerContext ctx, ProxyProtocolMessage msg, ByteBuf out) { + out.writeBytes(getHaProxyMessage(msg)); + } + private byte[] getHaProxyMessage(ProxyProtocolMessage msg) { + return String.format( + "%s %s %s %s %s %s\r\n", + msg.getCommand(), + msg.getProxiedProtocol(), + msg.getSourceAddress(), + msg.getDestinationAddress(), + msg.getSourcePort(), + msg.getDestinationPort()) + .getBytes(); + } } diff --git a/src/main/java/org/littleshoot/proxy/extras/LogFormat.java b/src/main/java/org/littleshoot/proxy/extras/LogFormat.java new file mode 100644 index 00000000..0244d88b --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/extras/LogFormat.java @@ -0,0 +1,34 @@ +package org.littleshoot.proxy.extras; + +/** Enumeration of supported log formats for the {@link ActivityLogger}. */ +public enum LogFormat { + /** Common Log Format (CLF). host ident authuser [date] "request" status bytes */ + CLF, + + /** + * Extended Log Format (ELF). Similar to w3c but often customizable. We'll use a standard extended + * set. + */ + ELF, + + /** JSON Format. Structured log in JSON. */ + JSON, + + /** + * Squid Native access log format. time elapsed remotehost code/status bytes method URL rfc931 + * peerstatus/peerhost type + */ + SQUID, + + /** W3C Extended Log File Format. */ + W3C, + + /** Labeled Tab-Separated Values (LTSV). label:value\tlabel2:value2 */ + LTSV, + + /** Comma-Separated Values (CSV). "timestamp","host","method","uri",... */ + CSV, + + /** HAProxy HTTP Log Format. Includes detailed timing information. */ + HAPROXY +} diff --git a/src/main/java/org/littleshoot/proxy/extras/ProxyProtocolMessage.java b/src/main/java/org/littleshoot/proxy/extras/ProxyProtocolMessage.java index cee89ad0..227d2386 100644 --- a/src/main/java/org/littleshoot/proxy/extras/ProxyProtocolMessage.java +++ b/src/main/java/org/littleshoot/proxy/extras/ProxyProtocolMessage.java @@ -7,60 +7,66 @@ public class ProxyProtocolMessage { - private HAProxyProtocolVersion protocolVersion; - private HAProxyCommand command; - private HAProxyProxiedProtocol proxiedProtocol; - private String sourceAddress; - private String destinationAddress; - private int sourcePort; - private int destinationPort; + private final HAProxyProtocolVersion protocolVersion; + private final HAProxyCommand command; + private final HAProxyProxiedProtocol proxiedProtocol; + private final String sourceAddress; + private final String destinationAddress; + private final int sourcePort; + private final int destinationPort; - public ProxyProtocolMessage(HAProxyProtocolVersion protocolVersion, HAProxyCommand command, HAProxyProxiedProtocol proxiedProtocol, String sourceAddress, String destinationAddress - , int sourcePort, int destinationPort) { - this.protocolVersion = protocolVersion; - this.command = command; - this.proxiedProtocol = proxiedProtocol; - this.sourceAddress = sourceAddress; - this.destinationAddress = destinationAddress; - this.sourcePort = sourcePort; - this.destinationPort = destinationPort; - } + public ProxyProtocolMessage( + HAProxyProtocolVersion protocolVersion, + HAProxyCommand command, + HAProxyProxiedProtocol proxiedProtocol, + String sourceAddress, + String destinationAddress, + int sourcePort, + int destinationPort) { + this.protocolVersion = protocolVersion; + this.command = command; + this.proxiedProtocol = proxiedProtocol; + this.sourceAddress = sourceAddress; + this.destinationAddress = destinationAddress; + this.sourcePort = sourcePort; + this.destinationPort = destinationPort; + } - public ProxyProtocolMessage(HAProxyMessage haProxyMessage) { - this.protocolVersion = haProxyMessage.protocolVersion(); - this.command = haProxyMessage.command(); - this.proxiedProtocol = haProxyMessage.proxiedProtocol(); - this.sourceAddress = haProxyMessage.sourceAddress(); - this.destinationAddress = haProxyMessage.destinationAddress(); - this.sourcePort = haProxyMessage.sourcePort(); - this.destinationPort = haProxyMessage.destinationPort(); - } + public ProxyProtocolMessage(HAProxyMessage haProxyMessage) { + protocolVersion = haProxyMessage.protocolVersion(); + command = haProxyMessage.command(); + proxiedProtocol = haProxyMessage.proxiedProtocol(); + sourceAddress = haProxyMessage.sourceAddress(); + destinationAddress = haProxyMessage.destinationAddress(); + sourcePort = haProxyMessage.sourcePort(); + destinationPort = haProxyMessage.destinationPort(); + } - public HAProxyProtocolVersion getProtocolVersion() { - return protocolVersion; - } + public HAProxyProtocolVersion getProtocolVersion() { + return protocolVersion; + } - public HAProxyCommand getCommand() { - return command; - } + public HAProxyCommand getCommand() { + return command; + } - public HAProxyProxiedProtocol getProxiedProtocol() { - return proxiedProtocol; - } + public HAProxyProxiedProtocol getProxiedProtocol() { + return proxiedProtocol; + } - public String getSourceAddress() { - return sourceAddress; - } + public String getSourceAddress() { + return sourceAddress; + } - public String getDestinationAddress() { - return destinationAddress; - } + public String getDestinationAddress() { + return destinationAddress; + } - public int getSourcePort() { - return sourcePort; - } + public int getSourcePort() { + return sourcePort; + } - public int getDestinationPort() { - return destinationPort; - } + public int getDestinationPort() { + return destinationPort; + } } diff --git a/src/main/java/org/littleshoot/proxy/extras/SelfSignedMitmManager.java b/src/main/java/org/littleshoot/proxy/extras/SelfSignedMitmManager.java index 190eacaa..2f6b366b 100644 --- a/src/main/java/org/littleshoot/proxy/extras/SelfSignedMitmManager.java +++ b/src/main/java/org/littleshoot/proxy/extras/SelfSignedMitmManager.java @@ -1,37 +1,39 @@ package org.littleshoot.proxy.extras; import io.netty.handler.codec.http.HttpRequest; -import org.littleshoot.proxy.MitmManager; - import javax.net.ssl.SSLEngine; import javax.net.ssl.SSLSession; +import org.littleshoot.proxy.MitmManager; -/** - * {@link MitmManager} that uses self-signed certs for everything. - */ +/** {@link MitmManager} that uses self-signed certs for everything. */ public class SelfSignedMitmManager implements MitmManager { - private final SelfSignedSslEngineSource selfSignedSslEngineSource; - - public SelfSignedMitmManager() { - this.selfSignedSslEngineSource = new SelfSignedSslEngineSource(true); - } - - public SelfSignedMitmManager(SelfSignedSslEngineSource selfSignedSslEngineSource) { - this.selfSignedSslEngineSource = selfSignedSslEngineSource; - } - - @Override - public SSLEngine serverSslEngine(String peerHost, int peerPort) { - return selfSignedSslEngineSource.newSslEngine(peerHost, peerPort); - } - - @Override - public SSLEngine serverSslEngine() { - return selfSignedSslEngineSource.newSslEngine(); - } - - @Override - public SSLEngine clientSslEngineFor(HttpRequest httpRequest, SSLSession serverSslSession) { - return selfSignedSslEngineSource.newSslEngine(); - } + private final SelfSignedSslEngineSource selfSignedSslEngineSource; + + public SelfSignedMitmManager(String keyStorePath) { + selfSignedSslEngineSource = new SelfSignedSslEngineSource(keyStorePath, true, true); + } + + public SelfSignedMitmManager(String keyStorePath, boolean trustAllServers, boolean sendCerts) { + selfSignedSslEngineSource = + new SelfSignedSslEngineSource(keyStorePath, trustAllServers, sendCerts); + } + + public SelfSignedMitmManager(SelfSignedSslEngineSource selfSignedSslEngineSource) { + this.selfSignedSslEngineSource = selfSignedSslEngineSource; + } + + @Override + public SSLEngine serverSslEngine(String peerHost, int peerPort) { + return selfSignedSslEngineSource.newSslEngine(peerHost, peerPort); + } + + @Override + public SSLEngine serverSslEngine() { + return selfSignedSslEngineSource.newSslEngine(); + } + + @Override + public SSLEngine clientSslEngineFor(HttpRequest httpRequest, SSLSession serverSslSession) { + return selfSignedSslEngineSource.newSslEngine(); + } } diff --git a/src/main/java/org/littleshoot/proxy/extras/SelfSignedSslEngineSource.java b/src/main/java/org/littleshoot/proxy/extras/SelfSignedSslEngineSource.java index 3d435059..0efa224c 100644 --- a/src/main/java/org/littleshoot/proxy/extras/SelfSignedSslEngineSource.java +++ b/src/main/java/org/littleshoot/proxy/extras/SelfSignedSslEngineSource.java @@ -1,191 +1,224 @@ package org.littleshoot.proxy.extras; -import com.google.common.io.ByteStreams; -import org.littleshoot.proxy.SslEngineSource; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; +import static java.lang.System.nanoTime; +import static java.nio.charset.StandardCharsets.UTF_8; +import static java.util.Arrays.asList; +import static java.util.Objects.requireNonNullElse; +import static java.util.concurrent.TimeUnit.NANOSECONDS; -import javax.net.ssl.*; +import com.google.common.io.ByteStreams; import java.io.File; import java.io.IOException; import java.io.InputStream; import java.net.URL; +import java.nio.file.Path; +import java.nio.file.Paths; import java.security.GeneralSecurityException; import java.security.KeyStore; import java.security.Security; -import java.security.cert.X509Certificate; -import java.util.Arrays; +import javax.net.ssl.KeyManager; +import javax.net.ssl.KeyManagerFactory; +import javax.net.ssl.SSLContext; +import javax.net.ssl.SSLEngine; +import javax.net.ssl.TrustManager; +import javax.net.ssl.TrustManagerFactory; +import org.littleshoot.proxy.SslEngineSource; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; /** - * Basic {@link SslEngineSource} for testing. The {@link SSLContext} uses - * self-signed certificates that are generated lazily if the given key store - * file doesn't yet exist. + * Basic {@link SslEngineSource} for testing. The {@link SSLContext} uses self-signed certificates + * that are generated lazily if the given key store file doesn't yet exist. */ public class SelfSignedSslEngineSource implements SslEngineSource { - private static final Logger LOG = LoggerFactory - .getLogger(SelfSignedSslEngineSource.class); - - private static final String PROTOCOL = "TLS"; - - private final String alias; - private final String password; - private final String keyStoreFile; - private final boolean trustAllServers; - private final boolean sendCerts; - - private SSLContext sslContext; - - public SelfSignedSslEngineSource(String keyStorePath, boolean trustAllServers, boolean sendCerts, - String alias, String password) { - this.trustAllServers = trustAllServers; - this.sendCerts = sendCerts; - this.keyStoreFile = keyStorePath; - this.alias = alias; - this.password = password; - initializeSSLContext(); + private static final Logger LOG = LoggerFactory.getLogger(SelfSignedSslEngineSource.class); + + private static final String PROTOCOL = "TLS"; + + private final String alias; + private final String password; + private final String keyStoreFile; + private final boolean trustAllServers; + private final boolean sendCerts; + + private SSLContext sslContext; + + public SelfSignedSslEngineSource( + String keyStorePath, + boolean trustAllServers, + boolean sendCerts, + String alias, + String password) { + this.trustAllServers = trustAllServers; + this.sendCerts = sendCerts; + this.keyStoreFile = keyStorePath; + this.alias = alias; + this.password = password; + initializeSSLContext(); + } + + public SelfSignedSslEngineSource( + String keyStorePath, boolean trustAllServers, boolean sendCerts) { + this(keyStorePath, trustAllServers, sendCerts, "littleproxy", "Be Your Own Lantern"); + } + + public SelfSignedSslEngineSource(String keyStorePath) { + this(keyStorePath, false, true); + } + + public SelfSignedSslEngineSource(boolean trustAllServers) { + this(trustAllServers, true); + } + + public SelfSignedSslEngineSource(boolean trustAllServers, boolean sendCerts) { + this("littleproxy_keystore.jks", trustAllServers, sendCerts); + } + + public SelfSignedSslEngineSource() { + this(false); + } + + @Override + public SSLEngine newSslEngine() { + return sslContext.createSSLEngine(); + } + + @Override + public SSLEngine newSslEngine(String peerHost, int peerPort) { + return sslContext.createSSLEngine(peerHost, peerPort); + } + + public SSLContext getSslContext() { + return sslContext; + } + + private void initializeKeyStore(File keyStoreLocalFile) { + initializeKeyStore(keyStoreLocalFile, "littleproxy_cert"); + } + + private void initializeKeyStore(File keyStoreLocalFile, String certificateFileName) { + File keyStoreLocalAbsoluteFile = keyStoreLocalFile.getAbsoluteFile(); + + nativeCall( + "keytool", + "-genkey", + "-alias", + alias, + "-keysize", + "4096", + "-validity", + "36500", + "-keyalg", + "RSA", + "-dname", + "CN=littleproxy", + "-keypass", + password, + "-storepass", + password, + "-keystore", + keyStoreLocalAbsoluteFile.getPath()); + + LOG.info("Generated LittleProxy keystore in {}", keyStoreLocalAbsoluteFile); + + Path certificateFile = Paths.get(keyStoreLocalAbsoluteFile.getParent(), certificateFileName); + nativeCall( + "keytool", + "-exportcert", + "-alias", + alias, + "-keystore", + keyStoreLocalAbsoluteFile.getPath(), + "-storepass", + password, + "-file", + certificateFile.toString()); + LOG.info("Generated LittleProxy certificate in {}", certificateFile); + } + + private void initializeSSLContext() { + String algorithm = + requireNonNullElse(Security.getProperty("ssl.KeyManagerFactory.algorithm"), "SunX509"); + + try { + final KeyStore ks = loadKeyStore(); + + // Set up key manager factory to use our key store + final KeyManagerFactory kmf = KeyManagerFactory.getInstance(algorithm); + kmf.init(ks, password.toCharArray()); + + // Set up a trust manager factory to use our key store + TrustManagerFactory tmf = TrustManagerFactory.getInstance(algorithm); + tmf.init(ks); + + TrustManager[] trustManagers = createTrustManagers(tmf); + KeyManager[] keyManagers = sendCerts ? kmf.getKeyManagers() : new KeyManager[0]; + + // Initialize the SSLContext to work with our key managers. + sslContext = SSLContext.getInstance(PROTOCOL); + sslContext.init(keyManagers, trustManagers, null); + } catch (IOException | GeneralSecurityException e) { + throw new RuntimeException("Failed to initialize the server-side SSLContext", e); } - - public SelfSignedSslEngineSource(String keyStorePath, boolean trustAllServers, boolean sendCerts) { - this(keyStorePath, trustAllServers, sendCerts, "littleproxy", "Be Your Own Lantern"); - } - - public SelfSignedSslEngineSource(String keyStorePath) { - this(keyStorePath, false, true); - } - - public SelfSignedSslEngineSource(boolean trustAllServers) { - this(trustAllServers, true); - } - - public SelfSignedSslEngineSource(boolean trustAllServers, boolean sendCerts) { - this("littleproxy_keystore.jks", trustAllServers, sendCerts); + } + + private TrustManager[] createTrustManagers(TrustManagerFactory tmf) { + return trustAllServers + ? new TrustManager[] {new TrustingTrustManager()} + : tmf.getTrustManagers(); + } + + private KeyStore loadKeyStore() throws IOException, GeneralSecurityException { + URL resourceUrl = getClass().getResource(keyStoreFile); + if (resourceUrl != null) { + return loadKeyStore(resourceUrl); + } else { + File keyStoreLocalFile = new File(keyStoreFile); + if (!keyStoreLocalFile.isFile()) { + initializeKeyStore(keyStoreLocalFile); + } + return loadKeyStore(keyStoreLocalFile.toURI().toURL()); } + } - public SelfSignedSslEngineSource() { - this(false); + private KeyStore loadKeyStore(URL url) throws IOException, GeneralSecurityException { + KeyStore keyStore = KeyStore.getInstance("JKS"); + try (InputStream is = url.openStream()) { + keyStore.load(is, password.toCharArray()); } - - @Override - public SSLEngine newSslEngine() { - return sslContext.createSSLEngine(); - } - - @Override - public SSLEngine newSslEngine(String peerHost, int peerPort) { - return sslContext.createSSLEngine(peerHost, peerPort); - } - - public SSLContext getSslContext() { - return sslContext; - } - - private void initializeKeyStore(String filename) { - nativeCall("keytool", "-genkey", "-alias", alias, "-keysize", - "4096", "-validity", "36500", "-keyalg", "RSA", "-dname", - "CN=littleproxy", "-keypass", password, "-storepass", - password, "-keystore", filename); - - nativeCall("keytool", "-exportcert", "-alias", alias, "-keystore", - filename, "-storepass", password, "-file", - "littleproxy_cert"); - } - - private void initializeSSLContext() { - String algorithm = Security - .getProperty("ssl.KeyManagerFactory.algorithm"); - if (algorithm == null) { - algorithm = "SunX509"; - } - - try { - final KeyStore ks = loadKeyStore(); - - // Set up key manager factory to use our key store - final KeyManagerFactory kmf = - KeyManagerFactory.getInstance(algorithm); - kmf.init(ks, password.toCharArray()); - - // Set up a trust manager factory to use our key store - TrustManagerFactory tmf = TrustManagerFactory - .getInstance(algorithm); - tmf.init(ks); - - TrustManager[] trustManagers; - if (!trustAllServers) { - trustManagers = tmf.getTrustManagers(); - } else { - trustManagers = new TrustManager[] { new X509TrustManager() { - // TrustManager that trusts all servers - @Override - public void checkClientTrusted(X509Certificate[] arg0, String arg1) { - } - - @Override - public void checkServerTrusted(X509Certificate[] arg0, String arg1) { - } - - @Override - public X509Certificate[] getAcceptedIssuers() { - return null; - } - } }; - } - - KeyManager[] keyManagers; - if (sendCerts) { - keyManagers = kmf.getKeyManagers(); - } else { - keyManagers = new KeyManager[0]; - } - - // Initialize the SSLContext to work with our key managers. - sslContext = SSLContext.getInstance(PROTOCOL); - sslContext.init(keyManagers, trustManagers, null); - } catch (final Exception e) { - throw new Error( - "Failed to initialize the server-side SSLContext", e); - } - } - - private KeyStore loadKeyStore() throws IOException, GeneralSecurityException { - final KeyStore keyStore = KeyStore.getInstance("JKS"); - URL resourceUrl = getClass().getResource(keyStoreFile); - if(resourceUrl != null) { - loadKeyStore(keyStore, resourceUrl); - } else { - File keyStoreLocalFile = new File(keyStoreFile); - if(!keyStoreLocalFile.isFile()) { - initializeKeyStore(keyStoreLocalFile.getName()); - } - loadKeyStore(keyStore, keyStoreLocalFile.toURI().toURL()); - } - return keyStore; - } - - private void loadKeyStore(KeyStore keyStore, URL url) throws IOException, GeneralSecurityException { - try(InputStream is = url.openStream()) { - keyStore.load(is, password.toCharArray()); - } - } - - private String nativeCall(final String... commands) { - LOG.info("Running '{}'", Arrays.asList(commands)); - final ProcessBuilder pb = new ProcessBuilder(commands); - try { - final Process process = pb.start(); - byte[] data; - try (InputStream is = process.getInputStream()) { - data = ByteStreams.toByteArray(is); - } - String dataAsString = new String(data); - - LOG.info("Completed native call: '{}'\nResponse: '" + dataAsString + "'", - Arrays.asList(commands)); - return dataAsString; - } catch (final IOException e) { - LOG.error("Error running commands: " + Arrays.asList(commands), e); - return ""; - } + LOG.debug("Loaded LittleProxy keystore from {}", url); + return keyStore; + } + + private void nativeCall(final String... commands) { + long start = nanoTime(); + LOG.info("Running '{}'", asList(commands)); + final ProcessBuilder pb = new ProcessBuilder(commands); + // Merge stderr into stdout so we only need to read one stream + pb.redirectErrorStream(true); + try { + final Process process = pb.start(); + byte[] data; + try (InputStream is = process.getInputStream()) { + data = ByteStreams.toByteArray(is); + } + int exitCode = process.waitFor(); + String dataAsString = new String(data, UTF_8); + LOG.info( + "Completed native call '{}' in {} ms (exit: {})\nResponse: '{}'", + asList(commands), + duration(start), + exitCode, + dataAsString); + } catch (IOException e) { + LOG.error("Error running commands {} after {} ms", asList(commands), duration(start), e); + } catch (InterruptedException e) { + LOG.error("Error running commands {} after {} ms", asList(commands), duration(start), e); + Thread.currentThread().interrupt(); } + } + private long duration(long start) { + return NANOSECONDS.toMillis(nanoTime() - start); + } } diff --git a/src/main/java/org/littleshoot/proxy/extras/TrustingTrustManager.java b/src/main/java/org/littleshoot/proxy/extras/TrustingTrustManager.java new file mode 100644 index 00000000..e7ff656c --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/extras/TrustingTrustManager.java @@ -0,0 +1,18 @@ +package org.littleshoot.proxy.extras; + +import java.security.cert.X509Certificate; +import javax.net.ssl.X509TrustManager; + +/** TrustManager that trusts all servers */ +class TrustingTrustManager implements X509TrustManager { + @Override + public void checkClientTrusted(X509Certificate[] arg0, String arg1) {} + + @Override + public void checkServerTrusted(X509Certificate[] arg0, String arg1) {} + + @Override + public X509Certificate[] getAcceptedIssuers() { + return null; + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/CategorizedThreadFactory.java b/src/main/java/org/littleshoot/proxy/impl/CategorizedThreadFactory.java index 54c07ab7..70ab09c8 100644 --- a/src/main/java/org/littleshoot/proxy/impl/CategorizedThreadFactory.java +++ b/src/main/java/org/littleshoot/proxy/impl/CategorizedThreadFactory.java @@ -1,47 +1,56 @@ package org.littleshoot.proxy.impl; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - import java.util.concurrent.ThreadFactory; import java.util.concurrent.atomic.AtomicInteger; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; -/** - * A ThreadFactory that adds LittleProxy-specific information to the threads' names. - */ +/** A ThreadFactory that adds LittleProxy-specific information to the threads' names. */ public class CategorizedThreadFactory implements ThreadFactory { - private static final Logger log = LoggerFactory.getLogger(CategorizedThreadFactory.class); - - private final String name; - private final String category; - private final int uniqueServerGroupId; - - private AtomicInteger threadCount = new AtomicInteger(0); - - /** - * Exception handler for proxy threads. Logs the name of the thread and the exception that was caught. - */ - private static final Thread.UncaughtExceptionHandler UNCAUGHT_EXCEPTION_HANDLER = (t, e) -> log.error("Uncaught throwable in thread: {}", t.getName(), e); - - - /** - * @param name the user-supplied name of this proxy - * @param category the type of threads this factory is creating (acceptor, client-to-proxy worker, proxy-to-server worker) - * @param uniqueServerGroupId a unique number for the server group creating this thread factory, to differentiate multiple proxy instances with the same name - */ - public CategorizedThreadFactory(String name, String category, int uniqueServerGroupId) { - this.category = category; - this.name = name; - this.uniqueServerGroupId = uniqueServerGroupId; - } - - @Override - public Thread newThread(Runnable r) { - Thread t = new Thread(r, name + "-" + uniqueServerGroupId + "-" + category + "-" + threadCount.getAndIncrement()); - - t.setUncaughtExceptionHandler(UNCAUGHT_EXCEPTION_HANDLER); - - return t; - } - + private static final Logger log = LoggerFactory.getLogger(CategorizedThreadFactory.class); + + private final String name; + private final String category; + private final int uniqueServerGroupId; + + private final AtomicInteger threadCount = new AtomicInteger(0); + + /** + * Exception handler for proxy threads. Logs the name of the thread and the exception that was + * caught. + */ + private static final Thread.UncaughtExceptionHandler UNCAUGHT_EXCEPTION_HANDLER = + (t, e) -> log.error("Uncaught throwable in thread: {}", t.getName(), e); + + /** + * @param name the user-supplied name of this proxy + * @param category the type of threads this factory is creating (acceptor, client-to-proxy worker, + * proxy-to-server worker) + * @param uniqueServerGroupId a unique number for the server group creating this thread factory, + * to differentiate multiple proxy instances with the same name + */ + public CategorizedThreadFactory(String name, String category, int uniqueServerGroupId) { + this.category = category; + this.name = name; + this.uniqueServerGroupId = uniqueServerGroupId; + } + + @Override + public Thread newThread(Runnable r) { + Thread t = + new Thread( + r, + name + + "-" + + uniqueServerGroupId + + "-" + + category + + "-" + + threadCount.getAndIncrement()); + + t.setDaemon(true); + t.setUncaughtExceptionHandler(UNCAUGHT_EXCEPTION_HANDLER); + + return t; + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ClientDetails.java b/src/main/java/org/littleshoot/proxy/impl/ClientDetails.java index aa47c276..dda9aac8 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ClientDetails.java +++ b/src/main/java/org/littleshoot/proxy/impl/ClientDetails.java @@ -2,34 +2,28 @@ import java.net.InetSocketAddress; -/** - * Contains information about the client. - */ +/** Contains information about the client. */ public class ClientDetails { - /** - * The user name that was used for authentication, or null if authentication wasn't performed. - */ - private volatile String userName; + /** The username that was used for authentication, or null if authentication wasn't performed. */ + private volatile String userName; - /** - * The client's address - */ - private volatile InetSocketAddress clientAddress; + /** The client's address */ + private volatile InetSocketAddress clientAddress; - public String getUserName() { - return userName; - } + public String getUserName() { + return userName; + } - void setUserName(String userName) { - this.userName = userName; - } + void setUserName(String userName) { + this.userName = userName; + } - public InetSocketAddress getClientAddress() { - return clientAddress; - } + public InetSocketAddress getClientAddress() { + return clientAddress; + } - void setClientAddress(InetSocketAddress clientAddress) { - this.clientAddress = clientAddress; - } + void setClientAddress(InetSocketAddress clientAddress) { + this.clientAddress = clientAddress; + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ClientToProxyConnection.java b/src/main/java/org/littleshoot/proxy/impl/ClientToProxyConnection.java index 6397e7e9..00a6d29e 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ClientToProxyConnection.java +++ b/src/main/java/org/littleshoot/proxy/impl/ClientToProxyConnection.java @@ -1,1453 +1,1867 @@ package org.littleshoot.proxy.impl; -import com.google.common.io.BaseEncoding; +import static java.nio.charset.StandardCharsets.UTF_8; +import static java.time.format.DateTimeFormatter.ofPattern; +import static java.util.Objects.requireNonNull; +import static java.util.Objects.requireNonNullElse; +import static java.util.Optional.ofNullable; +import static org.littleshoot.proxy.HttpFiltersAdapter.NOOP_FILTER; +import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_CHUNK; +import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_INITIAL; +import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_PROXY_AUTHENTICATION; +import static org.littleshoot.proxy.impl.ConnectionState.DISCONNECT_REQUESTED; +import static org.littleshoot.proxy.impl.ConnectionState.NEGOTIATING_CONNECT; + +import com.google.errorprone.annotations.CheckReturnValue; import io.netty.buffer.ByteBuf; import io.netty.buffer.Unpooled; import io.netty.channel.Channel; import io.netty.channel.ChannelPipeline; import io.netty.handler.codec.haproxy.HAProxyMessage; import io.netty.handler.codec.haproxy.HAProxyMessageDecoder; -import io.netty.handler.codec.http.*; +import io.netty.handler.codec.http.DefaultHttpRequest; +import io.netty.handler.codec.http.FullHttpRequest; +import io.netty.handler.codec.http.FullHttpResponse; +import io.netty.handler.codec.http.HttpContent; +import io.netty.handler.codec.http.HttpHeaderNames; +import io.netty.handler.codec.http.HttpHeaderValues; +import io.netty.handler.codec.http.HttpHeaders; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpObject; +import io.netty.handler.codec.http.HttpObjectAggregator; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpRequestDecoder; +import io.netty.handler.codec.http.HttpResponse; +import io.netty.handler.codec.http.HttpResponseEncoder; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpUtil; +import io.netty.handler.codec.http.HttpVersion; import io.netty.handler.timeout.IdleStateHandler; import io.netty.handler.traffic.GlobalTrafficShapingHandler; -import io.netty.util.concurrent.Future; import io.netty.util.ReferenceCounted; -import org.apache.commons.lang3.StringUtils; -import org.littleshoot.proxy.*; - -import javax.net.ssl.SSLSession; +import io.netty.util.concurrent.Future; import java.io.IOException; import java.net.InetSocketAddress; import java.net.UnknownHostException; -import java.nio.charset.Charset; -import java.util.*; +import java.time.LocalDateTime; +import java.time.ZoneId; +import java.util.Arrays; +import java.util.Base64; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import java.util.Queue; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.RejectedExecutionException; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; import java.util.regex.Pattern; - -import static org.littleshoot.proxy.impl.ConnectionState.*; +import javax.net.ssl.SSLEngine; +import javax.net.ssl.SSLSession; +import org.apache.commons.lang3.StringUtils; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.ActivityTracker; +import org.littleshoot.proxy.ChainedProxy; +import org.littleshoot.proxy.ChainedProxyManager; +import org.littleshoot.proxy.ChainedProxyType; +import org.littleshoot.proxy.FlowContext; +import org.littleshoot.proxy.FullFlowContext; +import org.littleshoot.proxy.HttpFilters; +import org.littleshoot.proxy.ProxyAuthenticator; +import org.littleshoot.proxy.SslEngineSource; /** - *

- * Represents a connection from a client to our proxy. Each - * ClientToProxyConnection can have multiple {@link ProxyToServerConnection}s, - * at most one per outbound host:port. - *

- * - *

- * Once a ProxyToServerConnection has been created for a given server, it is - * continually reused. The ProxyToServerConnection goes through its own - * lifecycle of connects and disconnects, with different underlying - * {@link Channel}s, but only a single ProxyToServerConnection object is used - * per server. The one exception to this is CONNECT tunneling - if a connection - * has been used for CONNECT tunneling, that connection will never be reused. - *

- * - *

- * As the ProxyToServerConnections receive responses from their servers, they - * feed these back to the client by calling - * {@link #respond(ProxyToServerConnection, HttpFilters, HttpRequest, HttpResponse, HttpObject)} - * . - *

+ * Represents a connection from a client to our proxy. Each ClientToProxyConnection can have + * multiple {@link ProxyToServerConnection}s, at most one per outbound host:port. + * + *

Once a ProxyToServerConnection has been created for a given server, it is continually reused. + * The ProxyToServerConnection goes through its own lifecycle of connects and disconnects, with + * different underlying {@link Channel}s, but only a single ProxyToServerConnection object is used + * per server. The one exception to this is CONNECT tunneling - if a connection has been used for + * CONNECT tunneling, that connection will never be reused. + * + *

As the ProxyToServerConnections receive responses from their servers, they feed these back to + * the client by calling {@link #respond(ProxyToServerConnection, HttpFilters, HttpRequest, + * HttpResponse, HttpObject)} . */ +@NullMarked public class ClientToProxyConnection extends ProxyConnection { - private static final HttpResponseStatus CONNECTION_ESTABLISHED = new HttpResponseStatus( - 200, "Connection established"); + private static final HttpResponseStatus CONNECTION_ESTABLISHED = + new HttpResponseStatus(200, "Connection established"); - /** - * Used for case-insensitive comparisons when checking direct proxy request. - */ - private static final Pattern HTTP_SCHEME = Pattern.compile("^http://.*", Pattern.CASE_INSENSITIVE); + // Pipeline handler names: + private static final String HTTP_ENCODER_NAME = "encoder"; + private static final String HTTP_DECODER_NAME = "decoder"; + private static final String HTTP_PROXY_DECODER_NAME = "proxy-protocol-decoder"; + private static final String HTTP_REQUEST_READ_MONITOR_NAME = "requestReadMonitor"; + private static final String HTTP_RESPONSE_WRITTEN_MONITOR_NAME = "responseWrittenMonitor"; + private static final String MAIN_HANDLER_NAME = "handler"; - /** - * Keep track of all ProxyToServerConnections by host+port. - */ - private final Map serverConnectionsByHostAndPort = new ConcurrentHashMap<>(); + /** Used for case-insensitive comparisons when checking direct proxy request. */ + private static final Pattern ABSOLUTE_URI_PATTERN = + Pattern.compile("^(http|ws)://.*", Pattern.CASE_INSENSITIVE); - /** - * Keep track of how many servers are currently in the process of - * connecting. - */ - private final AtomicInteger numberOfCurrentlyConnectingServers = new AtomicInteger( - 0); + /** Keep track of all ProxyToServerConnections by host+port. */ + private final Map serverConnectionsByHostAndPort = + new ConcurrentHashMap<>(); - /** - * Keep track of proxy protocol header - */ - private HAProxyMessage haProxyMessage = null; + /** Keep track of how many servers are currently in the process of connecting. */ + private final AtomicInteger numberOfCurrentlyConnectingServers = new AtomicInteger(0); - /** - * Keep track of how many servers are currently connected. - */ - private final AtomicInteger numberOfCurrentlyConnectedServers = new AtomicInteger( - 0); + /** Keep track of proxy protocol header */ + @Nullable private volatile HAProxyMessage haProxyMessage; - /** - * Keep track of how many times we were able to reuse a connection. - */ - private final AtomicInteger numberOfReusedServerConnections = new AtomicInteger( - 0); + /** Keep track of how many servers are currently connected. */ + private final AtomicInteger numberOfCurrentlyConnectedServers = new AtomicInteger(0); - /** - * This is the current server connection that we're using while transferring - * chunked data. - */ - private volatile ProxyToServerConnection currentServerConnection; + /** Keep track of how many times we were able to reuse a connection. */ + private final AtomicInteger numberOfReusedServerConnections = new AtomicInteger(0); - /** - * The current filters to apply to incoming requests/chunks. - */ - private volatile HttpFilters currentFilters = HttpFiltersAdapter.NOOP_FILTER; + /** This is the current server connection that we're using while transferring chunked data. */ + @Nullable private volatile ProxyToServerConnection currentServerConnection; - private volatile SSLSession clientSslSession; + private final Map serverFlowContexts = + new ConcurrentHashMap<>(); - /** - * Tracks whether or not this ClientToProxyConnection is current doing MITM. - */ - private volatile boolean mitming = false; + /** The current filters to apply to incoming requests/chunks. */ + private volatile HttpFilters currentFilters = NOOP_FILTER; - private AtomicBoolean authenticated = new AtomicBoolean(); + @Nullable private volatile SSLSession clientSslSession; - private final GlobalTrafficShapingHandler globalTrafficShapingHandler; + /** Tracks whether this ClientToProxyConnection is current doing MITM. */ + private volatile boolean mitming; - /** - * The current HTTP request that this connection is currently servicing. - */ - private volatile HttpRequest currentRequest; - - private final ClientDetails clientDetails = new ClientDetails(); - - ClientToProxyConnection( - final DefaultHttpProxyServer proxyServer, - SslEngineSource sslEngineSource, - boolean authenticateClients, - ChannelPipeline pipeline, - GlobalTrafficShapingHandler globalTrafficShapingHandler) { - super(AWAITING_INITIAL, proxyServer, false); - - initChannelPipeline(pipeline); - - if (sslEngineSource != null) { - LOG.debug("Enabling encryption of traffic from client to proxy"); - encrypt(pipeline, sslEngineSource.newSslEngine(), - authenticateClients) - .addListener( - future -> { - if (future.isSuccess()) { - clientSslSession = sslEngine.getSession(); - recordClientSSLHandshakeSucceeded(); - } - }); - } - this.globalTrafficShapingHandler = globalTrafficShapingHandler; + private final AtomicBoolean authenticated = new AtomicBoolean(); - LOG.debug("Created ClientToProxyConnection"); - } + /** Ensures {@link #recordClientConnected()} fires at most once per connection. */ + private static final boolean CLIENT_CONNECTED_NOT_YET_RECORDED = false; - @Override - protected void readHAProxyMessage(HAProxyMessage msg) { - haProxyMessage = msg; - } + private static final boolean CLIENT_CONNECTED_RECORDED = true; + private final AtomicBoolean clientConnectedRecorded = new AtomicBoolean(); - /* ************************************************************************* - * Reading - **************************************************************************/ + private final GlobalTrafficShapingHandler globalTrafficShapingHandler; - @Override - protected ConnectionState readHTTPInitial(HttpRequest httpRequest) { - LOG.debug("Received raw request: {}", httpRequest); + /** Cached FlowContext for consistent timing data across client lifecycle events. */ + private final FlowContext clientFlowContext; - // if we cannot parse the request, immediately return a 400 and close the connection, since we do not know what state - // the client thinks the connection is in - if (httpRequest.decoderResult().isFailure()) { - LOG.debug("Could not parse request from client. Decoder result: {}", httpRequest.decoderResult().toString()); + /** The current HTTP request that this connection is currently servicing. */ + @Nullable private volatile HttpRequest currentRequest; - FullHttpResponse response = ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, - HttpResponseStatus.BAD_REQUEST, - "Unable to parse HTTP request"); - HttpUtil.setKeepAlive(response, false); + private final ClientDetails clientDetails = new ClientDetails(); - respondWithShortCircuitResponse(response); + ClientToProxyConnection( + final DefaultHttpProxyServer proxyServer, + @Nullable SslEngineSource sslEngineSource, + boolean authenticateClients, + ChannelPipeline pipeline, + GlobalTrafficShapingHandler globalTrafficShapingHandler) { + super(AWAITING_INITIAL, proxyServer, false); + this.clientFlowContext = new FlowContext(this); - return DISCONNECT_REQUESTED; - } + initChannelPipeline(pipeline, sslEngineSource, authenticateClients); - boolean authenticationRequired = authenticationRequired(httpRequest); + this.globalTrafficShapingHandler = globalTrafficShapingHandler; - if (authenticationRequired) { - LOG.debug("Not authenticated!!"); - return AWAITING_PROXY_AUTHENTICATION; - } else { - return doReadHTTPInitial(httpRequest); - } - } + LOG.debug("Created ClientToProxyConnection"); + } - /** - *

- * Reads an {@link HttpRequest}. - *

- * - *

- * If we don't yet have a {@link ProxyToServerConnection} for the desired - * server, this takes care of creating it. - *

- * - *

- * Note - the "server" could be a chained proxy, not the final endpoint for - * the request. - *

- */ - private ConnectionState doReadHTTPInitial(HttpRequest httpRequest) { - // Make a copy of the original request - this.currentRequest = copy(httpRequest); - - // Set up our filters based on the original request. If the HttpFiltersSource returns null (meaning the request/response - // should not be filtered), fall back to the default no-op filter source. - HttpFilters filterInstance = proxyServer.getFiltersSource().filterRequest(currentRequest, ctx); - if (filterInstance != null) { - currentFilters = filterInstance; - } else { - currentFilters = HttpFiltersAdapter.NOOP_FILTER; - } + @Override + protected void readHAProxyMessage(HAProxyMessage msg) { + haProxyMessage = msg; + // PROXY header available: fire the deferred clientConnected now (guarded). + recordClientConnected(); + } - // Send the request through the clientToProxyRequest filter, and respond with the short-circuit response if required - HttpResponse clientToProxyFilterResponse = currentFilters.clientToProxyRequest(httpRequest); + /* ************************************************************************* + * Reading + **************************************************************************/ - if (clientToProxyFilterResponse != null) { - LOG.debug("Responding to client with short-circuit response from filter: {}", clientToProxyFilterResponse); + @Override + ConnectionState readHTTPInitial(HttpRequest httpRequest) { + LOG.debug("Received raw request: {}", httpRequest); - boolean keepAlive = respondWithShortCircuitResponse(clientToProxyFilterResponse); - if (keepAlive) { - return AWAITING_INITIAL; - } else { - return DISCONNECT_REQUESTED; - } - } + // Earliest point to report connected for connections with no PROXY header (guarded; no-op if + // already fired). + recordClientConnected(); - // if origin-form requests are not explicitly enabled, short-circuit requests that treat the proxy as the - // origin server, to avoid infinite loops - if (!proxyServer.isAllowRequestsToOriginServer() && isRequestToOriginServer(httpRequest)) { - boolean keepAlive = writeBadRequest(httpRequest); - if (keepAlive) { - return AWAITING_INITIAL; - } else { - return DISCONNECT_REQUESTED; - } - } + // if we cannot parse the request, immediately return a 400 and close the connection, since we + // do not know what state + // the client thinks the connection is in + if (httpRequest.decoderResult().isFailure()) { + LOG.debug( + "Could not parse request from client. Decoder result: {}", + httpRequest.decoderResult().toString()); - // Identify our server and chained proxy - String serverHostAndPort = identifyHostAndPort(httpRequest); - - LOG.debug("Ensuring that hostAndPort are available in {}", - httpRequest.uri()); - if (serverHostAndPort == null || StringUtils.isBlank(serverHostAndPort)) { - LOG.warn("No host and port found in {}", httpRequest.uri()); - boolean keepAlive = writeBadGateway(httpRequest); - if (keepAlive) { - return AWAITING_INITIAL; - } else { - return DISCONNECT_REQUESTED; - } - } + FullHttpResponse response = + ProxyUtils.createFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.BAD_REQUEST, "Unable to parse HTTP request"); + HttpUtil.setKeepAlive(response, false); - LOG.debug("Finding ProxyToServerConnection for: {}", serverHostAndPort); - currentServerConnection = isMitming() || isTunneling() ? - this.currentServerConnection - : this.serverConnectionsByHostAndPort.get(serverHostAndPort); - - boolean newConnectionRequired = false; - if (ProxyUtils.isCONNECT(httpRequest)) { - LOG.debug( - "Not reusing existing ProxyToServerConnection because request is a CONNECT for: {}", - serverHostAndPort); - newConnectionRequired = true; - } else if (currentServerConnection == null) { - LOG.debug("Didn't find existing ProxyToServerConnection for: {}", - serverHostAndPort); - newConnectionRequired = true; - } + respondWithShortCircuitResponse(response); - if (newConnectionRequired) { - try { - currentServerConnection = ProxyToServerConnection.create( - proxyServer, - this, - serverHostAndPort, - currentFilters, - httpRequest, - globalTrafficShapingHandler); - if (currentServerConnection == null) { - LOG.debug("Unable to create server connection, probably no chained proxies available"); - boolean keepAlive = writeBadGateway(httpRequest); - resumeReading(); - if (keepAlive) { - return AWAITING_INITIAL; - } else { - return DISCONNECT_REQUESTED; - } - } - // Remember the connection for later - serverConnectionsByHostAndPort.put(serverHostAndPort, - currentServerConnection); - } catch (UnknownHostException uhe) { - LOG.info("Bad Host {}", httpRequest.uri()); - boolean keepAlive = writeBadGateway(httpRequest); - resumeReading(); - if (keepAlive) { - return AWAITING_INITIAL; - } else { - return DISCONNECT_REQUESTED; - } - } - } else { - LOG.debug("Reusing existing server connection: {}", - currentServerConnection); - numberOfReusedServerConnections.incrementAndGet(); - } + return DISCONNECT_REQUESTED; + } - modifyRequestHeadersToReflectProxying(httpRequest); + boolean authenticationRequired = authenticationRequired(httpRequest); + + if (authenticationRequired) { + LOG.debug("Not authenticated!!"); + return AWAITING_PROXY_AUTHENTICATION; + } else { + return doReadHTTPInitial(httpRequest); + } + } + + /** + * Reads an {@link HttpRequest}. + * + *

If we don't yet have a {@link ProxyToServerConnection} for the desired server, this takes + * care of creating it. + * + *

Note - the "server" could be a chained proxy, not the final endpoint for the request. + */ + private ConnectionState doReadHTTPInitial(HttpRequest httpRequest) { + resetCurrentRequest(); + // Make a copy of the original request + currentRequest = copy(httpRequest); + + // Set up our filters based on the original request. If the HttpFiltersSource returns null + // (meaning the request/response + // should not be filtered), fall back to the default no-op filter source. + HttpFilters filterInstance = + proxyServer.getFiltersSource().filterRequest(requireNonNull(currentRequest), ctx); + currentFilters = requireNonNullElse(filterInstance, NOOP_FILTER); + + // Send the request through the clientToProxyRequest filter, and respond with the short-circuit + // response if required + HttpResponse clientToProxyFilterResponse = currentFilters.clientToProxyRequest(httpRequest); + + if (clientToProxyFilterResponse != null) { + LOG.debug( + "Responding to client with short-circuit response from filter: {}", + clientToProxyFilterResponse); + + boolean keepAlive = respondWithShortCircuitResponse(clientToProxyFilterResponse); + if (keepAlive) { + return AWAITING_INITIAL; + } else { + return DISCONNECT_REQUESTED; + } + } - HttpResponse proxyToServerFilterResponse = currentFilters.proxyToServerRequest(httpRequest); - if (proxyToServerFilterResponse != null) { - LOG.debug("Responding to client with short-circuit response from filter: {}", proxyToServerFilterResponse); + // if origin-form requests are not explicitly enabled, short-circuit requests that treat the + // proxy as the + // origin server, to avoid infinite loops + if (!proxyServer.isAllowRequestsToOriginServer() && isRequestToOriginServer(httpRequest)) { + boolean keepAlive = writeBadRequest(httpRequest); + if (keepAlive) { + return AWAITING_INITIAL; + } else { + return DISCONNECT_REQUESTED; + } + } - boolean keepAlive = respondWithShortCircuitResponse(proxyToServerFilterResponse); - if (keepAlive) { - return AWAITING_INITIAL; - } else { - return DISCONNECT_REQUESTED; - } - } + // Identify our server and chained proxy + String serverHostAndPort = identifyHostAndPort(httpRequest); - LOG.debug("Writing request to ProxyToServerConnection"); - currentServerConnection.write(httpRequest, currentFilters); + LOG.debug("Ensuring that hostAndPort are available in {}", httpRequest.uri()); + if (StringUtils.isBlank(serverHostAndPort)) { + LOG.warn("No host and port found in {}", httpRequest.uri()); + boolean keepAlive = writeBadGateway(httpRequest); + return keepAlive ? AWAITING_INITIAL : DISCONNECT_REQUESTED; + } + + LOG.debug("Finding ProxyToServerConnection for: {}", serverHostAndPort); + + // Use the shared connection pool if enabled (disabled by default for backwards compatibility) + ServerConnectionPool pool = proxyServer.getServerConnectionPool(); + boolean usePool = pool != null; + boolean isConnect = ProxyUtils.isCONNECT(httpRequest); + boolean poolSharedMitm = usePool && proxyServer.isPoolSharedMitmConnections(); + boolean poolPerRequest = usePool && proxyServer.isPoolPerRequestInMitm(); + boolean useSharedPool = + usePool + && !isTunneling() + && !ProxyUtils.isSwitchingToWebSocketProtocol(httpRequest) + && (poolPerRequest || !isMitming()) + && (poolSharedMitm || !isConnect); + + boolean newConnectionRequired = false; + if (!useSharedPool) { + // For non-pooled mode, CONNECT, tunneling (WebSocket), or MITM, use dedicated connections + LOG.debug( + "Not using shared pool for: {} (pool enabled: {}, is CONNECT: {}, is Tunneling: {}, is MITM: {})", + serverHostAndPort, + usePool, + ProxyUtils.isCONNECT(httpRequest), + isTunneling(), + isMitming()); + currentServerConnection = serverConnectionsByHostAndPort.get(serverHostAndPort); + if (currentServerConnection == null) { + newConnectionRequired = true; + } + } else { + // For regular requests with pool enabled, get from shared pool + newConnectionRequired = true; + } - // Figure out our next state - if (ProxyUtils.isCONNECT(httpRequest)) { - return NEGOTIATING_CONNECT; - } else if (ProxyUtils.isChunked(httpRequest)) { - return AWAITING_CHUNK; + if (newConnectionRequired) { + try { + // Create dedicated connection for non-pooled/CONNECT/tunneling/MITM, or use pool + if (!useSharedPool) { + currentServerConnection = + ProxyToServerConnection.create( + proxyServer, + this, + serverHostAndPort, + currentFilters, + httpRequest, + globalTrafficShapingHandler); } else { + // Use the shared pool for regular requests + // Resolve ChainedProxy address before pooling to segregate connections by upstream route + ChainedProxy chainedProxy = null; + ChainedProxyManager chainedProxyManager = proxyServer.getChainProxyManager(); + if (chainedProxyManager != null) { + Queue chainedProxies = new ConcurrentLinkedQueue<>(); + chainedProxyManager.lookupChainedProxies( + httpRequest, chainedProxies, getClientDetails()); + if (!chainedProxies.isEmpty()) { + chainedProxy = chainedProxies.poll(); + } + } + InetSocketAddress chainedProxyAddress = + chainedProxy != null ? chainedProxy.getChainedProxyAddress() : null; + currentServerConnection = + pool.getOrCreateConnection( + serverHostAndPort, chainedProxyAddress, this, currentFilters, httpRequest); + } + + if (currentServerConnection == null) { + LOG.debug("Unable to create server connection, probably no chained proxies available"); + boolean keepAlive = writeBadGateway(httpRequest); + resumeReading(); + if (keepAlive) { return AWAITING_INITIAL; + } else { + return DISCONNECT_REQUESTED; + } } - } - /** - * Returns true if the specified request is a request to an origin server, rather than to a proxy server. If this - * request is being MITM'd, this method always returns false. The format of requests to a proxy server are defined - * in RFC 7230, section 5.3.2 (all other requests are considered requests to an origin server): -

-         When making a request to a proxy, other than a CONNECT or server-wide
-         OPTIONS request (as detailed below), a client MUST send the target
-         URI in absolute-form as the request-target.
-         [...]
-         An example absolute-form of request-line would be:
-         GET http://www.example.org/pub/WWW/TheProject.html HTTP/1.1
-         To allow for transition to the absolute-form for all requests in some
-         future version of HTTP, a server MUST accept the absolute-form in
-         requests, even though HTTP/1.1 clients will only send them in
-         requests to proxies.
-     
- * - * @param httpRequest the request to evaluate - * @return true if the specified request is a request to an origin server, otherwise false - */ - private boolean isRequestToOriginServer(HttpRequest httpRequest) { - // while MITMing, all HTTPS requests are requests to the origin server, since the client does not know - // the request is being MITM'd by the proxy - if (httpRequest.method() == HttpMethod.CONNECT || isMitming()) { - return false; + // Remember the connection for tracking (for non-pooled connections or + // CONNECT/tunneling/MITM). In per-request MITM mode, CONNECT connections are + // released to pool after flow completes, so they don't need session tracking. + if (!useSharedPool || (isConnect && !poolPerRequest)) { + serverConnectionsByHostAndPort.put( + serverHostAndPort, requireNonNull(currentServerConnection)); } - - // direct requests to the proxy have the path only without a scheme - String uri = httpRequest.uri(); - return !HTTP_SCHEME.matcher(uri).matches(); + } catch (UnknownHostException uhe) { + LOG.info("Bad Host {}", httpRequest.uri()); + boolean keepAlive = writeBadGateway(httpRequest); + resumeReading(); + if (keepAlive) { + return AWAITING_INITIAL; + } else { + return DISCONNECT_REQUESTED; + } + } + } else { + LOG.debug("Reusing existing server connection: {}", currentServerConnection); + numberOfReusedServerConnections.incrementAndGet(); } - @Override - protected void readHTTPChunk(HttpContent chunk) { - currentFilters.clientToProxyRequest(chunk); - currentFilters.proxyToServerRequest(chunk); - - currentServerConnection.write(chunk); + // For pooled connections, set the current client connection before writing + // This allows the server connection to know where to send the response + // Set for pooled requests including CONNECT and MITM when pool flags are active, + // so that reused connections route responses to the correct client + if (usePool + && !isTunneling() + && !ProxyUtils.isSwitchingToWebSocketProtocol(httpRequest) + && currentServerConnection != null + && (poolPerRequest + || poolSharedMitm + || (!ProxyUtils.isCONNECT(httpRequest) && !isMitming()))) { + currentServerConnection.setCurrentClientConnectionForRequest(this); } - @Override - protected void readRaw(ByteBuf buf) { - currentServerConnection.write(buf); + modifyRequestHeadersToReflectProxying(httpRequest); + + HttpResponse proxyToServerFilterResponse = currentFilters.proxyToServerRequest(httpRequest); + if (proxyToServerFilterResponse != null) { + LOG.debug( + "Responding to client with short-circuit response from filter: {}", + proxyToServerFilterResponse); + + if (usePool && currentServerConnection != null) { + currentServerConnection.setCurrentClientConnectionForRequest(null); + currentServerConnection.releaseToPool(); + } + + boolean keepAlive = respondWithShortCircuitResponse(proxyToServerFilterResponse); + if (keepAlive) { + return AWAITING_INITIAL; + } else { + return DISCONNECT_REQUESTED; + } } - /* ************************************************************************* - * Writing - **************************************************************************/ + LOG.debug("Writing request to ProxyToServerConnection"); + requireNonNull(currentServerConnection).write(httpRequest, currentFilters); - /** - * Send a response to the client. - * - * @param serverConnection - * the ProxyToServerConnection that's responding - * @param filters - * the filters to apply to the response - * @param currentHttpRequest - * the HttpRequest that prompted this response - * @param currentHttpResponse - * the HttpResponse corresponding to this data (when doing - * chunked transfers, this is the initial HttpResponse object - * that came in before the other chunks) - * @param httpObject - * the data with which to respond - */ - void respond(ProxyToServerConnection serverConnection, HttpFilters filters, - HttpRequest currentHttpRequest, HttpResponse currentHttpResponse, - HttpObject httpObject) { - // we are sending a response to the client, so we are done handling this request - if (currentRequest != null && currentRequest instanceof ReferenceCounted) { - ((ReferenceCounted)currentRequest).release(); - } - this.currentRequest = null; - - httpObject = filters.serverToProxyResponse(httpObject); - if (httpObject == null) { - forceDisconnect(serverConnection); - return; - } + // Figure out our next state + if (ProxyUtils.isCONNECT(httpRequest)) { + return NEGOTIATING_CONNECT; + } else if (ProxyUtils.isChunked(httpRequest)) { + return AWAITING_CHUNK; + } else { + return AWAITING_INITIAL; + } + } + + /** + * Returns true if the specified request is a request to an origin server, rather than to a proxy + * server. If this request is being MITM'd, this method always returns false. The format of + * requests to a proxy server are defined in RFC 7230, section 5.3.2 (all other requests are + * considered requests to an origin server): + * + *
+   * When making a request to a proxy, other than a CONNECT or server-wide
+   * OPTIONS request (as detailed below), a client MUST send the target
+   * URI in absolute-form as the request-target.
+   * [...]
+   * An example absolute-form of request-line would be:
+   * GET https://www.example.org/pub/WWW/TheProject.html HTTP/1.1
+   * To allow for transition to the absolute-form for all requests in some
+   * future version of HTTP, a server MUST accept the absolute-form in
+   * requests, even though HTTP/1.1 clients will only send them in
+   * requests to proxies.
+   * 
+ * + * @param httpRequest the request to evaluate + * @return true if the specified request is a request to an origin server, otherwise false + */ + private boolean isRequestToOriginServer(HttpRequest httpRequest) { + // while MITMing, all HTTPS requests are requests to the origin server, since the client does + // not know + // the request is being MITM'd by the proxy + if (httpRequest.method() == HttpMethod.CONNECT || isMitming()) { + return false; + } - if (httpObject instanceof HttpResponse) { - HttpResponse httpResponse = (HttpResponse) httpObject; - - // if this HttpResponse does not have any means of signaling the end of the message body other than closing - // the connection, convert the message to a "Transfer-Encoding: chunked" HTTP response. This avoids the need - // to close the client connection to indicate the end of the message. (Responses to HEAD requests "must be" empty.) - if (!ProxyUtils.isHEAD(currentHttpRequest) && !ProxyUtils.isResponseSelfTerminating(httpResponse)) { - // if this is not a FullHttpResponse, duplicate the HttpResponse from the server before sending it to - // the client. this allows us to set the Transfer-Encoding to chunked without interfering with netty's - // handling of the response from the server. if we modify the original HttpResponse from the server, - // netty will not generate the appropriate LastHttpContent when it detects the connection closure from - // the server (see HttpObjectDecoder#decodeLast). (This does not apply to FullHttpResponses, for which - // netty already generates the empty final chunk when Transfer-Encoding is chunked.) - if (!(httpResponse instanceof FullHttpResponse)) { - HttpResponse duplicateResponse = ProxyUtils.duplicateHttpResponse(httpResponse); - - // set the httpObject and httpResponse to the duplicated response, to allow all other standard processing - // (filtering, header modification for proxying, etc.) to be applied. - httpObject = httpResponse = duplicateResponse; - } + // direct requests to the proxy have the path only without a scheme + String uri = httpRequest.uri(); + return !ABSOLUTE_URI_PATTERN.matcher(uri).matches(); + } - HttpUtil.setTransferEncodingChunked(httpResponse, true); - } + @Override + protected void readHTTPChunk(HttpContent chunk) { + if (currentServerConnection == null) { + LOG.warn("Cannot forward HTTP chunk: no server connection"); + return; + } + currentFilters.clientToProxyRequest(chunk); + currentFilters.proxyToServerRequest(chunk); - fixHttpVersionHeaderIfNecessary(httpResponse); - modifyResponseHeadersToReflectProxying(httpResponse); - } + currentServerConnection.write(chunk); + } - httpObject = filters.proxyToClientResponse(httpObject); - if (httpObject == null) { - forceDisconnect(serverConnection); - return; - } + @Override + protected void readRaw(ByteBuf buf) { + if (currentServerConnection == null) { + LOG.warn("Cannot forward raw data: no server connection"); + return; + } + currentServerConnection.write(buf); + } + + /* ************************************************************************* + * Writing + **************************************************************************/ + + /** + * Send a response to the client. + * + * @param serverConnection the ProxyToServerConnection that's responding + * @param filters the filters to apply to the response + * @param currentHttpRequest the HttpRequest that prompted this response + * @param currentHttpResponse the HttpResponse corresponding to this data (when doing chunked + * transfers, this is the initial HttpResponse object that came in before the other chunks) + * @param httpObject the data with which to respond + */ + void respond( + ProxyToServerConnection serverConnection, + HttpFilters filters, + HttpRequest currentHttpRequest, + HttpResponse currentHttpResponse, + HttpObject httpObject) { + // we are sending a response to the client, so we are done handling this request + resetCurrentRequest(); + + httpObject = filters.serverToProxyResponse(httpObject); + if (httpObject == null) { + forceDisconnect(serverConnection); + return; + } - write(httpObject); + final boolean isSwitchingToWebSocketProtocol; + if (httpObject instanceof HttpResponse) { + HttpResponse httpResponse = (HttpResponse) httpObject; + + isSwitchingToWebSocketProtocol = ProxyUtils.isSwitchingToWebSocketProtocol(httpResponse); + + // if this HttpResponse does not have any means of signaling the end of the message body other + // than closing + // the connection, convert the message to a "Transfer-Encoding: chunked" HTTP response. This + // avoids the need + // to close the client connection to indicate the end of the message. (Responses to HEAD + // requests "must be" empty.) + if (!ProxyUtils.isHEAD(currentHttpRequest) + && !ProxyUtils.isResponseSelfTerminating(httpResponse)) { + // if this is not a FullHttpResponse, duplicate the HttpResponse from the server before + // sending it to + // the client. this allows us to set the Transfer-Encoding to chunked without interfering + // with netty's + // handling of the response from the server. if we modify the original HttpResponse from the + // server, + // netty will not generate the appropriate LastHttpContent when it detects the connection + // closure from + // the server (see HttpObjectDecoder#decodeLast). (This does not apply to FullHttpResponses, + // for which + // netty already generates the empty final chunk when Transfer-Encoding is chunked.) + if (!(httpResponse instanceof FullHttpResponse)) { + HttpResponse duplicateResponse = ProxyUtils.duplicateHttpResponse(httpResponse); + + // set the httpObject and httpResponse to the duplicated response, to allow all other + // standard processing + // (filtering, header modification for proxying, etc.) to be applied. + httpObject = httpResponse = duplicateResponse; + } + + HttpUtil.setTransferEncodingChunked(httpResponse, true); + } + + fixHttpVersionHeaderIfNecessary(httpResponse); + modifyResponseHeadersToReflectProxying(httpResponse); + + // modifyResponseHeadersToReflectProxying strips hop-by-hop headers (Upgrade, Connection), + // but a WebSocket upgrade response requires both to be present for the client to switch + // protocols. Re-add them after the general stripping. + if (isSwitchingToWebSocketProtocol) { + httpResponse.headers().set(HttpHeaderNames.UPGRADE, "websocket"); + httpResponse.headers().set(HttpHeaderNames.CONNECTION, "Upgrade"); + } + } else { + isSwitchingToWebSocketProtocol = false; + } - if (ProxyUtils.isLastChunk(httpObject)) { - writeEmptyBuffer(); - } + final HttpObject filteredhttpObject = filters.proxyToClientResponse(httpObject); + if (filteredhttpObject == null) { + forceDisconnect(serverConnection); + return; + } - closeConnectionsAfterWriteIfNecessary(serverConnection, - currentHttpRequest, currentHttpResponse, httpObject); + if (isSwitchingToWebSocketProtocol) { + serverConnection.switchToWebSocketProtocol(); + } + write(filteredhttpObject) + .addListener( + l -> { + if (isSwitchingToWebSocketProtocol) { + switchToWebSocketProtocol(serverConnection); + } else if (ProxyUtils.isLastChunk(filteredhttpObject)) { + writeEmptyBuffer(); + } + + closeConnectionsAfterWriteIfNecessary( + serverConnection, currentHttpRequest, currentHttpResponse, filteredhttpObject); + }); + } + + private void resetCurrentRequest() { + if (currentRequest != null && currentRequest instanceof ReferenceCounted) { + ((ReferenceCounted) currentRequest).release(); + } + currentRequest = null; + } + + private void switchToWebSocketProtocol(final ProxyToServerConnection serverConnection) { + final List orderedHandlersToRemove = + Arrays.asList( + HTTP_REQUEST_READ_MONITOR_NAME, + HTTP_RESPONSE_WRITTEN_MONITOR_NAME, + HTTP_PROXY_DECODER_NAME, + HTTP_ENCODER_NAME, + HTTP_DECODER_NAME); + if (channel.pipeline().get(MAIN_HANDLER_NAME) != null) { + channel + .pipeline() + .replace( + MAIN_HANDLER_NAME, + "pipe-to-server", + new WebSocketFramePipeHandler(serverConnection, currentFilters, true)); } + orderedHandlersToRemove.forEach(this::removeHandlerIfPresent); + } - /* ************************************************************************* - * Connection Lifecycle - **************************************************************************/ + /* ************************************************************************* + * Connection Lifecycle + **************************************************************************/ - /** - * Tells the Client that its HTTP CONNECT request was successful. - */ - ConnectionFlowStep RespondCONNECTSuccessful = new ConnectionFlowStep( - this, NEGOTIATING_CONNECT) { + /** Tells the Client that its HTTP CONNECT request was successful. */ + final ConnectionFlowStep RespondCONNECTSuccessful = + new ConnectionFlowStep<>(this, NEGOTIATING_CONNECT) { @Override boolean shouldSuppressInitialRequest() { - return true; + return true; } protected Future execute() { - LOG.debug("Responding with CONNECT successful"); - HttpResponse response = ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, - CONNECTION_ESTABLISHED); - response.headers().set(HttpHeaderNames.CONNECTION, HttpHeaderValues.KEEP_ALIVE); - ProxyUtils.addVia(response, proxyServer.getProxyAlias()); - return writeToChannel(response); - } - }; - - /** - * On connect of the client, start waiting for an initial - * {@link HttpRequest}. - */ - @Override - protected void connected() { - super.connected(); - become(AWAITING_INITIAL); - recordClientConnected(); - } - - void timedOut(ProxyToServerConnection serverConnection) { - if (currentServerConnection == serverConnection && this.lastReadTime > currentServerConnection.lastReadTime) { - // the idle timeout fired on the active server connection. send a timeout response to the client. - LOG.warn("Server timed out: {}", currentServerConnection); - currentFilters.serverToProxyResponseTimedOut(); - writeGatewayTimeout(currentRequest); - } + LOG.debug("Responding with CONNECT successful"); + HttpResponse response = + ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, CONNECTION_ESTABLISHED); + ProxyUtils.addVia(response, proxyServer.getProxyAlias()); + return writeToChannel(response); + } + }; + + /** On connect of the client, start waiting for an initial {@link HttpRequest}. */ + @Override + protected void connected() { + super.connected(); + become(AWAITING_INITIAL); + // recordClientConnected() is deferred, not called here: with PROXY protocol it must wait for + // the + // header so it reports the real client address (readHAProxyMessage); otherwise it fires on the + // first request (readHTTPInitial). The header isn't available yet at channel-active time. + } + + void timedOut(ProxyToServerConnection serverConnection) { + if (currentServerConnection == serverConnection + && lastReadTime > currentServerConnection.lastReadTime) { + // the idle timeout fired on the active server connection. send a timeout response to the + // client. + LOG.warn("Server timed out: {}", currentServerConnection); + currentFilters.serverToProxyResponseTimedOut(); + writeGatewayTimeout(currentRequest); } - - @Override - protected void timedOut() { - // idle timeout fired on the client channel. if we aren't waiting on a response from a server, hang up - if (currentServerConnection == null || this.lastReadTime <= currentServerConnection.lastReadTime) { - super.timedOut(); - } + } + + @Override + protected void timedOut() { + // idle timeout fired on the client channel. if we aren't waiting on a response from a server, + // hang up + // + // The original issue: when server.lastReadTime == 0 (server has never read) and + // client.lastReadTime > 0, the comparison lastReadTime <= server.lastReadTime evaluates to + // FALSE, preventing timeout even when the client IS idle. + // + // However, we need to be careful not to close when: + // 1. A request has been sent but the server hasn't responded yet (normal) + // 2. A request was sent, proxy generated 504 (server didn't respond), and we're waiting for + // the next request from client + // + // We distinguish these by checking if the server connection has an "initialRequest" that has + // been written to the server but not yet reset. If initialRequest is not null, a request + // has been written and we're waiting for response. + boolean requestHasBeenWritten = false; + if (currentServerConnection != null) { + // Check if a request has been written to the server but not yet completed + HttpRequest initialRequest = currentServerConnection.getInitialRequest(); + requestHasBeenWritten = initialRequest != null; } - /** - * On disconnect of the client, disconnect all server connections. - */ - @Override - protected void disconnected() { - super.disconnected(); - for (ProxyToServerConnection serverConnection : serverConnectionsByHostAndPort - .values()) { - serverConnection.disconnect(); - } - recordClientDisconnected(); + // no request has been transmitted to server + if (currentServerConnection == null + || + // server has never read anything yet + (currentServerConnection.lastReadTime == 0 + // no initial request has been written to the server + && !requestHasBeenWritten + // there are no current request from the client + && currentRequest == null) + || + // The client hasn't sent data as recently as the server + // - Both sides are idle (no activity on either end) + // - After a response is complete and neither client nor server has sent anything new + lastReadTime <= currentServerConnection.lastReadTime) { + super.timedOut(); + recordConnectionTimedOut(); } - - /** - * Called when {@link ProxyToServerConnection} starts its connection flow. - */ - protected void serverConnectionFlowStarted( - ProxyToServerConnection serverConnection) { - stopReading(); - this.numberOfCurrentlyConnectingServers.incrementAndGet(); + } + + /** On disconnect of the client, disconnect all server connections. */ + @Override + protected void disconnected() { + super.disconnected(); + boolean poolSharedMitm = proxyServer.isPoolSharedMitmConnections(); + boolean poolPerRequest = proxyServer.isPoolPerRequestInMitm(); + for (ProxyToServerConnection serverConnection : serverConnectionsByHostAndPort.values()) { + // Phase 1: release pooled MITM connections back to the pool + // Phase 2: connections are already in pool after each response, no release needed + if (poolSharedMitm && !poolPerRequest && serverConnection.isManagedByPool()) { + serverConnection.releaseToPool(); + } else { + serverConnection.disconnect(); + } } - - /** - * If the {@link ProxyToServerConnection} completes its connection lifecycle - * successfully, this method is called to let us know about it. - */ - protected void serverConnectionSucceeded( - ProxyToServerConnection serverConnection, - boolean shouldForwardInitialRequest) { - LOG.debug("Connection to server succeeded: {}", - serverConnection.getRemoteAddress()); - resumeReadingIfNecessary(); - become(shouldForwardInitialRequest ? getCurrentState() - : AWAITING_INITIAL); - numberOfCurrentlyConnectedServers.incrementAndGet(); - } - - /** - * If the {@link ProxyToServerConnection} fails to complete its connection - * lifecycle successfully, this method is called to let us know about it. - * - *

- * After failing to connect to the server, one of two things can happen: - *

- * - *
    - *
  1. If the server was a chained proxy, we fall back to connecting to the - * ultimate endpoint directly.
  2. - *
  3. If the server was the ultimate endpoint, we return a 502 Bad Gateway - * to the client.
  4. - *
- * - * @param serverConnection - * @param lastStateBeforeFailure - * @param cause - * what caused the failure - * - * @return true if we're falling back to a another chained proxy (or direct - * connection) and trying again - */ - protected boolean serverConnectionFailed( - ProxyToServerConnection serverConnection, - ConnectionState lastStateBeforeFailure, - Throwable cause) { - resumeReadingIfNecessary(); - HttpRequest initialRequest = serverConnection.getInitialRequest(); - try { - boolean retrying = serverConnection.connectionFailed(cause); - if (retrying) { - LOG.debug("Failed to connect to upstream server or chained proxy. Retrying connection. Last state before failure: {}", - lastStateBeforeFailure, cause); - return true; - } else { - LOG.debug( - "Connection to upstream server or chained proxy failed: {}. Last state before failure: {}", - serverConnection.getRemoteAddress(), - lastStateBeforeFailure, - cause); - connectionFailedUnrecoverably(initialRequest, serverConnection); - return false; - } - } catch (UnknownHostException uhe) { - connectionFailedUnrecoverably(initialRequest, serverConnection); - return false; - } + recordClientDisconnected(); + } + + /** Called when {@link ProxyToServerConnection} starts its connection flow. */ + protected void serverConnectionFlowStarted(ProxyToServerConnection serverConnection) { + stopReading(); + numberOfCurrentlyConnectingServers.incrementAndGet(); + } + + /** + * If the {@link ProxyToServerConnection} completes its connection lifecycle successfully, this + * method is called to let us know about it. + */ + protected void serverConnectionSucceeded( + ProxyToServerConnection serverConnection, boolean shouldForwardInitialRequest) { + LOG.debug("Connection to server succeeded: {}", serverConnection.getRemoteAddress()); + resumeReadingIfNecessary(); + become(shouldForwardInitialRequest ? getCurrentState() : AWAITING_INITIAL); + numberOfCurrentlyConnectedServers.incrementAndGet(); + } + + /** + * If the {@link ProxyToServerConnection} fails to complete its connection lifecycle successfully, + * this method is called to let us know about it. + * + *

After failing to connect to the server, one of two things can happen: + * + *

    + *
  1. If the server was a chained proxy, we fall back to connecting to the ultimate endpoint + * directly. + *
  2. If the server was the ultimate endpoint, we return a 502 Bad Gateway to the client. + *
+ * + * @param serverConnection + * @param lastStateBeforeFailure + * @param cause what caused the failure + * @return true if we're falling back to another chained proxy (or direct connection) and trying + * again + */ + protected boolean serverConnectionFailed( + ProxyToServerConnection serverConnection, + ConnectionState lastStateBeforeFailure, + Throwable cause) { + resumeReadingIfNecessary(); + HttpRequest initialRequest = serverConnection.getInitialRequest(); + try { + boolean retrying = serverConnection.connectionFailed(cause); + if (retrying) { + LOG.debug( + "Failed to connect to upstream server or chained proxy. Retrying connection. Last state before failure: {}", + lastStateBeforeFailure, + cause); + return true; + } else { + LOG.debug( + "Connection to upstream server or chained proxy failed: {}. Last state before failure: {}", + serverConnection.getRemoteAddress(), + lastStateBeforeFailure, + cause); + connectionFailedUnrecoverably(initialRequest, serverConnection); + return false; + } + } catch (UnknownHostException uhe) { + connectionFailedUnrecoverably(initialRequest, serverConnection); + return false; } - - private void connectionFailedUnrecoverably(HttpRequest initialRequest, ProxyToServerConnection serverConnection) { - // the connection to the server failed, so disconnect the server and remove the ProxyToServerConnection from the - // map of open server connections - serverConnection.disconnect(); - this.serverConnectionsByHostAndPort.remove(serverConnection.getServerHostAndPort()); - - boolean keepAlive = writeBadGateway(initialRequest); - if (keepAlive) { - become(AWAITING_INITIAL); - } else { - become(DISCONNECT_REQUESTED); - } + } + + private void connectionFailedUnrecoverably( + HttpRequest initialRequest, ProxyToServerConnection serverConnection) { + // the connection to the server failed, so disconnect the server and remove the + // ProxyToServerConnection from the + // map of open server connections + serverConnection.disconnect(); + serverConnectionsByHostAndPort.remove(serverConnection.getServerHostAndPort()); + + boolean keepAlive = writeBadGateway(initialRequest); + if (keepAlive) { + become(AWAITING_INITIAL); + } else { + become(DISCONNECT_REQUESTED); } + } - private void resumeReadingIfNecessary() { - if (this.numberOfCurrentlyConnectingServers.decrementAndGet() == 0) { - LOG.debug("All servers have finished attempting to connect, resuming reading from client."); - resumeReading(); - } + private void resumeReadingIfNecessary() { + if (numberOfCurrentlyConnectingServers.decrementAndGet() == 0) { + LOG.debug("All servers have finished attempting to connect, resuming reading from client."); + resumeReading(); } - - /* ************************************************************************* - * Other Lifecycle - **************************************************************************/ - - /** - * On disconnect of the server, track that we have one fewer connected - * servers and then disconnect the client if necessary. - */ - protected void serverDisconnected(ProxyToServerConnection serverConnection) { - numberOfCurrentlyConnectedServers.decrementAndGet(); - - // for non-SSL connections, do not disconnect the client from the proxy, even if this was the last server connection. - // this allows clients to continue to use the open connection to the proxy to make future requests. for SSL - // connections, whether we are tunneling or MITMing, we need to disconnect the client because there is always - // exactly one ClientToProxyConnection per ProxyToServerConnection, and vice versa. - if (isTunneling() || isMitming()) { - disconnect(); - } + } + + /* ************************************************************************* + * Other Lifecycle + **************************************************************************/ + + /** + * On disconnect of the server, track that we have one fewer connected servers and then disconnect + * the client if necessary. + */ + protected void serverDisconnected(ProxyToServerConnection serverConnection) { + numberOfCurrentlyConnectedServers.decrementAndGet(); + + // for non-SSL connections, do not disconnect the client from the proxy, even if this was the + // last server connection. + // this allows clients to continue to use the open connection to the proxy to make future + // requests. for SSL + // connections, whether we are tunneling or MITMing, we need to disconnect the client because + // there is always + // exactly one ClientToProxyConnection per ProxyToServerConnection, and vice versa. + if (isTunneling() || isMitming()) { + disconnect(); } - - /** - * When the ClientToProxyConnection becomes saturated, stop reading on all - * associated ProxyToServerConnections. - */ - @Override - synchronized protected void becameSaturated() { - super.becameSaturated(); - for (ProxyToServerConnection serverConnection : serverConnectionsByHostAndPort - .values()) { - synchronized (serverConnection) { - if (this.isSaturated()) { - serverConnection.stopReading(); - } - } - } + } + + /** + * When the ClientToProxyConnection becomes saturated, stop reading on all associated + * ProxyToServerConnections. + */ + @Override + protected synchronized void becameSaturated() { + super.becameSaturated(); + recordConnectionSaturated(); + ProxyToServerConnection current = currentServerConnection; + for (ProxyToServerConnection serverConnection : serverConnectionsByHostAndPort.values()) { + synchronized (serverConnection) { + if (isSaturated()) { + serverConnection.stopReading(); + } + } } - - /** - * When the ClientToProxyConnection becomes writable, resume reading on all - * associated ProxyToServerConnections. - */ - @Override - synchronized protected void becameWritable() { - super.becameWritable(); - for (ProxyToServerConnection serverConnection : serverConnectionsByHostAndPort - .values()) { - synchronized (serverConnection) { - if (!this.isSaturated()) { - serverConnection.resumeReading(); - } - } + if (current != null) { + synchronized (current) { + if (isSaturated()) { + current.stopReading(); } + } } - - /** - * When a server becomes saturated, we stop reading from the client. - */ - synchronized protected void serverBecameSaturated( - ProxyToServerConnection serverConnection) { - if (serverConnection.isSaturated()) { - LOG.info("Connection to server became saturated, stopping reading"); - stopReading(); - } + } + + /** + * When the ClientToProxyConnection becomes writable, resume reading on all associated + * ProxyToServerConnections. + */ + @Override + protected synchronized void becameWritable() { + super.becameWritable(); + recordConnectionWritable(); + ProxyToServerConnection current = currentServerConnection; + for (ProxyToServerConnection serverConnection : serverConnectionsByHostAndPort.values()) { + synchronized (serverConnection) { + if (!isSaturated()) { + serverConnection.resumeReading(); + } + } } - - /** - * When a server becomes writeable, we check to see if all servers are - * writeable and if they are, we resume reading. - */ - synchronized protected void serverBecameWriteable( - ProxyToServerConnection serverConnection) { - boolean anyServersSaturated = false; - for (ProxyToServerConnection otherServerConnection : serverConnectionsByHostAndPort - .values()) { - if (otherServerConnection.isSaturated()) { - anyServersSaturated = true; - break; - } - } - if (!anyServersSaturated) { - LOG.info("All server connections writeable, resuming reading"); - resumeReading(); + if (current != null) { + synchronized (current) { + if (!isSaturated()) { + current.resumeReading(); } + } } + } - @Override - protected void exceptionCaught(Throwable cause) { - try { - if (cause instanceof IOException) { - // IOExceptions are expected errors, for example when a browser is killed and aborts a connection. - // rather than flood the logs with stack traces for these expected exceptions, we log the message at the - // INFO level and the stack trace at the DEBUG level. - LOG.info("An IOException occurred on ClientToProxyConnection: " + cause.getMessage()); - LOG.debug("An IOException occurred on ClientToProxyConnection", cause); - } else if (cause instanceof RejectedExecutionException) { - LOG.info("An executor rejected a read or write operation on the ClientToProxyConnection (this is normal if the proxy is shutting down). Message: " + cause.getMessage()); - LOG.debug("A RejectedExecutionException occurred on ClientToProxyConnection", cause); - } else { - LOG.error("Caught an exception on ClientToProxyConnection", cause); - } - } finally { - // always disconnect the client when an exception occurs on the channel - disconnect(); - } + /** When a server becomes saturated, we stop reading from the client. */ + protected synchronized void serverBecameSaturated(ProxyToServerConnection serverConnection) { + if (serverConnection.isSaturated()) { + LOG.info("Connection to server became saturated, stopping reading"); + stopReading(); + } + } + + /** + * When a server becomes writeable, we check to see if all servers are writeable and if they are, + * we resume reading. + */ + protected synchronized void serverBecameWriteable(ProxyToServerConnection serverConnection) { + boolean anyServersSaturated = false; + ProxyToServerConnection current = currentServerConnection; + for (ProxyToServerConnection otherServerConnection : serverConnectionsByHostAndPort.values()) { + if (otherServerConnection.isSaturated()) { + anyServersSaturated = true; + break; + } + } + if (!anyServersSaturated + && current != null + && current != serverConnection + && current.isSaturated()) { + anyServersSaturated = true; + } + if (!anyServersSaturated) { + LOG.info("All server connections writeable, resuming reading"); + resumeReading(); + } + } + + @Override + protected void exceptionCaught(Throwable cause) { + try { + recordConnectionExceptionCaught(cause); + if (cause instanceof IOException) { + // IOExceptions are expected errors, for example when a browser is killed and aborts a + // connection. + // rather than flood the logs with stack traces for these expected exceptions, we log the + // message at the + // INFO level and the stack trace at the DEBUG level. + LOG.info("An IOException occurred on ClientToProxyConnection: " + cause.getMessage()); + LOG.debug("An IOException occurred on ClientToProxyConnection", cause); + } else if (cause instanceof RejectedExecutionException) { + LOG.info( + "An executor rejected a read or write operation on the ClientToProxyConnection (this is normal if the proxy is shutting down). Message: " + + cause.getMessage()); + LOG.debug("A RejectedExecutionException occurred on ClientToProxyConnection", cause); + } else { + LOG.error("Caught an exception on ClientToProxyConnection", cause); + } + } finally { + // always disconnect the client when an exception occurs on the channel + disconnect(); + } + } + + /* ************************************************************************* + * Connection Management + **************************************************************************/ + + /** + * Initialize the {@link ChannelPipeline} for the client to proxy channel. LittleProxy acts like a + * server here. + * + *

A {@link ChannelPipeline} invokes the read (Inbound) handlers in ascending ordering of the + * list and then the write (Outbound) handlers in descending ordering. + * + *

Regarding the Javadoc of {@link HttpObjectAggregator} it's needed to have the {@link + * HttpResponseEncoder} or {@link io.netty.handler.codec.http.HttpRequestEncoder} before the + * {@link HttpObjectAggregator} in the {@link ChannelPipeline}. + * + *

If an {@link SslEngineSource} is provided, SSL encryption is enabled on the pipeline. When + * the proxy protocol is enabled, the {@link HAProxyMessageDecoder} is added after the SSL handler + * setup to ensure it is positioned before the {@link io.netty.handler.ssl.SslHandler} in the + * inbound pipeline, so that the PROXY protocol header is decoded before the TLS handshake begins. + * + * @param pipeline the {@link ChannelPipeline} to configure + * @param sslEngineSource the {@link SslEngineSource} for client-to-proxy encryption, or {@code + * null} if SSL is not enabled + * @param authenticateClients whether to require client certificate authentication + */ + private void initChannelPipeline( + ChannelPipeline pipeline, + @Nullable SslEngineSource sslEngineSource, + boolean authenticateClients) { + LOG.debug("Configuring ChannelPipeline"); + + pipeline.addLast("bytesReadMonitor", bytesReadMonitor); + pipeline.addLast("bytesWrittenMonitor", bytesWrittenMonitor); + + pipeline.addLast(HTTP_ENCODER_NAME, new HttpResponseEncoder()); + // We want to allow longer request lines, headers, and chunks + // respectively. + pipeline.addLast( + HTTP_DECODER_NAME, + new HttpRequestDecoder( + proxyServer.getMaxInitialLineLength(), + proxyServer.getMaxHeaderSize(), + proxyServer.getMaxChunkSize())); + + // Enable aggregation for filtering if necessary + int numberOfBytesToBuffer = proxyServer.getFiltersSource().getMaximumRequestBufferSizeInBytes(); + if (numberOfBytesToBuffer > 0) { + aggregateContentForFiltering(pipeline, numberOfBytesToBuffer); } - /* ************************************************************************* - * Connection Management - **************************************************************************/ - - /** - * Initialize the {@link ChannelPipeline} for the client to proxy channel. - * LittleProxy acts like a server here. - * - * A {@link ChannelPipeline} invokes the read (Inbound) handlers in - * ascending ordering of the list and then the write (Outbound) handlers in - * descending ordering. - * - * Regarding the Javadoc of {@link HttpObjectAggregator} it's needed to have - * the {@link HttpResponseEncoder} or {@link io.netty.handler.codec.http.HttpRequestEncoder} before the - * {@link HttpObjectAggregator} in the {@link ChannelPipeline}. - */ - private void initChannelPipeline(ChannelPipeline pipeline) { - LOG.debug("Configuring ChannelPipeline"); - - pipeline.addLast("bytesReadMonitor", bytesReadMonitor); - pipeline.addLast("bytesWrittenMonitor", bytesWrittenMonitor); - - pipeline.addLast("encoder", new HttpResponseEncoder()); - if (isAcceptProxyProtocol()) { - pipeline.addLast("proxy-protocol-decoder", new HAProxyMessageDecoder()); - } - // We want to allow longer request lines, headers, and chunks - // respectively. - pipeline.addLast("decoder", new HttpRequestDecoder( - proxyServer.getMaxInitialLineLength(), - proxyServer.getMaxHeaderSize(), - proxyServer.getMaxChunkSize())); - - // Enable aggregation for filtering if necessary - int numberOfBytesToBuffer = proxyServer.getFiltersSource() - .getMaximumRequestBufferSizeInBytes(); - if (numberOfBytesToBuffer > 0) { - aggregateContentForFiltering(pipeline, numberOfBytesToBuffer); - } + pipeline.addLast(HTTP_REQUEST_READ_MONITOR_NAME, requestReadMonitor); + pipeline.addLast(HTTP_RESPONSE_WRITTEN_MONITOR_NAME, responseWrittenMonitor); - pipeline.addLast("requestReadMonitor", requestReadMonitor); - pipeline.addLast("responseWrittenMonitor", responseWrittenMonitor); + pipeline.addLast("idle", new IdleStateHandler(0, 0, proxyServer.getIdleConnectionTimeout())); - pipeline.addLast( - "idle", - new IdleStateHandler(0, 0, proxyServer - .getIdleConnectionTimeout())); + pipeline.addLast(MAIN_HANDLER_NAME, this); - pipeline.addLast("handler", this); + if (sslEngineSource != null) { + LOG.debug("Enabling encryption of traffic from client to proxy"); + SSLEngine sslEngine = sslEngineSource.newSslEngine(); + recordClientSSLHandshakeStarted(); + encrypt(pipeline, sslEngine, authenticateClients) + .addListener( + future -> { + if (future.isSuccess()) { + clientSslSession = sslEngine.getSession(); + recordClientSSLHandshakeSucceeded(); + } + }); } - /** - * Is the proxy server set to accept a proxy protocol header - * @return True if the proxy server set to accept a proxy protocol header. False otherwise - */ - boolean isAcceptProxyProtocol() { - return proxyServer.isAcceptProxyProtocol(); + if (isAcceptProxyProtocol()) { + pipeline.addFirst(HTTP_PROXY_DECODER_NAME, new HAProxyMessageDecoder()); } - - /** - * Is the proxy server set to send a proxy protocol header - * @return True if the proxy server set to send a proxy protocol header. False otherwise - */ - boolean isSendProxyProtocol() { - return proxyServer.isSendProxyProtocol(); + } + + private void removeHandlerIfPresent(String name) { + removeHandlerIfPresent(channel.pipeline(), name); + } + + /** + * Is the proxy server set to accept a proxy protocol header + * + * @return True if the proxy server set to accept a proxy protocol header. False otherwise + */ + boolean isAcceptProxyProtocol() { + return proxyServer.isAcceptProxyProtocol(); + } + + /** + * Is the proxy server set to send a proxy protocol header + * + * @return True if the proxy server set to send a proxy protocol header. False otherwise + */ + boolean isSendProxyProtocol() { + return proxyServer.isSendProxyProtocol(); + } + + /** + * This method takes care of closing client to proxy and/or proxy to server connections after + * finishing writing. + */ + private void closeConnectionsAfterWriteIfNecessary( + ProxyToServerConnection serverConnection, + HttpRequest currentHttpRequest, + HttpResponse currentHttpResponse, + HttpObject httpObject) { + boolean closeServerConnection = + shouldCloseServerConnection(currentHttpRequest, currentHttpResponse, httpObject); + boolean closeClientConnection = + shouldCloseClientConnection(currentHttpRequest, currentHttpResponse, httpObject); + + if (closeServerConnection) { + LOG.debug("Closing remote connection after writing to client"); + serverConnection.disconnect(); } - /** - * This method takes care of closing client to proxy and/or proxy to server - * connections after finishing a write. - */ - private void closeConnectionsAfterWriteIfNecessary( - ProxyToServerConnection serverConnection, - HttpRequest currentHttpRequest, HttpResponse currentHttpResponse, - HttpObject httpObject) { - boolean closeServerConnection = shouldCloseServerConnection( - currentHttpRequest, currentHttpResponse, httpObject); - boolean closeClientConnection = shouldCloseClientConnection( - currentHttpRequest, currentHttpResponse, httpObject); - - if (closeServerConnection) { - LOG.debug("Closing remote connection after writing to client"); - serverConnection.disconnect(); - } - - if (closeClientConnection) { - LOG.debug("Closing connection to client after writes"); - disconnect(); + if (closeClientConnection) { + LOG.debug("Closing connection to client after writes"); + disconnect(); + } + } + + private void forceDisconnect(ProxyToServerConnection serverConnection) { + LOG.debug("Forcing disconnect"); + serverConnection.disconnect(); + disconnect(); + } + + /** Determine whether the client connection should be closed. */ + private boolean shouldCloseClientConnection( + HttpRequest req, HttpResponse res, HttpObject httpObject) { + if (ProxyUtils.isChunked(res)) { + // If the response is chunked, we want to return false unless it's + // the last chunk. If it is the last chunk, then we want to pass + // through to the same close semantics we'd otherwise use. + if (httpObject != null) { + if (!ProxyUtils.isLastChunk(httpObject)) { + String uri = null; + if (req != null) { + uri = req.uri(); + } + LOG.debug("Not closing client connection on middle chunk for {}", uri); + return false; + } else { + LOG.debug("Handling last chunk. Using normal client connection closing rules."); } + } } - private void forceDisconnect(ProxyToServerConnection serverConnection) { - LOG.debug("Forcing disconnect"); - serverConnection.disconnect(); - disconnect(); + if (!HttpUtil.isKeepAlive(req)) { + LOG.debug("Closing client connection since request is not keep alive: {}", req); + // Here we simply want to close the connection because the + // client itself has requested it be closed in the request. + return true; } - /** - * Determine whether or not the client connection should be closed. - */ - private boolean shouldCloseClientConnection(HttpRequest req, - HttpResponse res, HttpObject httpObject) { - if (ProxyUtils.isChunked(res)) { - // If the response is chunked, we want to return false unless it's - // the last chunk. If it is the last chunk, then we want to pass - // through to the same close semantics we'd otherwise use. - if (httpObject != null) { - if (!ProxyUtils.isLastChunk(httpObject)) { - String uri = null; - if (req != null) { - uri = req.uri(); - } - LOG.debug("Not closing client connection on middle chunk for {}", uri); - return false; - } else { - LOG.debug("Handling last chunk. Using normal client connection closing rules."); - } - } - } - - if (!HttpUtil.isKeepAlive(req)) { - LOG.debug("Closing client connection since request is not keep alive: {}", req); - // Here we simply want to close the connection because the - // client itself has requested it be closed in the request. - return true; + // ignore the response's keep-alive; we can keep this client connection open as long as the + // client allows it. + + LOG.debug("Not closing client connection for request: {}", req); + return false; + } + + /** + * Determines if the remote connection should be closed based on the request and response pair. If + * the request is HTTP 1.0 with no keep-alive header, for example, the connection should be + * closed. + * + *

This in part determines if we should close the connection. Here's the relevant section of + * RFC 2616: + * + *

"HTTP/1.1 defines the "close" connection option for the sender to signal that the connection + * will be closed after completion of the response. For example, + * + *

Connection: close + * + *

in either the request or the response header fields indicates that the connection SHOULD NOT + * be considered "persistent" (section 8.1) after the current request/response is complete." + * + * @param req The request. + * @param res The response. + * @param msg The message. + * @return Returns true if the connection should close. + */ + private boolean shouldCloseServerConnection(HttpRequest req, HttpResponse res, HttpObject msg) { + if (ProxyUtils.isChunked(res)) { + // If the response is chunked, we want to return false unless it's + // the last chunk. If it is the last chunk, then we want to pass + // through to the same close semantics we'd otherwise use. + if (msg != null) { + if (!ProxyUtils.isLastChunk(msg)) { + String uri = null; + if (req != null) { + uri = req.uri(); + } + LOG.debug("Not closing server connection on middle chunk for {}", uri); + return false; + } else { + LOG.debug("Handling last chunk. Using normal server connection closing rules."); } + } + } - // ignore the response's keep-alive; we can keep this client connection open as long as the client allows it. + // ignore the request's keep-alive; we can keep this server connection open as long as the + // server allows it. - LOG.debug("Not closing client connection for request: {}", req); - return false; + if (!HttpUtil.isKeepAlive(res)) { + LOG.debug("Closing server connection since response is not keep alive: {}", res); + // In this case, we want to honor the Connection: close header + // from the remote server and close that connection. We don't + // necessarily want to close the connection to the client, however + // as it's possible it has other connections open. + return true; } - /** - * Determines if the remote connection should be closed based on the request - * and response pair. If the request is HTTP 1.0 with no keep-alive header, - * for example, the connection should be closed. - * - * This in part determines if we should close the connection. Here's the - * relevant section of RFC 2616: - * - * "HTTP/1.1 defines the "close" connection option for the sender to signal - * that the connection will be closed after completion of the response. For - * example, - * - * Connection: close - * - * in either the request or the response header fields indicates that the - * connection SHOULD NOT be considered `persistent' (section 8.1) after the - * current request/response is complete." - * - * @param req - * The request. - * @param res - * The response. - * @param msg - * The message. - * @return Returns true if the connection should close. - */ - private boolean shouldCloseServerConnection(HttpRequest req, - HttpResponse res, HttpObject msg) { - if (ProxyUtils.isChunked(res)) { - // If the response is chunked, we want to return false unless it's - // the last chunk. If it is the last chunk, then we want to pass - // through to the same close semantics we'd otherwise use. - if (msg != null) { - if (!ProxyUtils.isLastChunk(msg)) { - String uri = null; - if (req != null) { - uri = req.uri(); - } - LOG.debug("Not closing server connection on middle chunk for {}", uri); - return false; - } else { - LOG.debug("Handling last chunk. Using normal server connection closing rules."); - } - } - } + LOG.debug("Not closing server connection for response: {}", res); + return false; + } + + /* ************************************************************************* + * Authentication + **************************************************************************/ + + /** + * Checks whether the given HttpRequest requires authentication. + * + *

If the request contains credentials, these are checked. + * + *

If authentication is still required, either because no credentials were provided or the + * credentials were wrong, this writes a 407 response to the client. + */ + private boolean authenticationRequired(HttpRequest request) { + + if (authenticated.get()) { + return false; + } - // ignore the request's keep-alive; we can keep this server connection open as long as the server allows it. + final ProxyAuthenticator authenticator = proxyServer.getProxyAuthenticator(); - if (!HttpUtil.isKeepAlive(res)) { - LOG.debug("Closing server connection since response is not keep alive: {}", res); - // In this case, we want to honor the Connection: close header - // from the remote server and close that connection. We don't - // necessarily want to close the connection to the client, however - // as it's possible it has other connections open. - return true; - } + if (authenticator == null) return false; - LOG.debug("Not closing server connection for response: {}", res); - return false; + if (!request.headers().contains(HttpHeaderNames.PROXY_AUTHORIZATION)) { + writeAuthenticationRequired(authenticator.getRealm()); + return true; } - /* ************************************************************************* - * Authentication - **************************************************************************/ - - /** - *

- * Checks whether the given HttpRequest requires authentication. - *

- * - *

- * If the request contains credentials, these are checked. - *

- * - *

- * If authentication is still required, either because no credentials were - * provided or the credentials were wrong, this writes a 407 response to the - * client. - *

- */ - private boolean authenticationRequired(HttpRequest request) { + List values = request.headers().getAll(HttpHeaderNames.PROXY_AUTHORIZATION); + String fullValue = values.iterator().next(); + String value = StringUtils.substringAfter(fullValue, "Basic ").trim(); - if (authenticated.get()) { - return false; - } + String decodedValue = new String(Base64.getDecoder().decode(value), UTF_8); - final ProxyAuthenticator authenticator = proxyServer - .getProxyAuthenticator(); + String userName = StringUtils.substringBefore(decodedValue, ":"); + String password = StringUtils.substringAfter(decodedValue, ":"); + if (!authenticator.authenticate(userName, password)) { + writeAuthenticationRequired(authenticator.getRealm()); + return true; + } + clientDetails.setUserName(userName); + + LOG.debug("Got proxy authorization!"); + // We need to remove the header before sending the request on. + String authentication = request.headers().get(HttpHeaderNames.PROXY_AUTHORIZATION); + LOG.debug(authentication); + request.headers().remove(HttpHeaderNames.PROXY_AUTHORIZATION); + authenticated.set(true); + return false; + } + + private void writeAuthenticationRequired(String realm) { + String body = + "\n" + + "\n" + + "407 Proxy Authentication Required\n" + + "\n" + + "

Proxy Authentication Required

\n" + + "

This server could not verify that you\n" + + "are authorized to access the document\n" + + "requested. Either you supplied the wrong\n" + + "credentials (e.g., bad password), or your\n" + + "browser doesn't understand how to supply\n" + + "the credentials required.

\n" + + "\n"; + FullHttpResponse response = + ProxyUtils.createFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.PROXY_AUTHENTICATION_REQUIRED, body); + response.headers().set(HttpHeaderNames.DATE, dateHeaderValue()); + response + .headers() + .set( + HttpHeaderNames.PROXY_AUTHENTICATE, + "Basic realm=\"" + (realm == null ? "Restricted Files" : realm) + "\""); + write(response); + } + + private String dateHeaderValue() { + return LocalDateTime.now() + .atZone(ZoneId.of("GMT")) + .format(ofPattern("EEE, dd MMM yyyy HH:mm:ss zzz")); + } + + /* ************************************************************************* + * Request/Response Rewriting + **************************************************************************/ + + /** Copy the given {@link HttpRequest} verbatim. */ + @NonNull + @CheckReturnValue + private HttpRequest copy(HttpRequest original) { + if (original instanceof FullHttpRequest) { + return ((FullHttpRequest) original).copy(); + } else { + HttpRequest request = + new DefaultHttpRequest(original.protocolVersion(), original.method(), original.uri()); + request.headers().set(original.headers()); + return request; + } + } + + /** + * Chunked encoding is an HTTP 1.1 feature, but sometimes we get a chunked response that reports + * its HTTP version as 1.0. In this case, we change it to 1.1. + */ + private void fixHttpVersionHeaderIfNecessary(HttpResponse httpResponse) { + String te = httpResponse.headers().get(HttpHeaderNames.TRANSFER_ENCODING); + if (StringUtils.isNotBlank(te) && te.equalsIgnoreCase(HttpHeaderValues.CHUNKED.toString())) { + if (httpResponse.protocolVersion() != HttpVersion.HTTP_1_1) { + LOG.debug("Fixing HTTP version."); + httpResponse.setProtocolVersion(HttpVersion.HTTP_1_1); + } + } + } + + /** + * If and only if our proxy is not running in transparent mode, modify the request headers to + * reflect that it was proxied. + */ + private void modifyRequestHeadersToReflectProxying(HttpRequest httpRequest) { + if (isNextHopOriginServer()) { + /* + * We are making the request to the origin server, so must modify + * the 'absolute-URI' into the 'origin-form' as per RFC 7230 + * section 5.3.1. + * + * This must happen even for 'transparent' mode, otherwise the origin + * server could infer that the request came via a proxy server. + */ + LOG.debug("Modifying request for proxy chaining"); + // Strip host from uri + String uri = httpRequest.uri(); + String adjustedUri = ProxyUtils.stripHost(uri); + LOG.debug("Stripped host from uri: {} yielding: {}", uri, adjustedUri); + httpRequest.setUri(adjustedUri); + } + if (!proxyServer.isTransparent()) { + LOG.debug("Modifying request headers for proxying"); - if (authenticator == null) - return false; + HttpHeaders headers = httpRequest.headers(); - if (!request.headers().contains(HttpHeaderNames.PROXY_AUTHORIZATION)) { - writeAuthenticationRequired(authenticator.getRealm()); - return true; - } + // Remove sdch from encodings we accept since we can't decode it. + ProxyUtils.removeSdchEncoding(headers); + switchProxyConnectionHeader(headers); + stripConnectionTokens(headers); - List values = request.headers().getAll( - HttpHeaderNames.PROXY_AUTHORIZATION); - String fullValue = values.iterator().next(); - String value = StringUtils.substringAfter(fullValue, "Basic ").trim(); + stripHopByHopHeaders(headers); - byte[] decodedValue = BaseEncoding.base64().decode(value); + // If we're forwarding to an upstream proxy that requires authentication, add the credentials + if (shouldPreserveProxyAuthorizationForUpstream()) { + addUpstreamProxyAuthorization(headers); + } - String decodedString = new String(decodedValue, Charset.forName("UTF-8")); - - String userName = StringUtils.substringBefore(decodedString, ":"); - String password = StringUtils.substringAfter(decodedString, ":"); - if (!authenticator.authenticate(userName, password)) { - writeAuthenticationRequired(authenticator.getRealm()); - return true; - } - clientDetails.setUserName(userName); - - LOG.debug("Got proxy authorization!"); - // We need to remove the header before sending the request on. - String authentication = request.headers().get( - HttpHeaderNames.PROXY_AUTHORIZATION); - LOG.debug(authentication); - request.headers().remove(HttpHeaderNames.PROXY_AUTHORIZATION); - authenticated.set(true); - return false; + ProxyUtils.addVia(httpRequest, proxyServer.getProxyAlias()); } - - private void writeAuthenticationRequired(String realm) { - String body = "\n" - + "\n" - + "407 Proxy Authentication Required\n" - + "\n" - + "

Proxy Authentication Required

\n" - + "

This server could not verify that you\n" - + "are authorized to access the document\n" - + "requested. Either you supplied the wrong\n" - + "credentials (e.g., bad password), or your\n" - + "browser doesn't understand how to supply\n" - + "the credentials required.

\n" + "\n"; - FullHttpResponse response = ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, - HttpResponseStatus.PROXY_AUTHENTICATION_REQUIRED, body); - response.headers().set(HttpHeaderNames.DATE, new Date()); - response.headers().set(HttpHeaderNames.PROXY_AUTHENTICATE, - "Basic realm=\"" + (realm == null ? "Restricted Files" : realm) + "\""); - write(response); - } - - /* ************************************************************************* - * Request/Response Rewriting - **************************************************************************/ - - /** - * Copy the given {@link HttpRequest} verbatim. - */ - private HttpRequest copy(HttpRequest original) { - if (original instanceof FullHttpRequest) { - return ((FullHttpRequest) original).copy(); - } else { - HttpRequest request = new DefaultHttpRequest(original.protocolVersion(), - original.method(), original.uri()); - request.headers().set(original.headers()); - return request; - } + } + + /** + * Checks if we should preserve Proxy-Authorization headers for upstream proxy authentication. + * + * @return true if we're forwarding to an upstream proxy that requires authentication + */ + boolean shouldPreserveProxyAuthorizationForUpstream() { + if (!currentServerConnection.hasUpstreamChainedProxy()) { + return false; } - /** - * Chunked encoding is an HTTP 1.1 feature, but sometimes we get a chunked - * response that reports its HTTP version as 1.0. In this case, we change it - * to 1.1. - */ - private void fixHttpVersionHeaderIfNecessary(HttpResponse httpResponse) { - String te = httpResponse.headers().get( - HttpHeaderNames.TRANSFER_ENCODING); - if (StringUtils.isNotBlank(te) - && te.equalsIgnoreCase(HttpHeaderValues.CHUNKED.toString())) { - if (httpResponse.protocolVersion() != HttpVersion.HTTP_1_1) { - LOG.debug("Fixing HTTP version."); - httpResponse.setProtocolVersion(HttpVersion.HTTP_1_1); - } - } + ChainedProxy chainedProxy = currentServerConnection.getChainedProxy(); + if (chainedProxy == null) { + return false; } - /** - * If and only if our proxy is not running in transparent mode, modify the - * request headers to reflect that it was proxied. - */ - private void modifyRequestHeadersToReflectProxying(HttpRequest httpRequest) { - if (isNextHopOriginServer()) { - /* - * We are making the request to the origin server, so must modify - * the 'absolute-URI' into the 'origin-form' as per RFC 7230 - * section 5.3.1. - * - * This must happen even for 'transparent' mode, otherwise the origin - * server could infer that the request came via a proxy server. - */ - LOG.debug("Modifying request for proxy chaining"); - // Strip host from uri - String uri = httpRequest.uri(); - String adjustedUri = ProxyUtils.stripHost(uri); - LOG.debug("Stripped host from uri: {} yielding: {}", uri, - adjustedUri); - httpRequest.setUri(adjustedUri); - } - if (!proxyServer.isTransparent()) { - LOG.debug("Modifying request headers for proxying"); - - HttpHeaders headers = httpRequest.headers(); - - // Remove sdch from encodings we accept since we can't decode it. - ProxyUtils.removeSdchEncoding(headers); - switchProxyConnectionHeader(headers); - stripConnectionTokens(headers); - stripHopByHopHeaders(headers); - ProxyUtils.addVia(httpRequest, proxyServer.getProxyAlias()); - } + // Only preserve for HTTP proxies (not SOCKS) + if (chainedProxy.getChainedProxyType() != ChainedProxyType.HTTP) { + return false; } - private boolean isNextHopOriginServer() { - // If there is no upstream chained proxy, the next hop must be the origin server. - if (!currentServerConnection.hasUpstreamChainedProxy()) { - return true; - } - - /* - * Upstream SOCKS proxies are a special case because they do not - * parse or modify the HTTP request in any way. If the upstream - * chained proxy is a SOCKS proxy, we should treat it as if we - * are connecting directly to the origin server. - */ - switch (currentServerConnection.getChainedProxyType()) { - case HTTP: - return false; - case SOCKS4: - case SOCKS5: - return true; - default: - LOG.warn("Assuming upstream chained proxy of unknown type " - + currentServerConnection.getChainedProxyType() - + " should not be treated as an origin server"); - return false; - } + // Check if the upstream proxy requires authentication + return chainedProxy.getUsername() != null && chainedProxy.getPassword() != null; + } + + /** + * Handles upstream proxy 407 (Proxy Authentication Required) responses. This method checks if the + * response is a 407 from an upstream proxy and handles the authentication challenge + * appropriately. + * + * @param httpResponse the response from the upstream proxy + * @return true if this is an upstream proxy 407 that should be handled, false otherwise + */ + boolean handleUpstreamProxyAuthenticationRequired(HttpResponse httpResponse) { + // Check if this is a 407 response + if (httpResponse.status() != HttpResponseStatus.PROXY_AUTHENTICATION_REQUIRED) { + return false; } - /** - * If and only if our proxy is not running in transparent mode, modify the - * response headers to reflect that it was proxied. - */ - private void modifyResponseHeadersToReflectProxying( - HttpResponse httpResponse) { - if (!proxyServer.isTransparent()) { - HttpHeaders headers = httpResponse.headers(); - - stripConnectionTokens(headers); - stripHopByHopHeaders(headers); - ProxyUtils.addVia(httpResponse, proxyServer.getProxyAlias()); - - /* - * RFC2616 Section 14.18 - * - * A received message that does not have a Date header field MUST be - * assigned one by the recipient if the message will be cached by - * that recipient or gatewayed via a protocol which requires a Date. - */ - if (!headers.contains(HttpHeaderNames.DATE)) { - headers.set(HttpHeaderNames.DATE, new Date()); - } - } + // Check if we have an upstream chained proxy + if (!currentServerConnection.hasUpstreamChainedProxy()) { + return false; } - /** - * Switch the de-facto standard "Proxy-Connection" header to "Connection" - * when we pass it along to the remote host. This is largely undocumented - * but seems to be what most browsers and servers expect. - * - * @param headers - * The headers to modify - */ - private void switchProxyConnectionHeader(HttpHeaders headers) { - String proxyConnectionKey = "Proxy-Connection"; - if (headers.contains(proxyConnectionKey)) { - String header = headers.get(proxyConnectionKey); - headers.remove(proxyConnectionKey); - headers.set(HttpHeaderNames.CONNECTION, header); - } + ChainedProxy chainedProxy = currentServerConnection.getChainedProxy(); + if (chainedProxy == null) { + return false; } - /** - * RFC2616 Section 14.10 - * - * HTTP/1.1 proxies MUST parse the Connection header field before a message - * is forwarded and, for each connection-token in this field, remove any - * header field(s) from the message with the same name as the - * connection-token. - * - * @param headers - * The headers to modify - */ - private void stripConnectionTokens(HttpHeaders headers) { - if (headers.contains(HttpHeaderNames.CONNECTION)) { - for (String headerValue : headers.getAll(HttpHeaderNames.CONNECTION)) { - for (String connectionToken : ProxyUtils.splitCommaSeparatedHeaderValues(headerValue)) { - // do not strip out the Transfer-Encoding header if it is specified in the Connection header, since LittleProxy does not - // normally modify the Transfer-Encoding of the message. - if (!HttpHeaderNames.TRANSFER_ENCODING.toString().equals(connectionToken.toLowerCase(Locale.US))) { - headers.remove(connectionToken); - } - } - } - } + // Only handle for HTTP proxies + if (chainedProxy.getChainedProxyType() != ChainedProxyType.HTTP) { + return false; } - /** - * Removes all headers that should not be forwarded. See RFC 2616 13.5.1 - * End-to-end and Hop-by-hop Headers. - * - * @param headers - * The headers to modify - */ - private void stripHopByHopHeaders(HttpHeaders headers) { - Set headerNames = headers.names(); - for (String headerName : headerNames) { - if (ProxyUtils.shouldRemoveHopByHopHeader(headerName)) { - headers.remove(headerName); - } - } + // Check if the upstream proxy requires authentication + if (chainedProxy.getUsername() == null || chainedProxy.getPassword() == null) { + // Upstream proxy doesn't have credentials configured, pass the 407 to client + return false; } - /* ************************************************************************* - * Miscellaneous - **************************************************************************/ - - /** - * Tells the client that something went wrong trying to proxy its request. If the Bad Gateway is a response to - * an HTTP HEAD request, the response will contain no body, but the Content-Length header will be set to the - * value it would have been if this 502 Bad Gateway were in response to a GET. - * - * @param httpRequest the HttpRequest that is resulting in the Bad Gateway response - * @return true if the connection will be kept open, or false if it will be disconnected - */ - private boolean writeBadGateway(HttpRequest httpRequest) { - String body = "Bad Gateway: " + httpRequest.uri(); - FullHttpResponse response = ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.BAD_GATEWAY, body); - - if (ProxyUtils.isHEAD(httpRequest)) { - // don't allow any body content in response to a HEAD request - response.content().clear(); - } - - return respondWithShortCircuitResponse(response); + // This is an upstream proxy 407 that we can handle + LOG.debug("Received 407 from upstream proxy, will retry with authentication"); + + // We need to retry the request with proper authentication + // This would require more complex logic to retry the request + // For now, we'll pass the 407 to the client as the current architecture + // doesn't easily support retrying requests + + return true; + } + + /** + * Adds Proxy-Authorization header for upstream proxy authentication. + * + * @param headers the headers to modify + */ + void addUpstreamProxyAuthorization(HttpHeaders headers) { + ChainedProxy chainedProxy = currentServerConnection.getChainedProxy(); + if (chainedProxy != null) { + String username = chainedProxy.getUsername(); + String password = chainedProxy.getPassword(); + + if (username != null && password != null) { + String credentials = username + ":" + password; + String base64Credentials = Base64.getEncoder().encodeToString(credentials.getBytes(UTF_8)); + String authHeader = "Basic " + base64Credentials; + + headers.set(HttpHeaderNames.PROXY_AUTHORIZATION, authHeader); + } } + } - /** - * Tells the client that the request was malformed or erroneous. If the Bad Request is a response to - * an HTTP HEAD request, the response will contain no body, but the Content-Length header will be set to the - * value it would have been if this Bad Request were in response to a GET. - * - * @return true if the connection will be kept open, or false if it will be disconnected - */ - private boolean writeBadRequest(HttpRequest httpRequest) { - String body = "Bad Request to URI: " + httpRequest.uri(); - FullHttpResponse response = ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.BAD_REQUEST, body); - - if (ProxyUtils.isHEAD(httpRequest)) { - // don't allow any body content in response to a HEAD request - response.content().clear(); - } - - return respondWithShortCircuitResponse(response); + private boolean isNextHopOriginServer() { + // If there is no upstream chained proxy, the next hop must be the origin server. + if (!currentServerConnection.hasUpstreamChainedProxy()) { + return true; } - /** - * Tells the client that the connection to the server, or possibly to some intermediary service (such as DNS), timed out. - * If the Gateway Timeout is a response to an HTTP HEAD request, the response will contain no body, but the - * Content-Length header will be set to the value it would have been if this 504 Gateway Timeout were in response to a GET. - * - * @param httpRequest the HttpRequest that is resulting in the Gateway Timeout response - * @return true if the connection will be kept open, or false if it will be disconnected + /* + * Upstream SOCKS proxies are a special case because they do not + * parse or modify the HTTP request in any way. If the upstream + * chained proxy is a SOCKS proxy, we should treat it as if we + * are connecting directly to the origin server. */ - private boolean writeGatewayTimeout(HttpRequest httpRequest) { - String body = "Gateway Timeout"; - FullHttpResponse response = ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, - HttpResponseStatus.GATEWAY_TIMEOUT, body); - - if (httpRequest != null && ProxyUtils.isHEAD(httpRequest)) { - // don't allow any body content in response to a HEAD request - response.content().clear(); - } - - return respondWithShortCircuitResponse(response); + switch (currentServerConnection.getChainedProxyType()) { + case HTTP: + return false; + case SOCKS4: + case SOCKS5: + return true; + default: + LOG.warn( + "Assuming upstream chained proxy of unknown type " + + currentServerConnection.getChainedProxyType() + + " should not be treated as an origin server"); + return false; + } + } + + /** + * If and only if our proxy is not running in transparent mode, modify the response headers to + * reflect that it was proxied. + */ + private void modifyResponseHeadersToReflectProxying(HttpResponse httpResponse) { + if (!proxyServer.isTransparent()) { + HttpHeaders headers = httpResponse.headers(); + + stripConnectionTokens(headers); + stripHopByHopHeaders(headers); + ProxyUtils.addVia(httpResponse, proxyServer.getProxyAlias()); + + /* + * RFC2616 Section 14.18 + * + * A received message that does not have a Date header field MUST be + * assigned one by the recipient if the message will be cached by + * that recipient or gatewayed via a protocol which requires a Date. + */ + if (!headers.contains(HttpHeaderNames.DATE)) { + headers.set(HttpHeaderNames.DATE, dateHeaderValue()); + } + } + } + + /** + * Switch the de-facto standard "Proxy-Connection" header to "Connection" when we pass it along to + * the remote host. This is largely undocumented but seems to be what most browsers and servers + * expect. + * + * @param headers The headers to modify + */ + private void switchProxyConnectionHeader(HttpHeaders headers) { + String proxyConnectionKey = "Proxy-Connection"; + if (headers.contains(proxyConnectionKey)) { + String header = headers.get(proxyConnectionKey); + headers.remove(proxyConnectionKey); + headers.set(HttpHeaderNames.CONNECTION, header); + } + } + + /** + * RFC2616 Section 14.10 + * + *

HTTP/1.1 proxies MUST parse the Connection header field before a message is forwarded and, + * for each connection-token in this field, remove any header field(s) from the message with the + * same name as the connection-token. + * + * @param headers The headers to modify + */ + private void stripConnectionTokens(HttpHeaders headers) { + if (headers.contains(HttpHeaderNames.CONNECTION)) { + for (String headerValue : headers.getAll(HttpHeaderNames.CONNECTION)) { + for (String connectionToken : ProxyUtils.splitCommaSeparatedHeaderValues(headerValue)) { + // do not strip out the Transfer-Encoding header if it is specified in the Connection + // header, since LittleProxy does not + // normally modify the Transfer-Encoding of the message. + if (!HttpHeaderNames.TRANSFER_ENCODING + .toString() + .equals(connectionToken.toLowerCase(Locale.US))) { + headers.remove(connectionToken); + } + } + } + } + } + + /** + * Removes all headers that should not be forwarded. See RFC 2616 13.5.1 End-to-end and Hop-by-hop + * Headers. + * + * @param headers The headers to modify + */ + void stripHopByHopHeaders(HttpHeaders headers) { + ProxyUtils.stripHopByHopHeaders(headers); + } + + /* ************************************************************************* + * Miscellaneous + **************************************************************************/ + + /** + * Tells the client that something went wrong trying to proxy its request. If the Bad Gateway is a + * response to an HTTP HEAD request, the response will contain no body, but the Content-Length + * header will be set to the value it would have been if this 502 Bad Gateway were in response to + * a GET. + * + * @param httpRequest the HttpRequest that is resulting in the Bad Gateway response + * @return true if the connection will be kept open, or false if it will be disconnected + */ + private boolean writeBadGateway(HttpRequest httpRequest) { + String body = "Bad Gateway: " + httpRequest.uri(); + FullHttpResponse response = + ProxyUtils.createFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.BAD_GATEWAY, body); + + if (ProxyUtils.isHEAD(httpRequest)) { + // don't allow any body content in response to a HEAD request + response.content().clear(); } - /** - * Responds to the client with the specified "short-circuit" response. The response will be sent through the - * {@link HttpFilters#proxyToClientResponse(HttpObject)} filter method before writing it to the client. The client - * will not be disconnected, unless the response includes a "Connection: close" header, or the filter returns - * a null HttpResponse (in which case no response will be written to the client and the connection will be - * disconnected immediately). If the response is not a Bad Gateway or Gateway Timeout response, the response's headers - * will be modified to reflect proxying, including adding a Via header, Date header, etc. - * - * @param httpResponse the response to return to the client - * @return true if the connection will be kept open, or false if it will be disconnected. - */ - private boolean respondWithShortCircuitResponse(HttpResponse httpResponse) { - // we are sending a response to the client, so we are done handling this request - this.currentRequest = null; - - HttpResponse filteredResponse = (HttpResponse) currentFilters.proxyToClientResponse(httpResponse); - if (filteredResponse == null) { - disconnect(); - return false; - } - - // allow short-circuit messages to close the connection. normally the Connection header would be stripped when modifying - // the message for proxying, so save the keep-alive status before the modifications are made. - boolean isKeepAlive = HttpUtil.isKeepAlive(httpResponse); - - // if the response is not a Bad Gateway or Gateway Timeout, modify the headers "as if" the short-circuit response were proxied - int statusCode = httpResponse.status().code(); - if (statusCode != HttpResponseStatus.BAD_GATEWAY.code() && statusCode != HttpResponseStatus.GATEWAY_TIMEOUT.code()) { - modifyResponseHeadersToReflectProxying(httpResponse); - } - - // restore the keep alive status, if it was overwritten when modifying headers for proxying - HttpUtil.setKeepAlive(httpResponse, isKeepAlive); - - write(httpResponse); + return respondWithShortCircuitResponse(response); + } + + /** + * Tells the client that the request was malformed or erroneous. If the Bad Request is a response + * to an HTTP HEAD request, the response will contain no body, but the Content-Length header will + * be set to the value it would have been if this Bad Request were in response to a GET. + * + * @return true if the connection will be kept open, or false if it will be disconnected + */ + private boolean writeBadRequest(HttpRequest httpRequest) { + String body = "Bad Request to URI: " + httpRequest.uri(); + FullHttpResponse response = + ProxyUtils.createFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.BAD_REQUEST, body); + + if (ProxyUtils.isHEAD(httpRequest)) { + // don't allow any body content in response to a HEAD request + response.content().clear(); + } - if (ProxyUtils.isLastChunk(httpResponse)) { - writeEmptyBuffer(); - } + return respondWithShortCircuitResponse(response); + } + + /** + * Tells the client that the connection to the server, or possibly to some intermediary service + * (such as DNS), timed out. If the Gateway Timeout is a response to an HTTP HEAD request, the + * response will contain no body, but the Content-Length header will be set to the value it would + * have been if this 504 Gateway Timeout were in response to a GET. + * + * @param httpRequest the HttpRequest that is resulting in the Gateway Timeout response + * @return true if the connection will be kept open, or false if it will be disconnected + */ + private void writeGatewayTimeout(HttpRequest httpRequest) { + String body = "Gateway Timeout"; + FullHttpResponse response = + ProxyUtils.createFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.GATEWAY_TIMEOUT, body); + + if (ProxyUtils.isHEAD(httpRequest)) { + // don't allow any body content in response to a HEAD request + response.content().clear(); + } - if (!HttpUtil.isKeepAlive(httpResponse)) { - disconnect(); - return false; - } + respondWithShortCircuitResponse(response); + } + + /** + * Responds to the client with the specified "short-circuit" response. The response will be sent + * through the {@link HttpFilters#proxyToClientResponse(HttpObject)} filter method before writing + * it to the client. The client will not be disconnected, unless the response includes a + * "Connection: close" header, or the filter returns a null HttpResponse (in which case no + * response will be written to the client and the connection will be disconnected immediately). If + * the response is not a Bad Gateway or Gateway Timeout response, the response's headers will be + * modified to reflect proxying, including adding a Via header, Date header, etc. + * + * @param httpResponse the response to return to the client + * @return true if the connection will be kept open, or false if it will be disconnected. + */ + private boolean respondWithShortCircuitResponse(HttpResponse httpResponse) { + // we are sending a response to the client, so we are done handling this request + resetCurrentRequest(); + + // allow short-circuit messages to close the connection. normally the Connection header would be + // stripped when modifying + // the message for proxying, so save the keep-alive status before the modifications are made. + boolean isKeepAlive = HttpUtil.isKeepAlive(httpResponse); + + HttpResponse filteredResponse = + (HttpResponse) currentFilters.proxyToClientResponse(httpResponse); + if (filteredResponse == null) { + disconnect(); + return false; + } - return true; + // if the response is not a Bad Gateway or Gateway Timeout, modify the headers "as if" the + // short-circuit response were proxied + int statusCode = filteredResponse.status().code(); + if (statusCode != HttpResponseStatus.BAD_GATEWAY.code() + && statusCode != HttpResponseStatus.GATEWAY_TIMEOUT.code()) { + + // Handle upstream proxy authentication challenges + if (handleUpstreamProxyAuthenticationRequired(filteredResponse)) { + // This is a 407 from upstream proxy that we can handle + // For now, we'll modify the response to indicate we're handling it + // In a more complete implementation, we would retry the request with auth + LOG.debug("Handling upstream proxy 407 response"); + } + + modifyResponseHeadersToReflectProxying(filteredResponse); } - /** - * Identify the host and port for a request. - */ - private String identifyHostAndPort(HttpRequest httpRequest) { - String hostAndPort = ProxyUtils.parseHostAndPort(httpRequest); - if (StringUtils.isBlank(hostAndPort)) { - List hosts = httpRequest.headers().getAll( - HttpHeaderNames.HOST); - if (hosts != null && !hosts.isEmpty()) { - hostAndPort = hosts.get(0); - } - } + // restore the keep alive status, if it was overwritten when modifying headers for proxying + HttpUtil.setKeepAlive(filteredResponse, isKeepAlive); - return hostAndPort; - } - - /** - * Write an empty buffer at the end of a chunked transfer. We need to do - * this to handle the way Netty creates HttpChunks from responses that - * aren't in fact chunked from the remote server using Transfer-Encoding: - * chunked. Netty turns these into pseudo-chunked responses in cases where - * the response would otherwise fill up too much memory or where the length - * of the response body is unknown. This is handy because it means we can - * start streaming response bodies back to the client without reading the - * entire response. The problem is that in these pseudo-cases the last chunk - * is encoded to null, and this thwarts normal ChannelFutures from - * propagating operationComplete events on writes to appropriate channel - * listeners. We work around this by writing an empty buffer in those cases - * and using the empty buffer's future instead to handle any operations we - * need to when responses are fully written back to clients. - */ - private void writeEmptyBuffer() { - write(Unpooled.EMPTY_BUFFER); + write(filteredResponse); + + if (ProxyUtils.isLastChunk(filteredResponse)) { + writeEmptyBuffer(); } - public boolean isMitming() { - return mitming; + if (!HttpUtil.isKeepAlive(filteredResponse)) { + disconnect(); + return false; } - protected void setMitming(boolean isMitming) { - this.mitming = isMitming; + return true; + } + + /** Identify the host and port for a request. */ + @NonNull + @CheckReturnValue + private String identifyHostAndPort(@NonNull HttpRequest httpRequest) { + String hostAndPort = ProxyUtils.parseHostAndPort(httpRequest); + if (StringUtils.isBlank(hostAndPort)) { + List hosts = httpRequest.headers().getAll(HttpHeaderNames.HOST); + if (hosts != null && !hosts.isEmpty()) { + hostAndPort = hosts.get(0); + } } - /* ************************************************************************* - * Activity Tracking/Statistics - * - * We track statistics on bytes, requests and responses by adding handlers - * at the appropriate parts of the pipeline (see initChannelPipeline()). - **************************************************************************/ - private final BytesReadMonitor bytesReadMonitor = new BytesReadMonitor() { + return hostAndPort; + } + + /** + * Write an empty buffer at the end of a chunked transfer. We need to do this to handle the way + * Netty creates HttpChunks from responses that aren't in fact chunked from the remote server + * using Transfer-Encoding: chunked. Netty turns these into pseudo-chunked responses in cases + * where the response would otherwise fill up too much memory or where the length of the response + * body is unknown. This is handy because it means we can start streaming response bodies back to + * the client without reading the entire response. The problem is that in these pseudo-cases the + * last chunk is encoded to null, and this thwarts normal ChannelFutures from propagating + * operationComplete events on writes to appropriate channel listeners. We work around this by + * writing an empty buffer in those cases and using the empty buffer's future instead to handle + * any operations we need to when responses are fully written back to clients. + */ + private void writeEmptyBuffer() { + write(Unpooled.EMPTY_BUFFER); + } + + public boolean isMitming() { + return mitming; + } + + protected void setMitming(boolean isMitming) { + mitming = isMitming; + } + + /* ************************************************************************* + * Activity Tracking/Statistics + * + * We track statistics on bytes, requests and responses by adding handlers + * at the appropriate parts of the pipeline (see initChannelPipeline()). + **************************************************************************/ + private final BytesReadMonitor bytesReadMonitor = + new BytesReadMonitor() { @Override protected void bytesRead(int numberOfBytes) { - FlowContext flowContext = flowContext(); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.bytesReceivedFromClient(flowContext, numberOfBytes); - } + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.bytesReceivedFromClient(flowContext, numberOfBytes); + } } - }; + }; - private RequestReadMonitor requestReadMonitor = new RequestReadMonitor() { + private final RequestReadMonitor requestReadMonitor = + new RequestReadMonitor() { @Override protected void requestRead(HttpRequest httpRequest) { - FlowContext flowContext = flowContext(); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.requestReceivedFromClient(flowContext, httpRequest); - } + recordClientConnected(); + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.requestReceivedFromClient(flowContext, httpRequest); + } } - }; + }; - private BytesWrittenMonitor bytesWrittenMonitor = new BytesWrittenMonitor() { + private final BytesWrittenMonitor bytesWrittenMonitor = + new BytesWrittenMonitor() { @Override protected void bytesWritten(int numberOfBytes) { - FlowContext flowContext = flowContext(); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.bytesSentToClient(flowContext, numberOfBytes); - } + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.bytesSentToClient(flowContext, numberOfBytes); + } } - }; + }; - private ResponseWrittenMonitor responseWrittenMonitor = new ResponseWrittenMonitor() { + private final ResponseWrittenMonitor responseWrittenMonitor = + new ResponseWrittenMonitor() { @Override protected void responseWritten(HttpResponse httpResponse) { - FlowContext flowContext = flowContext(); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.responseSentToClient(flowContext, - httpResponse); - } + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.responseSentToClient(flowContext, httpResponse); + } } - }; - - private void recordClientConnected() { - try { - InetSocketAddress clientAddress = getClientAddress(); - clientDetails.setClientAddress(clientAddress); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.clientConnected(clientAddress); - } - } catch (Exception e) { - LOG.error("Unable to recordClientConnected", e); - } - } + }; - private void recordClientSSLHandshakeSucceeded() { - try { - InetSocketAddress clientAddress = getClientAddress(); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.clientSSLHandshakeSucceeded( - clientAddress, clientSslSession); - } - } catch (Exception e) { - LOG.error("Unable to recorClientSSLHandshakeSucceeded", e); - } + private void recordClientConnected() { + if (!clientConnectedRecorded.compareAndSet( + CLIENT_CONNECTED_NOT_YET_RECORDED, CLIENT_CONNECTED_RECORDED)) { + return; } - - private void recordClientDisconnected() { - try { - InetSocketAddress clientAddress = getClientAddress(); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.clientDisconnected( - clientAddress, clientSslSession); - } - } catch (Exception e) { - LOG.error("Unable to recordClientDisconnected", e); - } + try { + FlowContext flowContext = flowContext(); + // Resolve via FlowContext so ClientDetails (used for chained-proxy routing) sees the real + // client IP, not the TCP peer. + clientDetails.setClientAddress(flowContext.getClientAddress()); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.clientConnected(flowContext); + } + } catch (Exception e) { + LOG.error("Unable to recordClientConnected", e); } - - public InetSocketAddress getClientAddress() { - if (channel == null) { - return null; - } - return (InetSocketAddress) channel.remoteAddress(); + } + + private void recordClientSSLHandshakeStarted() { + try { + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.clientSSLHandshakeStarted(flowContext); + } + } catch (Exception e) { + LOG.error("Unable to recordClientSSLHandshakeStarted", e); } - - private FlowContext flowContext() { - if (currentServerConnection != null) { - return new FullFlowContext(this, currentServerConnection); - } else { - return new FlowContext(this); - } + } + + private void recordClientSSLHandshakeSucceeded() { + try { + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.clientSSLHandshakeSucceeded(flowContext, clientSslSession); + } + } catch (Exception e) { + LOG.error("Unable to recordClientSSLHandshakeSucceeded", e); } - - public HAProxyMessage getHaProxyMessage() { - return haProxyMessage; + } + + private void recordClientDisconnected() { + // Ensure clientConnected was reported before clientDisconnected, even for silent connections + // (guarded). + recordClientConnected(); + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + try { + tracker.clientDisconnected(flowContext, clientSslSession); + } catch (Exception e) { + LOG.error("Unable to recordClientDisconnected", e); + } } - - public ClientDetails getClientDetails() { - return clientDetails; + } + + private void recordConnectionSaturated() { + try { + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.connectionSaturated(flowContext); + } + } catch (Exception e) { + LOG.error("Unable to recordConnectionSaturated", e); } - + } + + private void recordConnectionWritable() { + try { + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.connectionWritable(flowContext); + } + } catch (Exception e) { + LOG.error("Unable to recordConnectionWritable", e); + } + } + + private void recordConnectionTimedOut() { + try { + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.connectionTimedOut(flowContext); + } + } catch (Exception e) { + LOG.error("Unable to recordConnectionTimedOut", e); + } + } + + private void recordConnectionExceptionCaught(Throwable cause) { + try { + FlowContext flowContext = flowContext(); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.connectionExceptionCaught(flowContext, cause); + } + } catch (Exception e) { + LOG.error("Unable to recordConnectionExceptionCaught", e); + } + } + + @Nullable + public InetSocketAddress getClientAddress() { + return ofNullable(channel) + .map(c -> c.remoteAddress()) + .filter(InetSocketAddress.class::isInstance) + .map(InetSocketAddress.class::cast) + .orElse(null); + } + + FlowContext flowContext() { + FlowContext cached = clientFlowContext; + if (currentServerConnection != null && !(cached instanceof FullFlowContext)) { + cached = flowContextForServerConnection(currentServerConnection); + } + return cached; + } + + FullFlowContext flowContextForServerConnection(ProxyToServerConnection serverConnection) { + return serverFlowContexts.computeIfAbsent( + serverConnection, sc -> new FullFlowContext(this, sc)); + } + + void clearFlowContextForServerConnection(ProxyToServerConnection serverConnection) { + serverFlowContexts.remove(serverConnection); + } + + public @Nullable HAProxyMessage getHaProxyMessage() { + return haProxyMessage; + } + + public ClientDetails getClientDetails() { + return clientDetails; + } + + /** + * Gets the authenticated status of this connection. + * + * @return the authenticated status + */ + public AtomicBoolean getAuthenticated() { + return authenticated; + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ConcurrentMapServerConnectionPool.java b/src/main/java/org/littleshoot/proxy/impl/ConcurrentMapServerConnectionPool.java new file mode 100644 index 00000000..b6ec549f --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/ConcurrentMapServerConnectionPool.java @@ -0,0 +1,474 @@ +package org.littleshoot.proxy.impl; + +import io.netty.channel.Channel; +import io.netty.handler.codec.http.HttpRequest; +import java.net.InetSocketAddress; +import java.time.Duration; +import java.util.Queue; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentLinkedQueue; +import java.util.concurrent.ConcurrentMap; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.ScheduledFuture; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.ChainedProxy; +import org.littleshoot.proxy.ChainedProxyManager; +import org.littleshoot.proxy.HttpFilters; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +public class ConcurrentMapServerConnectionPool implements ServerConnectionPool { + private static final Logger LOG = + LoggerFactory.getLogger(ConcurrentMapServerConnectionPool.class); + + static final int DEFAULT_MAX_CONNECTIONS_PER_HOST = 10; + static final int DEFAULT_MAX_TOTAL_CONNECTIONS = 200; + + private final ConcurrentMap> + connectionsByHostAndPort = new ConcurrentHashMap<>(); + private final ConcurrentMap> availableConnectionsByHostAndPort = + new ConcurrentHashMap<>(); + private final ConcurrentMap connectionCountByHostAndPort = + new ConcurrentHashMap<>(); + private final ConcurrentMap> pendingRequestsByChannel = + new ConcurrentHashMap<>(); + private final AtomicInteger totalConnectionsCreated = new AtomicInteger(0); + + @Nullable private volatile Duration idleTimeout; + private volatile boolean connectionValidationEnabled = false; + private final ScheduledExecutorService evictionScheduler = + Executors.newSingleThreadScheduledExecutor( + r -> { + Thread t = new Thread(r, "connection-pool-eviction"); + t.setDaemon(true); + return t; + }); + @Nullable private volatile ScheduledFuture evictionTask; + + private final int maxConnectionsPerHost; + private final int maxConnections; + private final DefaultHttpProxyServer proxyServer; + private final io.netty.handler.traffic.GlobalTrafficShapingHandler globalTrafficShapingHandler; + + private final ConcurrentMap connectionKeys = + new ConcurrentHashMap<>(); + + private final java.util.concurrent.atomic.AtomicLong borrowCount = + new java.util.concurrent.atomic.AtomicLong(0); + private final java.util.concurrent.atomic.AtomicLong returnCount = + new java.util.concurrent.atomic.AtomicLong(0); + private final java.util.concurrent.atomic.AtomicLong evictionCount = + new java.util.concurrent.atomic.AtomicLong(0); + private final java.util.concurrent.atomic.AtomicLong validationFailureCount = + new java.util.concurrent.atomic.AtomicLong(0); + + public ConcurrentMapServerConnectionPool( + DefaultHttpProxyServer proxyServer, + io.netty.handler.traffic.GlobalTrafficShapingHandler globalTrafficShapingHandler) { + this( + proxyServer, + globalTrafficShapingHandler, + DEFAULT_MAX_CONNECTIONS_PER_HOST, + DEFAULT_MAX_TOTAL_CONNECTIONS); + } + + public ConcurrentMapServerConnectionPool( + DefaultHttpProxyServer proxyServer, + io.netty.handler.traffic.GlobalTrafficShapingHandler globalTrafficShapingHandler, + int maxConnectionsPerHost) { + this( + proxyServer, + globalTrafficShapingHandler, + maxConnectionsPerHost, + DEFAULT_MAX_TOTAL_CONNECTIONS); + } + + public ConcurrentMapServerConnectionPool( + DefaultHttpProxyServer proxyServer, + io.netty.handler.traffic.GlobalTrafficShapingHandler globalTrafficShapingHandler, + int maxConnectionsPerHost, + int maxConnections) { + this.proxyServer = proxyServer; + this.globalTrafficShapingHandler = globalTrafficShapingHandler; + this.maxConnectionsPerHost = + maxConnectionsPerHost > 0 ? maxConnectionsPerHost : DEFAULT_MAX_CONNECTIONS_PER_HOST; + this.maxConnections = maxConnections > 0 ? maxConnections : DEFAULT_MAX_TOTAL_CONNECTIONS; + } + + @Override + @Nullable + public ProxyToServerConnection getOrCreateConnection( + String serverHostAndPort, + @Nullable InetSocketAddress chainedProxyAddress, + ClientToProxyConnection clientConnection, + HttpFilters initialFilters, + HttpRequest initialHttpRequest) { + ChainedProxy chainedProxy = resolveChainedProxy(initialHttpRequest, clientConnection); + String poolKey = computePoolKey(serverHostAndPort, chainedProxyAddress); + ProxyToServerConnection available = borrowAvailableConnection(poolKey); + if (available != null) { + borrowCount.incrementAndGet(); + return available; + } + + int currentHostCount = + connectionCountByHostAndPort.computeIfAbsent(poolKey, k -> new AtomicInteger(0)).get(); + + if (currentHostCount >= maxConnectionsPerHost) { + LOG.warn( + "Per-host connection limit reached for {}: {} connections, max is {}", + poolKey, + currentHostCount, + maxConnectionsPerHost); + return null; + } + + if (totalConnectionsCreated.get() >= maxConnections) { + LOG.warn( + "Pool exhausted: {} connections created, max is {}", + totalConnectionsCreated.get(), + maxConnections); + return null; + } + + synchronized (this) { + ProxyToServerConnection existingAvailable = borrowAvailableConnection(poolKey); + if (existingAvailable != null) { + return existingAvailable; + } + + int hostCount = + connectionCountByHostAndPort.computeIfAbsent(poolKey, k -> new AtomicInteger(0)).get(); + + if (hostCount >= maxConnectionsPerHost) { + LOG.warn( + "Per-host limit reached after sync for {}: {} connections, max is {}", + poolKey, + hostCount, + maxConnectionsPerHost); + return null; + } + + if (totalConnectionsCreated.get() >= maxConnections) { + LOG.warn( + "Pool exhausted after sync: {} connections created, max is {}", + totalConnectionsCreated.get(), + maxConnections); + return null; + } + + try { + ProxyToServerConnection newConnection = + ProxyToServerConnection.createForPool( + proxyServer, + this, + clientConnection, + serverHostAndPort, + chainedProxy, + initialFilters, + initialHttpRequest, + globalTrafficShapingHandler); + + if (newConnection != null) { + connectionsByHostAndPort + .computeIfAbsent(poolKey, k -> new ConcurrentHashMap<>()) + .put(newConnection, Boolean.TRUE); + connectionCountByHostAndPort + .computeIfAbsent(poolKey, k -> new AtomicInteger(0)) + .incrementAndGet(); + connectionKeys.put(newConnection, poolKey); + totalConnectionsCreated.incrementAndGet(); + borrowCount.incrementAndGet(); + return newConnection; + } + } catch (java.net.UnknownHostException e) { + LOG.warn("Failed to resolve host for {}", serverHostAndPort, e); + } + } + return null; + } + + @Override + public void releaseConnection(ProxyToServerConnection connection) { + if (connection == null) { + return; + } + String poolKey = connectionKeys.get(connection); + if (poolKey == null) { + return; + } + ConcurrentMap connections = + connectionsByHostAndPort.get(poolKey); + if (connections == null || !connections.containsKey(connection)) { + return; + } + if (!connection.isConnected()) { + removeConnection(connection); + return; + } + returnCount.incrementAndGet(); + availableConnectionsByHostAndPort + .computeIfAbsent(poolKey, k -> new ConcurrentLinkedQueue<>()) + .add(new PooledConnection(connection, System.currentTimeMillis())); + } + + @Override + public void registerPendingRequest( + Channel channel, + ClientToProxyConnection clientConnection, + HttpRequest request, + HttpFilters filters) { + pendingRequestsByChannel + .computeIfAbsent(channel, k -> new ConcurrentLinkedQueue<>()) + .add(new PendingRequest(clientConnection, request, filters)); + } + + @Override + @Nullable + public PendingRequest removePendingRequest(Channel channel) { + final PendingRequest[] result = new PendingRequest[1]; + pendingRequestsByChannel.computeIfPresent( + channel, + (k, queue) -> { + result[0] = queue.poll(); + return queue.isEmpty() ? null : queue; + }); + return result[0]; + } + + @Override + @Nullable + public PendingRequest peekPendingRequest(Channel channel) { + Queue queue = pendingRequestsByChannel.get(channel); + if (queue == null || queue.isEmpty()) { + return null; + } + return queue.peek(); + } + + @Override + public void drainPendingRequests(Channel channel) { + pendingRequestsByChannel.remove(channel); + } + + @Override + public void removeConnection(ProxyToServerConnection connection) { + if (connection == null) { + return; + } + String poolKey = connectionKeys.remove(connection); + if (poolKey == null) { + return; + } + ConcurrentMap connections = + connectionsByHostAndPort.get(poolKey); + Queue available = availableConnectionsByHostAndPort.get(poolKey); + if (connections != null) { + if (connections.remove(connection) != null) { + AtomicInteger count = connectionCountByHostAndPort.get(poolKey); + if (count != null) { + int newCount = count.decrementAndGet(); + if (newCount <= 0) { + connectionCountByHostAndPort.remove(poolKey, count); + connectionsByHostAndPort.remove(poolKey, connections); + if (available != null) { + availableConnectionsByHostAndPort.remove(poolKey, available); + } else { + availableConnectionsByHostAndPort.remove(poolKey); + } + } + } + totalConnectionsCreated.decrementAndGet(); + } + } + if (available != null) { + available.removeIf(p -> p.connection == connection); + } + } + + @Override + public void closeAll() { + stopEvictionTask(); + evictionScheduler.shutdown(); + for (ConcurrentMap connections : + connectionsByHostAndPort.values()) { + for (ProxyToServerConnection connection : connections.keySet()) { + connection.close(); + } + } + connectionsByHostAndPort.clear(); + availableConnectionsByHostAndPort.clear(); + connectionCountByHostAndPort.clear(); + connectionKeys.clear(); + pendingRequestsByChannel.clear(); + totalConnectionsCreated.set(0); + } + + @Override + public int getMaxConnectionsPerHost() { + return maxConnectionsPerHost; + } + + @Override + public int getMaxConnections() { + return maxConnections; + } + + @Override + public void setIdleTimeout(@Nullable Duration idleTimeout) { + this.idleTimeout = idleTimeout; + if (idleTimeout != null && idleTimeout.toMillis() > 0) { + startEvictionTask(); + } else { + stopEvictionTask(); + } + } + + @Override + @Nullable + public Duration getIdleTimeout() { + return idleTimeout; + } + + @Override + public void setConnectionValidationEnabled(boolean validationEnabled) { + this.connectionValidationEnabled = validationEnabled; + LOG.info("Connection validation enabled: {}", validationEnabled); + } + + @Override + public boolean isConnectionValidationEnabled() { + return connectionValidationEnabled; + } + + @Override + public PoolMetrics getMetrics() { + int total = totalConnectionsCreated.get(); + int idle = availableConnectionsByHostAndPort.values().stream().mapToInt(q -> q.size()).sum(); + return new PoolMetrics( + total, + total - idle, + idle, + borrowCount.get(), + returnCount.get(), + evictionCount.get(), + validationFailureCount.get()); + } + + private void startEvictionTask() { + if (evictionTask != null && !evictionTask.isCancelled()) { + return; + } + long intervalMillis = idleTimeout != null ? idleTimeout.toMillis() / 2 : 30_000; + evictionTask = + evictionScheduler.scheduleAtFixedRate( + this::evictIdleConnections, intervalMillis, intervalMillis, TimeUnit.MILLISECONDS); + LOG.info("Started idle connection eviction task with interval {}ms", intervalMillis); + } + + private void stopEvictionTask() { + if (evictionTask != null) { + evictionTask.cancel(false); + evictionTask = null; + LOG.info("Stopped idle connection eviction task"); + } + } + + private void evictIdleConnections() { + if (idleTimeout == null || idleTimeout.toMillis() <= 0) { + return; + } + long now = System.currentTimeMillis(); + long idleThreshold = now - idleTimeout.toMillis(); + int evicted = 0; + + for (String serverHostAndPort : availableConnectionsByHostAndPort.keySet()) { + Queue queue = availableConnectionsByHostAndPort.get(serverHostAndPort); + if (queue == null) { + continue; + } + Queue toRemove = new ConcurrentLinkedQueue<>(); + for (PooledConnection pooled : queue) { + if (pooled.releasedAt < idleThreshold) { + toRemove.add(pooled); + evicted++; + } + } + for (PooledConnection pooled : toRemove) { + queue.remove(pooled); + removeConnection(pooled.connection); + pooled.connection.close(); + } + } + if (evicted > 0) { + evictionCount.addAndGet(evicted); + LOG.debug("Evicted {} idle connections", evicted); + } + } + + @Nullable + private ProxyToServerConnection borrowAvailableConnection(String poolKey) { + Queue queue = availableConnectionsByHostAndPort.get(poolKey); + if (queue == null || queue.isEmpty()) { + return null; + } + int checked = 0; + int size = queue.size(); + while (checked < size) { + PooledConnection pooled = queue.poll(); + if (pooled == null) { + return null; + } + checked++; + if (!pooled.connection.isConnected()) { + removeConnection(pooled.connection); + continue; + } + if (connectionValidationEnabled && !isConnectionValid(pooled.connection)) { + validationFailureCount.incrementAndGet(); + removeConnection(pooled.connection); + pooled.connection.close(); + LOG.debug("Connection validation failed, removing connection to {}", poolKey); + continue; + } + if (pooled.connection.isAvailableForNewRequest()) { + return pooled.connection; + } + queue.add(pooled); + } + return null; + } + + private boolean isConnectionValid(ProxyToServerConnection connection) { + return connection.isConnected() && connection.isAvailableForNewRequest(); + } + + @Nullable + private ChainedProxy resolveChainedProxy( + HttpRequest httpRequest, ClientToProxyConnection clientConnection) { + ChainedProxyManager chainedProxyManager = proxyServer.getChainProxyManager(); + if (chainedProxyManager == null) { + return null; + } + Queue chainedProxies = new ConcurrentLinkedQueue<>(); + chainedProxyManager.lookupChainedProxies( + httpRequest, chainedProxies, clientConnection.getClientDetails()); + if (chainedProxies.isEmpty()) { + return null; + } + return chainedProxies.poll(); + } + + private static class PooledConnection { + final ProxyToServerConnection connection; + final long releasedAt; + + PooledConnection(ProxyToServerConnection connection, long releasedAt) { + this.connection = connection; + this.releasedAt = releasedAt; + } + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/ConnectionFlow.java b/src/main/java/org/littleshoot/proxy/impl/ConnectionFlow.java index b13295d5..efbe8807 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ConnectionFlow.java +++ b/src/main/java/org/littleshoot/proxy/impl/ConnectionFlow.java @@ -1,231 +1,182 @@ package org.littleshoot.proxy.impl; -import io.netty.handler.codec.haproxy.HAProxyCommand; -import io.netty.handler.codec.haproxy.HAProxyMessage; -import io.netty.handler.codec.haproxy.HAProxyProtocolVersion; -import io.netty.handler.codec.haproxy.HAProxyProxiedProtocol; import io.netty.util.concurrent.Future; -import io.netty.util.concurrent.GenericFutureListener; -import org.littleshoot.proxy.extras.ProxyProtocolMessage; - import java.util.Deque; import java.util.concurrent.ConcurrentLinkedDeque; -import java.net.InetSocketAddress; /** - * Coordinates the various steps involved in establishing a connection, such as - * establishing a socket connection, SSL handshaking, HTTP CONNECT request - * processing, and so on. + * Coordinates the various steps involved in establishing a connection, such as establishing a + * socket connection, SSL handshaking, HTTP CONNECT request processing, and so on. */ class ConnectionFlow { - private Deque steps = new ConcurrentLinkedDeque(); - - private final ClientToProxyConnection clientConnection; - private final ProxyToServerConnection serverConnection; - private volatile ConnectionFlowStep currentStep; - private volatile boolean suppressInitialRequest = false; - private final Object connectLock; - - /** - * Construct a new {@link ConnectionFlow} for the given client and server - * connections. - * - * @param clientConnection - * @param serverConnection - * @param connectLock - * an object that's shared by {@link ConnectionFlow} and - * {@link ProxyToServerConnection} and that is used for - * synchronizing the reader and writer threads that are both - * involved during the establishing of a connection. - */ - ConnectionFlow( - ClientToProxyConnection clientConnection, - ProxyToServerConnection serverConnection, - Object connectLock) { - super(); - this.clientConnection = clientConnection; - this.serverConnection = serverConnection; - this.connectLock = connectLock; + private final Deque> steps = new ConcurrentLinkedDeque<>(); + + private final ClientToProxyConnection clientConnection; + private final ProxyToServerConnection serverConnection; + private volatile ConnectionFlowStep currentStep; + private volatile boolean suppressInitialRequest; + private final Object connectLock; + + /** + * Construct a new {@link ConnectionFlow} for the given client and server connections. + * + * @param clientConnection + * @param serverConnection + * @param connectLock an object that's shared by {@link ConnectionFlow} and {@link + * ProxyToServerConnection} and that is used for synchronizing the reader and writer threads + * that are both involved during the establishing of a connection. + */ + ConnectionFlow( + ClientToProxyConnection clientConnection, + ProxyToServerConnection serverConnection, + Object connectLock) { + super(); + this.clientConnection = clientConnection; + this.serverConnection = serverConnection; + this.connectLock = connectLock; + } + + /** Add a {@link ConnectionFlowStep} to the beginning of this flow. */ + ConnectionFlow first(ConnectionFlowStep step) { + steps.addFirst(step); + return this; + } + + /** Add a {@link ConnectionFlowStep} to the end of this flow. */ + ConnectionFlow then(ConnectionFlowStep step) { + steps.addLast(step); + return this; + } + + /** + * While we're in the process of connecting, any messages read by the {@link + * ProxyToServerConnection} are passed to this method, which passes it on to {@link + * ConnectionFlowStep#read(ConnectionFlow, Object)} for the current {@link ConnectionFlowStep}. + */ + void read(Object msg) { + if (currentStep != null) { + currentStep.read(this, msg); } - - /** - * Add a {@link ConnectionFlowStep} to the beginning of this flow. - */ - ConnectionFlow first(ConnectionFlowStep step) { - steps.addFirst(step); - return this; + } + + /** + * Starts the connection flow, notifying the {@link ClientToProxyConnection} that we've started. + */ + void start() { + clientConnection.serverConnectionFlowStarted(serverConnection); + advance(); + } + + /** + * Advances the flow. {@link #advance()} will be called until we're either out of steps, or a step + * has failed. + */ + void advance() { + currentStep = steps.poll(); + if (currentStep == null) { + succeed(); + } else { + processCurrentStep(); } - - /** - * Add a {@link ConnectionFlowStep} to the end of this flow. - */ - ConnectionFlow then(ConnectionFlowStep step) { - steps.addLast(step); - return this; + } + + /** + * Process the current {@link ConnectionFlowStep}. With each step, we: + * + *

    + *
  1. Change the state of the associated {@link ProxyConnection} to the value of {@link + * ConnectionFlowStep#getState()} + *
  2. Call {@link ConnectionFlowStep#execute()} + *
  3. On completion of the {@link Future} returned by {@link ConnectionFlowStep#execute()}, + * check the success. + *
  4. If successful, we call back into {@link ConnectionFlowStep#onSuccess(ConnectionFlow)}. + *
  5. If unsuccessful, we call {@link #fail()}, stopping the connection flow + *
+ */ + private void processCurrentStep() { + final ProxyConnection connection = currentStep.getConnection(); + final ProxyConnectionLogger LOG = connection.getLOG(); + + LOG.debug("Processing connection flow step: {}", currentStep); + connection.become(currentStep.getState()); + suppressInitialRequest = suppressInitialRequest || currentStep.shouldSuppressInitialRequest(); + + if (currentStep.shouldExecuteOnEventLoop()) { + connection.ctx.executor().submit(() -> doProcessCurrentStep(LOG)); + } else { + doProcessCurrentStep(LOG); } - - /** - * While we're in the process of connecting, any messages read by the - * {@link ProxyToServerConnection} are passed to this method, which passes - * it on to {@link ConnectionFlowStep#read(ConnectionFlow, Object)} for the - * current {@link ConnectionFlowStep}. - */ - void read(Object msg) { - if (this.currentStep != null) { - this.currentStep.read(this, msg); - } + } + + /** + * Does the work of processing the current step, checking the result and handling success/failure. + */ + private void doProcessCurrentStep(final ProxyConnectionLogger LOG) { + currentStep + .execute() + .addListener( + future -> { + synchronized (connectLock) { + if (future.isSuccess()) { + LOG.debug("ConnectionFlowStep succeeded"); + currentStep.onSuccess(ConnectionFlow.this); + } else { + LOG.debug("ConnectionFlowStep failed", future.cause()); + fail(future.cause()); + } + } + }); + } + + /** + * Called when the flow is complete and successful. Notifies the {@link ProxyToServerConnection} + * that we succeeded. + */ + void succeed() { + synchronized (connectLock) { + serverConnection.getLOG().debug("Connection flow completed successfully: {}", currentStep); + serverConnection.connectionSucceeded(!suppressInitialRequest); + notifyThreadsWaitingForConnection(); } - - /** - * Starts the connection flow, notifying the {@link ClientToProxyConnection} - * that we've started. - */ - void start() { - clientConnection.serverConnectionFlowStarted(serverConnection); - advance(); - } - - /** - *

- * Advances the flow. {@link #advance()} will be called until we're either - * out of steps, or a step has failed. - *

- */ - void advance() { - currentStep = steps.poll(); - if (currentStep == null) { - succeed(); - } else { - processCurrentStep(); - } - } - - /** - *

- * Process the current {@link ConnectionFlowStep}. With each step, we: - *

- * - *
    - *
  1. Change the state of the associated {@link ProxyConnection} to the - * value of {@link ConnectionFlowStep#getState()}
  2. - *
  3. Call {@link ConnectionFlowStep#execute()}
  4. - *
  5. On completion of the {@link Future} returned by - * {@link ConnectionFlowStep#execute()}, check the success.
  6. - *
  7. If successful, we call back into - * {@link ConnectionFlowStep#onSuccess(ConnectionFlow)}.
  8. - *
  9. If unsuccessful, we call {@link #fail()}, stopping the connection - * flow
  10. - *
- */ - private void processCurrentStep() { - final ProxyConnection connection = currentStep.getConnection(); - final ProxyConnectionLogger LOG = connection.getLOG(); - - LOG.debug("Processing connection flow step: {}", currentStep); - connection.become(currentStep.getState()); - suppressInitialRequest = suppressInitialRequest - || currentStep.shouldSuppressInitialRequest(); - - if (currentStep.shouldExecuteOnEventLoop()) { - connection.ctx.executor().submit(() -> doProcessCurrentStep(LOG)); - } else { - doProcessCurrentStep(LOG); - } - } - - /** - * Does the work of processing the current step, checking the result and - * handling success/failure. - */ - @SuppressWarnings("unchecked") - private void doProcessCurrentStep(final ProxyConnectionLogger LOG) { - currentStep.execute().addListener( - future -> { - synchronized (connectLock) { - if (future.isSuccess()) { - LOG.debug("ConnectionFlowStep succeeded"); - currentStep.onSuccess(ConnectionFlow.this); - } else { - LOG.debug("ConnectionFlowStep failed", - future.cause()); - fail(future.cause()); - } - } - }); - } - - /** - * Called when the flow is complete and successful. Notifies the - * {@link ProxyToServerConnection} that we succeeded. - */ - void succeed() { - synchronized (connectLock) { - serverConnection.getLOG().debug( - "Connection flow completed successfully: {}", currentStep); - serverConnection.connectionSucceeded(!suppressInitialRequest); - relayProxyInformation(); - notifyThreadsWaitingForConnection(); - } - } - - private void relayProxyInformation() { - if (clientConnection.isSendProxyProtocol()) { - ProxyProtocolMessage proxyProtocolMessage = getHAProxyMessage(clientConnection.getClientAddress(), serverConnection.getRemoteAddress()); - if ( proxyProtocolMessage != null ){ - serverConnection.writeToChannel(proxyProtocolMessage); - } - } - } - - private ProxyProtocolMessage getHAProxyMessage(InetSocketAddress clientAddress, InetSocketAddress remoteAddress) { - HAProxyMessage haProxyMessage = clientConnection.getHaProxyMessage(); - if ( haProxyMessage != null ){ - return new ProxyProtocolMessage(haProxyMessage); - } - return new ProxyProtocolMessage(HAProxyProtocolVersion.V1, HAProxyCommand.PROXY, HAProxyProxiedProtocol.TCP4, clientAddress.getAddress().getHostAddress(), remoteAddress.getAddress().getHostAddress(), clientAddress.getPort(), remoteAddress.getPort()); - } - - /** - * Called when the flow fails at some {@link ConnectionFlowStep}. - * Disconnects the {@link ProxyToServerConnection} and informs the - * {@link ClientToProxyConnection} that our connection failed. - */ - @SuppressWarnings("unchecked") - void fail(final Throwable cause) { - final ConnectionState lastStateBeforeFailure = serverConnection - .getCurrentState(); - serverConnection.disconnect().addListener( - (GenericFutureListener) future -> { - synchronized (connectLock) { - if (!clientConnection.serverConnectionFailed( - serverConnection, - lastStateBeforeFailure, - cause)) { - // the connection to the server failed and we are not retrying, so transition to the - // DISCONNECTED state - serverConnection.become(ConnectionState.DISCONNECTED); - - // We are not retrying our connection, let anyone waiting for a connection know that we're done - notifyThreadsWaitingForConnection(); - } - } - }); - } - - /** - * Like {@link #fail(Throwable)} but with no cause. - */ - void fail() { - fail(null); - } - - /** - * Once we've finished recording our connection and written our initial - * request, we can notify anyone who is waiting on the connection that it's - * okay to proceed. - */ - private void notifyThreadsWaitingForConnection() { - connectLock.notifyAll(); - } - + } + + /** + * Called when the flow fails at some {@link ConnectionFlowStep}. Disconnects the {@link + * ProxyToServerConnection} and informs the {@link ClientToProxyConnection} that our connection + * failed. + */ + void fail(final Throwable cause) { + final ConnectionState lastStateBeforeFailure = serverConnection.getCurrentState(); + serverConnection + .disconnect() + .addListener( + future -> { + synchronized (connectLock) { + if (!clientConnection.serverConnectionFailed( + serverConnection, lastStateBeforeFailure, cause)) { + // the connection to the server failed, and we are not retrying, so transition to + // the + // DISCONNECTED state + serverConnection.become(ConnectionState.DISCONNECTED); + + // We are not retrying our connection, let anyone waiting for a connection know + // that we're done + notifyThreadsWaitingForConnection(); + } + } + }); + } + + /** Like {@link #fail(Throwable)} but with no cause. */ + void fail() { + fail(null); + } + + /** + * Once we've finished recording our connection and written our initial request, we can notify + * anyone who is waiting on the connection that it's okay to proceed. + */ + private void notifyThreadsWaitingForConnection() { + connectLock.notifyAll(); + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ConnectionFlowStep.java b/src/main/java/org/littleshoot/proxy/impl/ConnectionFlowStep.java index a112727f..2226be7b 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ConnectionFlowStep.java +++ b/src/main/java/org/littleshoot/proxy/impl/ConnectionFlowStep.java @@ -1,107 +1,81 @@ package org.littleshoot.proxy.impl; +import io.netty.handler.codec.http.HttpObject; import io.netty.util.concurrent.Future; -/** - * Represents a phase in a {@link ConnectionFlow}. - */ -abstract class ConnectionFlowStep { - private final ProxyConnectionLogger LOG; - private final ProxyConnection connection; - private final ConnectionState state; +/** Represents a phase in a {@link ConnectionFlow}. */ +abstract class ConnectionFlowStep { + private final ProxyConnectionLogger LOG; + private final ProxyConnection connection; + private final ConnectionState state; - /** - * Construct a new step in a connection flow. - * - * @param connection - * the connection that we're working on - * @param state - * the state that the connection will show while we're processing - * this step - */ - ConnectionFlowStep(ProxyConnection connection, - ConnectionState state) { - super(); - this.connection = connection; - this.state = state; - this.LOG = connection.getLOG(); - } + /** + * Construct a new step in a connection flow. + * + * @param connection the connection that we're working on + * @param state the state that the connection will show while we're processing this step + */ + ConnectionFlowStep(ProxyConnection connection, ConnectionState state) { + this.connection = connection; + this.state = state; + LOG = connection.getLOG(); + } - ProxyConnection getConnection() { - return connection; - } + ProxyConnection getConnection() { + return connection; + } - ConnectionState getState() { - return state; - } + ConnectionState getState() { + return state; + } - /** - * Indicates whether or not to suppress the initial request. Defaults to - * false, can be overridden. - */ - boolean shouldSuppressInitialRequest() { - return false; - } + /** Indicates whether to suppress the initial request. Defaults to false, can be overridden. */ + boolean shouldSuppressInitialRequest() { + return false; + } - /** - *

- * Indicates whether or not this step should be executed on the channel's - * event loop. Defaults to true, can be overridden. - *

- * - *

- * If this step modifies the pipeline, for example by adding/removing - * handlers, it's best to make it execute on the event loop. - *

- */ - boolean shouldExecuteOnEventLoop() { - return true; - } + /** + * Indicates whether this step should be executed on the channel's event loop. Defaults to true, + * can be overridden. + * + *

If this step modifies the pipeline, for example by adding/removing handlers, it's best to + * make it execute on the event loop. + */ + boolean shouldExecuteOnEventLoop() { + return true; + } - /** - * Implement this method to actually do the work involved in this step of - * the flow. - */ - protected abstract Future execute(); + /** Implement this method to actually do the work involved in this step of the flow. */ + protected abstract Future execute(); - /** - * When the flow determines that this step was successful, it calls into - * this method. The default implementation simply continues with the flow. - * Other implementations may choose to not continue and instead wait for a - * message or something like that. - */ - void onSuccess(ConnectionFlow flow) { - flow.advance(); - } + /** + * When the flow determines that this step was successful, it calls into this method. The default + * implementation simply continues with the flow. Other implementations may choose to not continue + * and instead wait for a message or something like that. + */ + void onSuccess(ConnectionFlow flow) { + flow.advance(); + } - /** - *

- * Any messages that are read from the underlying connection while we're at - * this step of the connection flow are passed to this method. - *

- * - *

- * The default implementation ignores the message and logs this, since we - * weren't really expecting a message here. - *

- * - *

- * Some {@link ConnectionFlowStep}s do need to read the messages, so they - * override this method as appropriate. - *

- * - * @param flow - * our {@link ConnectionFlow} - * @param msg - * the message read from the underlying connection - */ - void read(ConnectionFlow flow, Object msg) { - LOG.debug("Received message while in the middle of connecting: {}", msg); - } - - @Override - public String toString() { - return state.toString(); - } + /** + * Any messages that are read from the underlying connection while we're at this step of the + * connection flow are passed to this method. + * + *

The default implementation ignores the message and logs this, since we weren't really + * expecting a message here. + * + *

Some {@link ConnectionFlowStep}s do need to read the messages, so they override this method + * as appropriate. + * + * @param flow our {@link ConnectionFlow} + * @param msg the message read from the underlying connection + */ + void read(ConnectionFlow flow, Object msg) { + LOG.debug("Received message while in the middle of connecting: {}", msg); + } + @Override + public String toString() { + return state.toString(); + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ConnectionState.java b/src/main/java/org/littleshoot/proxy/impl/ConnectionState.java index 371fcd58..ef658b0a 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ConnectionState.java +++ b/src/main/java/org/littleshoot/proxy/impl/ConnectionState.java @@ -1,81 +1,65 @@ package org.littleshoot.proxy.impl; enum ConnectionState { - /** - * Connection attempting to connect. - */ - CONNECTING(true), + /** Connection attempting to connect. */ + CONNECTING(true), - /** - * In the middle of doing an SSL handshake. - */ - HANDSHAKING(true), + /** In the middle of doing an SSL handshake. */ + HANDSHAKING(true), - /** - * In the process of negotiating an HTTP CONNECT from the client. - */ - NEGOTIATING_CONNECT(true), + /** In the process of negotiating an HTTP CONNECT from the client. */ + NEGOTIATING_CONNECT(true), - /** - * When forwarding a CONNECT to a chained proxy, we await the CONNECTION_OK - * message from the proxy. - */ - AWAITING_CONNECT_OK(true), + /** + * When forwarding a CONNECT to a chained proxy, we await the CONNECTION_OK message from the + * proxy. + */ + AWAITING_CONNECT_OK(true), - /** - * Connected but waiting for proxy authentication. - */ - AWAITING_PROXY_AUTHENTICATION, + /** Connected but waiting for proxy authentication. */ + AWAITING_PROXY_AUTHENTICATION, - /** - * Connected and awaiting initial message (e.g. HttpRequest or - * HttpResponse). - */ - AWAITING_INITIAL, + /** Connected and awaiting initial message (e.g. HttpRequest or HttpResponse). */ + AWAITING_INITIAL, - /** - * Connected and awaiting HttpContent chunk. - */ - AWAITING_CHUNK, + /** Connected and awaiting HttpContent chunk. */ + AWAITING_CHUNK, - /** - * We've asked the client to disconnect, but it hasn't yet. - */ - DISCONNECT_REQUESTED(), + /** We've asked the client to disconnect, but it hasn't yet. */ + DISCONNECT_REQUESTED(), - /** - * Disconnected - */ - DISCONNECTED(); + /** Disconnected */ + DISCONNECTED(); - private final boolean partOfConnectionFlow; + private final boolean partOfConnectionFlow; - ConnectionState(boolean partOfConnectionFlow) { - this.partOfConnectionFlow = partOfConnectionFlow; - } + ConnectionState(boolean partOfConnectionFlow) { + this.partOfConnectionFlow = partOfConnectionFlow; + } - ConnectionState() { - this(false); - } + ConnectionState() { + this(false); + } - /** - * Indicates whether this ConnectionState corresponds to a step in a - * {@link ConnectionFlow}. This is useful to distinguish so that we know - * whether or not we're in the process of establishing a connection. - * - * @return true if part of connection flow, otherwise false - */ - public boolean isPartOfConnectionFlow() { - return partOfConnectionFlow; - } + /** + * Indicates whether this ConnectionState corresponds to a step in a {@link ConnectionFlow}. This + * is useful to distinguish so that we know whether we're in the process of establishing a + * connection. + * + * @return true if part of connection flow, otherwise false + */ + public boolean isPartOfConnectionFlow() { + return partOfConnectionFlow; + } - /** - * Indicates whether this ConnectionState is no longer waiting for messages and is either in the process of disconnecting - * or is already disconnected. - * - * @return true if the connection state is {@link #DISCONNECT_REQUESTED} or {@link #DISCONNECTED}, otherwise false - */ - public boolean isDisconnectingOrDisconnected() { - return this == DISCONNECT_REQUESTED || this == DISCONNECTED; - } + /** + * Indicates whether this ConnectionState is no longer waiting for messages and is either in the + * process of disconnecting or is already disconnected. + * + * @return true if the connection state is {@link #DISCONNECT_REQUESTED} or {@link #DISCONNECTED}, + * otherwise false + */ + public boolean isDisconnectingOrDisconnected() { + return this == DISCONNECT_REQUESTED || this == DISCONNECTED; + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServer.java b/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServer.java index 16253ef5..0dfcf4dd 100644 --- a/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServer.java +++ b/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServer.java @@ -1,912 +1,665 @@ package org.littleshoot.proxy.impl; import io.netty.bootstrap.ServerBootstrap; -import io.netty.channel.*; +import io.netty.channel.Channel; +import io.netty.channel.ChannelFuture; +import io.netty.channel.ChannelInitializer; +import io.netty.channel.EventLoopGroup; import io.netty.channel.group.ChannelGroup; import io.netty.channel.group.ChannelGroupFuture; import io.netty.channel.group.DefaultChannelGroup; import io.netty.channel.socket.nio.NioServerSocketChannel; -import io.netty.channel.udt.nio.NioUdtProvider; import io.netty.handler.traffic.GlobalTrafficShapingHandler; import io.netty.util.concurrent.GlobalEventExecutor; -import org.littleshoot.proxy.*; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - -import javax.net.ssl.SSLEngine; import java.io.File; import java.io.FileInputStream; import java.io.IOException; import java.io.InputStream; import java.net.InetSocketAddress; +import java.time.Duration; import java.util.Collection; import java.util.Properties; import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.*; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; /** - *

* Primary implementation of an {@link HttpProxyServer}. - *

* - *

- * {@link DefaultHttpProxyServer} is bootstrapped by calling - * {@link #bootstrap()} or {@link #bootstrapFromFile(String)}, and then calling - * {@link DefaultHttpProxyServerBootstrap#start()}. For example: - *

+ *

{@link DefaultHttpProxyServer} is bootstrapped by calling {@link #bootstrap()} or {@link + * #bootstrapFromFile(String)}, and then calling {@link + * org.littleshoot.proxy.impl.DefaultHttpProxyServerBootstrap#start()}. For example: * *

- * DefaultHttpProxyServer server =
- *         DefaultHttpProxyServer
- *                 .bootstrap()
- *                 .withPort(8090)
- *                 .start();
+ * DefaultHttpProxyServer server = DefaultHttpProxyServer
+ *         .bootstrap()
+ *         .withPort(8090)
+ *         .start();
  * 
- * */ public class DefaultHttpProxyServer implements HttpProxyServer { - private static final Logger LOG = LoggerFactory.getLogger(DefaultHttpProxyServer.class); - - /** - * The interval in ms at which the GlobalTrafficShapingHandler will run to compute and throttle the - * proxy-to-server bandwidth. - */ - private static final long TRAFFIC_SHAPING_CHECK_INTERVAL_MS = 250L; - - private static final int MAX_INITIAL_LINE_LENGTH_DEFAULT = 8192; - private static final int MAX_HEADER_SIZE_DEFAULT = 8192*2; - private static final int MAX_CHUNK_SIZE_DEFAULT = 8192*2; - - /** - * The proxy alias to use in the Via header if no explicit proxy alias is specified and the hostname of the local - * machine cannot be resolved. - */ - private static final String FALLBACK_PROXY_ALIAS = "littleproxy"; - - /** - * Our {@link ServerGroup}. Multiple proxy servers can share the same - * ServerGroup in order to reuse threads and other such resources. - */ - private final ServerGroup serverGroup; - - private final TransportProtocol transportProtocol; - /* - * The address that the server will attempt to bind to. - */ - private final InetSocketAddress requestedAddress; - /* - * The actual address to which the server is bound. May be different from the requestedAddress in some circumstances, - * for example when the requested port is 0. - */ - private volatile InetSocketAddress localAddress; - private volatile InetSocketAddress boundAddress; - private final SslEngineSource sslEngineSource; - private final boolean authenticateSslClients; - private final ProxyAuthenticator proxyAuthenticator; - private final ChainedProxyManager chainProxyManager; - private final MitmManager mitmManager; - private final HttpFiltersSource filtersSource; - private final boolean transparent; - private volatile int connectTimeout; - private volatile int idleConnectionTimeout; - private final HostResolver serverResolver; - private volatile GlobalTrafficShapingHandler globalTrafficShapingHandler; - private final int maxInitialLineLength; - private final int maxHeaderSize; - private final int maxChunkSize; - private final boolean allowRequestsToOriginServer; - private final boolean acceptProxyProtocol; - private final boolean sendProxyProtocol; - - /** - * The alias or pseudonym for this proxy, used when adding the Via header. - */ - private final String proxyAlias; - - /** - * True when the proxy has already been stopped by calling {@link #stop()} or {@link #abort()}. - */ - private final AtomicBoolean stopped = new AtomicBoolean(false); - - /** - * Track all ActivityTrackers for tracking proxying activity. - */ - private final Collection activityTrackers = new ConcurrentLinkedQueue<>(); - - /** - * Keep track of all channels created by this proxy server for later shutdown when the proxy is stopped. - */ - private final ChannelGroup allChannels = new DefaultChannelGroup("HTTP-Proxy-Server", GlobalEventExecutor.INSTANCE); - - /** - * JVM shutdown hook to shutdown this proxy server. Declared as a class-level variable to allow removing the shutdown hook when the - * proxy server is stopped normally. - */ - private final Thread jvmShutdownHook = new Thread(this::abort, "LittleProxy-JVM-shutdown-hook"); - - /** - * Bootstrap a new {@link DefaultHttpProxyServer} starting from scratch. - */ - public static HttpProxyServerBootstrap bootstrap() { - return new DefaultHttpProxyServerBootstrap(); - } - - /** - * Bootstrap a new {@link DefaultHttpProxyServer} using defaults from the - * given file. - */ - public static HttpProxyServerBootstrap bootstrapFromFile(String path) { - final File propsFile = new File(path); - Properties props = new Properties(); - - if (propsFile.isFile()) { - try (InputStream is = new FileInputStream(propsFile)) { - props.load(is); - } catch (final IOException e) { - LOG.warn("Could not load props file?", e); - } - } - - return new DefaultHttpProxyServerBootstrap(props); - } - - /** - * Creates a new proxy server. - * - * @param serverGroup - * our ServerGroup for shared thread pools and such - * @param transportProtocol - * The protocol to use for data transport - * @param requestedAddress - * The address on which this server will listen - * @param sslEngineSource - * (optional) if specified, this Proxy will encrypt inbound - * connections from clients using an {@link SSLEngine} obtained - * from this {@link SslEngineSource}. - * @param authenticateSslClients - * Indicate whether or not to authenticate clients when using SSL - * @param proxyAuthenticator - * (optional) If specified, requests to the proxy will be - * authenticated using HTTP BASIC authentication per the provided - * {@link ProxyAuthenticator} - * @param chainProxyManager - * The proxy to send requests to if chaining proxies. Typically - * null. - * @param mitmManager - * The {@link MitmManager} to use for man in the middle'ing - * CONNECT requests - * @param filtersSource - * Source for {@link HttpFilters} - * @param transparent - * If true, this proxy will run as a transparent proxy. This will - * not modify the response, and will only modify the request to - * amend the URI if the target is the origin server (to comply - * with RFC 7230 section 5.3.1). - * @param idleConnectionTimeout - * The timeout (in seconds) for auto-closing idle connections. - * @param activityTrackers - * for tracking activity on this proxy - * @param connectTimeout - * number of milliseconds to wait to connect to the upstream - * server - * @param serverResolver - * the {@link HostResolver} to use for resolving server addresses - * @param readThrottleBytesPerSecond - * read throttle bandwidth - * @param writeThrottleBytesPerSecond - * write throttle bandwidth - * @param maxInitialLineLength - * @param maxHeaderSize - * @param maxChunkSize - * @param allowRequestsToOriginServer - * when true, allow the proxy to handle requests that contain an origin-form URI, as defined in RFC 7230 5.3.1 - * @param acceptProxyProtocol when true, the proxy will accept a proxy protocol header from client - * @param sendProxyProtocol when true, the proxy will send a proxy protocol header to the server - */ - private DefaultHttpProxyServer(ServerGroup serverGroup, - TransportProtocol transportProtocol, - InetSocketAddress requestedAddress, - SslEngineSource sslEngineSource, - boolean authenticateSslClients, - ProxyAuthenticator proxyAuthenticator, - ChainedProxyManager chainProxyManager, - MitmManager mitmManager, - HttpFiltersSource filtersSource, - boolean transparent, - int idleConnectionTimeout, - Collection activityTrackers, - int connectTimeout, - HostResolver serverResolver, - long readThrottleBytesPerSecond, - long writeThrottleBytesPerSecond, - InetSocketAddress localAddress, - String proxyAlias, - int maxInitialLineLength, - int maxHeaderSize, - int maxChunkSize, - boolean allowRequestsToOriginServer, - boolean acceptProxyProtocol, - boolean sendProxyProtocol) { - this.serverGroup = serverGroup; - this.transportProtocol = transportProtocol; - this.requestedAddress = requestedAddress; - this.sslEngineSource = sslEngineSource; - this.authenticateSslClients = authenticateSslClients; - this.proxyAuthenticator = proxyAuthenticator; - this.chainProxyManager = chainProxyManager; - this.mitmManager = mitmManager; - this.filtersSource = filtersSource; - this.transparent = transparent; - this.idleConnectionTimeout = idleConnectionTimeout; - if (activityTrackers != null) { - this.activityTrackers.addAll(activityTrackers); - } - this.connectTimeout = connectTimeout; - this.serverResolver = serverResolver; - - if (writeThrottleBytesPerSecond > 0 || readThrottleBytesPerSecond > 0) { - this.globalTrafficShapingHandler = createGlobalTrafficShapingHandler(transportProtocol, readThrottleBytesPerSecond, writeThrottleBytesPerSecond); - } else { - this.globalTrafficShapingHandler = null; - } - this.localAddress = localAddress; - - if (proxyAlias == null) { - // attempt to resolve the name of the local machine. if it cannot be resolved, use the fallback name. - String hostname = ProxyUtils.getHostName(); - if (hostname == null) { - hostname = FALLBACK_PROXY_ALIAS; - } - this.proxyAlias = hostname; - } else { - this.proxyAlias = proxyAlias; - } - this.maxInitialLineLength = maxInitialLineLength; - this.maxHeaderSize = maxHeaderSize; - this.maxChunkSize = maxChunkSize; - this.allowRequestsToOriginServer = allowRequestsToOriginServer; - this.acceptProxyProtocol = acceptProxyProtocol; - this.sendProxyProtocol = sendProxyProtocol; - } - - /** - * Creates a new GlobalTrafficShapingHandler for this HttpProxyServer, using this proxy's proxyToServerEventLoop. - */ - private GlobalTrafficShapingHandler createGlobalTrafficShapingHandler(TransportProtocol transportProtocol, long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond) { - EventLoopGroup proxyToServerEventLoop = this.getProxyToServerWorkerFor(transportProtocol); - return new GlobalTrafficShapingHandler(proxyToServerEventLoop, - writeThrottleBytesPerSecond, - readThrottleBytesPerSecond, - TRAFFIC_SHAPING_CHECK_INTERVAL_MS, - Long.MAX_VALUE); - } - - boolean isTransparent() { - return transparent; - } - - @Override - public int getIdleConnectionTimeout() { - return idleConnectionTimeout; - } - - @Override - public void setIdleConnectionTimeout(int idleConnectionTimeout) { - this.idleConnectionTimeout = idleConnectionTimeout; - } - - @Override - public int getConnectTimeout() { - return connectTimeout; - } - - @Override - public void setConnectTimeout(int connectTimeoutMs) { - this.connectTimeout = connectTimeoutMs; - } - - public HostResolver getServerResolver() { - return serverResolver; - } - - public InetSocketAddress getLocalAddress() { - return localAddress; - } - - @Override - public InetSocketAddress getListenAddress() { - return boundAddress; - } - - @Override - public void setThrottle(long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond) { - if (globalTrafficShapingHandler != null) { - globalTrafficShapingHandler.configure(writeThrottleBytesPerSecond, readThrottleBytesPerSecond); - } else { - // don't create a GlobalTrafficShapingHandler if throttling was not enabled and is still not enabled - if (readThrottleBytesPerSecond > 0 || writeThrottleBytesPerSecond > 0) { - globalTrafficShapingHandler = createGlobalTrafficShapingHandler(transportProtocol, readThrottleBytesPerSecond, writeThrottleBytesPerSecond); - } - } - } - - public long getReadThrottle() { - return globalTrafficShapingHandler.getReadLimit(); - } - - public long getWriteThrottle() { - return globalTrafficShapingHandler.getWriteLimit(); - } - - public int getMaxInitialLineLength() { - return maxInitialLineLength; - } - - public int getMaxHeaderSize() { - return maxHeaderSize; - } - - public int getMaxChunkSize() { - return maxChunkSize; - } - - public boolean isAllowRequestsToOriginServer() { - return allowRequestsToOriginServer; - } - - public boolean isAcceptProxyProtocol() { - return acceptProxyProtocol; - } - - public boolean isSendProxyProtocol() { - return sendProxyProtocol; - } - - @Override - public HttpProxyServerBootstrap clone() { - return new DefaultHttpProxyServerBootstrap(serverGroup, - transportProtocol, - new InetSocketAddress(requestedAddress.getAddress(), - requestedAddress.getPort() == 0 ? 0 : requestedAddress.getPort() + 1), - sslEngineSource, - authenticateSslClients, - proxyAuthenticator, - chainProxyManager, - mitmManager, - filtersSource, - transparent, - idleConnectionTimeout, - activityTrackers, - connectTimeout, - serverResolver, - globalTrafficShapingHandler != null ? globalTrafficShapingHandler.getReadLimit() : 0, - globalTrafficShapingHandler != null ? globalTrafficShapingHandler.getWriteLimit() : 0, - localAddress, - proxyAlias, - maxInitialLineLength, - maxHeaderSize, - maxChunkSize, - allowRequestsToOriginServer); - } - - @Override - public void stop() { - doStop(true); - } - - @Override - public void abort() { - doStop(false); - } - - /** - * Performs cleanup necessary to stop the server. Closes all channels opened by the server and unregisters this - * server from the server group. - * - * @param graceful when true, waits for requests to terminate before stopping the server - */ - protected void doStop(boolean graceful) { - // only stop the server if it hasn't already been stopped - if (stopped.compareAndSet(false, true)) { - if (graceful) { - LOG.info("Shutting down proxy server gracefully"); - } else { - LOG.info("Shutting down proxy server immediately (non-graceful)"); - } - - closeAllChannels(graceful); - - serverGroup.unregisterProxyServer(this, graceful); - - // remove the shutdown hook that was added when the proxy was started, since it has now been stopped - try { - Runtime.getRuntime().removeShutdownHook(jvmShutdownHook); - } catch (IllegalStateException e) { - // ignore -- IllegalStateException means the VM is already shutting down - } - - LOG.info("Done shutting down proxy server"); - } - } - - /** - * Register a new {@link Channel} with this server, for later closing. - */ - protected void registerChannel(Channel channel) { - allChannels.add(channel); - } - - /** - * Closes all channels opened by this proxy server. - * - * @param graceful when false, attempts to shutdown all channels immediately and ignores any channel-closing exceptions - */ - protected void closeAllChannels(boolean graceful) { - LOG.info("Closing all channels " + (graceful ? "(graceful)" : "(non-graceful)")); - - ChannelGroupFuture future = allChannels.close(); - - // if this is a graceful shutdown, log any channel closing failures. if this isn't a graceful shutdown, ignore them. - if (graceful) { - try { - future.await(10, TimeUnit.SECONDS); - } catch (InterruptedException e) { - Thread.currentThread().interrupt(); - - LOG.warn("Interrupted while waiting for channels to shut down gracefully."); - } - - if (!future.isSuccess()) { - for (ChannelFuture cf : future) { - if (!cf.isSuccess()) { - LOG.info("Unable to close channel. Cause of failure for {} is {}", cf.channel(), cf.cause()); - } - } - } - } - } - - private HttpProxyServer start() { - if (!serverGroup.isStopped()) { - LOG.info("Starting proxy at address: " + this.requestedAddress); - - serverGroup.registerProxyServer(this); - - doStart(); - } else { - throw new IllegalStateException("Attempted to start proxy, but proxy's server group is already stopped"); - } - - return this; - } - - private void doStart() { - ServerBootstrap serverBootstrap = new ServerBootstrap().group( + private static final Logger LOG = LoggerFactory.getLogger(DefaultHttpProxyServer.class); + + /** + * The interval in ms at which the GlobalTrafficShapingHandler will run to compute and throttle + * the proxy-to-server bandwidth. + */ + private static final long TRAFFIC_SHAPING_CHECK_INTERVAL_MS = 250L; + + private static final int MAX_INITIAL_LINE_LENGTH_DEFAULT = 8192; + private static final int MAX_HEADER_SIZE_DEFAULT = 8192 * 2; + private static final int MAX_CHUNK_SIZE_DEFAULT = 8192 * 2; + + /** + * The proxy alias to use in the Via header if no explicit proxy alias is specified and the + * hostname of the local machine cannot be resolved. + */ + private static final String FALLBACK_PROXY_ALIAS = "littleproxy"; + + private static final String DEFAULT_LITTLE_PROXY_NAME = "LittleProxy"; + public static final String LOCAL_ADDRESS = "127.0.0.1"; + public static final int DEFAULT_PORT = 8080; + public static final String DEFAULT_NIC_VALUE = "0.0.0.0"; + public static final String CLIENT_TO_PROXY_WORKER_THREADS = "client_to_proxy_worker_threads"; + public static final String PROXY_TO_SERVER_WORKER_THREADS = "proxy_to_server_worker_threads"; + public static final String ACTIVITY_LOG_FORMAT = "activity_log_format"; + public static final String ACCEPTOR_THREADS = "acceptor_threads"; + public static final String SEND_PROXY_PROTOCOL = "send_proxy_protocol"; + public static final String ALLOW_PROXY_PROTOCOL = "allow_proxy_protocol"; + public static final String SERVER_CONNECTION_POOL_TYPE = "server_connection_pool_type"; + public static final String USE_SHARED_SERVER_CONNECTION_POOL = + "use_shared_server_connection_pool"; + public static final String MAX_TOTAL_CONNECTIONS = "max_total_connections"; + public static final String MAX_CONNECTIONS_PER_HOST = "max_connections_per_host"; + public static final String POOL_SHARED_MITM_CONNECTIONS = "pool_shared_mitm_connections"; + public static final String POOL_PER_REQUEST_IN_MITM = "pool_per_request_in_mitm"; + public static final String ALLOW_REQUESTS_TO_ORIGIN_SERVER = "allow_requests_to_origin_server"; + public static final String THROTTLE_WRITE_BYTES_PER_SECOND = "throttle_write_bytes_per_second"; + public static final String THROTTLE_READ_BYTES_PER_SECOND = "throttle_read_bytes_per_second"; + public static final String TRANSPARENT = "transparent"; + public static final String SSL_CLIENTS_KEYSTORE_PATH = "ssl_clients_keystore_path"; + public static final String SSL_CLIENTS_KEYSTORE_PASSWORD = "ssl_clients_keystore_password"; + public static final String SSL_CLIENTS_KEYSTORE_ALIAS = "ssl_clients_keystore_alias"; + public static final String SSL_CLIENTS_SEND_CERTS = "ssl_clients_send_certs"; + public static final String AUTHENTICATE_SSL_CLIENTS = "authenticate_ssl_clients"; + public static final String SSL_CLIENTS_TRUST_ALL_SERVERS = "ssl_clients_trust_all_servers"; + public static final String ALLOW_LOCAL_ONLY = "allow_local_only"; + public static final String PROXY_ALIAS = "proxy_alias"; + public static final String NIC = "nic"; + public static final String PORT = "port"; + public static final String ADDRESS = "address"; + public static final String NAME = "name"; + private static final String DEFAULT_JKS_KEYSTORE_PATH = "littleproxy_keystore.jks"; + + /** + * Our {@link ServerGroup}. Multiple proxy servers can share the same ServerGroup in order to + * reuse threads and other such resources. + */ + private final ServerGroup serverGroup; + + private final TransportProtocol transportProtocol; + /* + * The address that the server will attempt to bind to. + */ + private final InetSocketAddress requestedAddress; + /* + * The actual address to which the server is bound. May be different from the + * requestedAddress in some circumstances, + * for example when the requested port is 0. + */ + private final InetSocketAddress localAddress; + private volatile InetSocketAddress boundAddress; + private final SslEngineSource sslEngineSource; + private final boolean authenticateSslClients; + private final ProxyAuthenticator proxyAuthenticator; + private final ChainedProxyManager chainProxyManager; + private final MitmManager mitmManager; + private final HttpFiltersSource filtersSource; + private final boolean transparent; + private volatile int connectTimeout; + private volatile Duration idleConnectionTimeout; + private final HostResolver serverResolver; + private volatile GlobalTrafficShapingHandler globalTrafficShapingHandler; + private final int maxInitialLineLength; + private final int maxHeaderSize; + private final int maxChunkSize; + private final boolean allowRequestsToOriginServer; + private final boolean acceptProxyProtocol; + private final boolean sendProxyProtocol; + + /** The alias or pseudonym for this proxy, used when adding the Via header. */ + private final String proxyAlias; + + /** + * True when the proxy has already been stopped by calling {@link #stop()} or {@link #abort()}. + */ + private final AtomicBoolean stopped = new AtomicBoolean(false); + + /** Track all ActivityTrackers for tracking proxying activity. */ + private final Collection activityTrackers = new ConcurrentLinkedQueue<>(); + + /** + * Keep track of all channels created by this proxy server for later shutdown when the proxy is + * stopped. + */ + private final ChannelGroup allChannels = + new DefaultChannelGroup("HTTP-Proxy-Server", GlobalEventExecutor.INSTANCE, true); + + /** + * Shared pool of ProxyToServerConnection instances for all ClientToProxyConnection. This + * addresses the connection explosion issue (GitHub issue #83). + */ + private volatile ServerConnectionPool serverConnectionPool; + + /** + * Whether to use the shared server connection pool. Disabled by default for backwards + * compatibility. + */ + private final boolean useSharedServerConnectionPool; + + /** + * Maximum number of connections per host:port when using the shared connection pool. Default is + * 10. + */ + private final int maxConnectionsPerHost; + + /** Maximum total connections in the shared pool. */ + private final int maxConnections; + + /** Selected server connection pool implementation. */ + private final ServerConnectionPoolType serverConnectionPoolType; + + /** Configuration for the server connection pool. */ + private final ServerConnectionPoolConfig serverConnectionPoolConfig; + + /** + * JVM shutdown hook to shut down this proxy server. Declared as a class-level variable to allow + * removing the shutdown hook when the proxy server is stopped normally. + */ + private final Thread jvmShutdownHook = new Thread(this::abort, "LittleProxy-JVM-shutdown-hook"); + + /** Bootstrap a new {@link DefaultHttpProxyServer} starting from scratch. */ + public static HttpProxyServerBootstrap bootstrap() { + return new org.littleshoot.proxy.impl.DefaultHttpProxyServerBootstrap(); + } + + /** Bootstrap a new {@link DefaultHttpProxyServer} using defaults from the given file. */ + public static HttpProxyServerBootstrap bootstrapFromFile(String path) { + final File propsFile = new File(path); + Properties props = new Properties(); + + if (propsFile.isFile()) { + try (InputStream is = new FileInputStream(propsFile)) { + props.load(is); + } catch (final IOException e) { + LOG.error("Could not load props file", e); + throw new IllegalArgumentException("Could not load props file." + e.getMessage()); + } + } else { + String cause = !propsFile.exists() ? "absent" : "a directory"; + LOG.error("Could not load props file. file is {}", cause); + throw new IllegalArgumentException("Could not load props file. file is " + (cause)); + } + + return new org.littleshoot.proxy.impl.DefaultHttpProxyServerBootstrap(props); + } + + DefaultHttpProxyServer(ServerGroup serverGroup, DefaultHttpProxyServerConfig config) { + this.serverGroup = serverGroup; + this.transportProtocol = config.getTransportProtocol(); + this.requestedAddress = config.getRequestedAddress(); + this.sslEngineSource = config.getSslEngineSource(); + this.authenticateSslClients = config.isAuthenticateSslClients(); + this.proxyAuthenticator = config.getProxyAuthenticator(); + this.chainProxyManager = config.getChainProxyManager(); + this.mitmManager = config.getMitmManager(); + this.filtersSource = config.getFiltersSource(); + this.transparent = config.isTransparent(); + this.idleConnectionTimeout = config.getIdleConnectionTimeout(); + this.activityTrackers.addAll(config.getActivityTrackers()); + this.connectTimeout = config.getConnectTimeout(); + this.serverResolver = config.getServerResolver(); + + long readThrottleBytesPerSecond = config.getReadThrottleBytesPerSecond(); + long writeThrottleBytesPerSecond = config.getWriteThrottleBytesPerSecond(); + if (writeThrottleBytesPerSecond > 0 || readThrottleBytesPerSecond > 0) { + globalTrafficShapingHandler = + createGlobalTrafficShapingHandler( + config.getTransportProtocol(), + readThrottleBytesPerSecond, + writeThrottleBytesPerSecond); + } else { + globalTrafficShapingHandler = null; + } + this.localAddress = config.getLocalAddress(); + + String proxyAlias = config.getProxyAlias(); + if (proxyAlias == null) { + // attempt to resolve the name of the local machine. if it cannot be resolved, + // use the fallback name. + String hostname = ProxyUtils.getHostName(); + if (hostname == null) { + hostname = FALLBACK_PROXY_ALIAS; + } + this.proxyAlias = hostname; + } else { + this.proxyAlias = proxyAlias; + } + this.maxInitialLineLength = config.getMaxInitialLineLength(); + this.maxHeaderSize = config.getMaxHeaderSize(); + this.maxChunkSize = config.getMaxChunkSize(); + this.allowRequestsToOriginServer = config.isAllowRequestsToOriginServer(); + this.acceptProxyProtocol = config.isAcceptProxyProtocol(); + this.sendProxyProtocol = config.isSendProxyProtocol(); + this.serverConnectionPoolConfig = config.getServerConnectionPoolConfig(); + this.useSharedServerConnectionPool = this.serverConnectionPoolConfig.isEnabled(); + this.maxConnectionsPerHost = this.serverConnectionPoolConfig.getMaxConnectionsPerHost(); + this.serverConnectionPoolType = this.serverConnectionPoolConfig.getPoolType(); + this.maxConnections = this.serverConnectionPoolConfig.getMaxConnections(); + } + + /** + * Creates a new GlobalTrafficShapingHandler for this HttpProxyServer, using this proxy's + * proxyToServerEventLoop. + */ + private GlobalTrafficShapingHandler createGlobalTrafficShapingHandler( + TransportProtocol transportProtocol, + long readThrottleBytesPerSecond, + long writeThrottleBytesPerSecond) { + EventLoopGroup proxyToServerEventLoop = getProxyToServerWorkerFor(transportProtocol); + return new GlobalTrafficShapingHandler( + proxyToServerEventLoop, + writeThrottleBytesPerSecond, + readThrottleBytesPerSecond, + TRAFFIC_SHAPING_CHECK_INTERVAL_MS, + Long.MAX_VALUE); + } + + boolean isTransparent() { + return transparent; + } + + @Override + public int getIdleConnectionTimeout() { + return (int) idleConnectionTimeout.toSeconds(); + } + + @Override + public void setIdleConnectionTimeout(int idleConnectionTimeoutInSeconds) { + this.idleConnectionTimeout = Duration.ofSeconds(idleConnectionTimeoutInSeconds); + } + + @Override + public void setIdleConnectionTimeout(Duration idleConnectionTimeout) { + this.idleConnectionTimeout = idleConnectionTimeout; + } + + @Override + public int getConnectTimeout() { + return connectTimeout; + } + + @Override + public void setConnectTimeout(int connectTimeoutMs) { + connectTimeout = connectTimeoutMs; + } + + public HostResolver getServerResolver() { + return serverResolver; + } + + public InetSocketAddress getLocalAddress() { + return localAddress; + } + + @Override + public InetSocketAddress getListenAddress() { + return boundAddress; + } + + @Override + public void setThrottle(long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond) { + if (globalTrafficShapingHandler != null) { + globalTrafficShapingHandler.configure( + writeThrottleBytesPerSecond, readThrottleBytesPerSecond); + } else { + // don't create a GlobalTrafficShapingHandler if throttling was not enabled and + // is still not enabled + if (readThrottleBytesPerSecond > 0 || writeThrottleBytesPerSecond > 0) { + globalTrafficShapingHandler = + createGlobalTrafficShapingHandler( + transportProtocol, readThrottleBytesPerSecond, writeThrottleBytesPerSecond); + } + } + } + + public long getReadThrottle() { + if (globalTrafficShapingHandler != null) { + return globalTrafficShapingHandler.getReadLimit(); + } else { + return 0; + } + } + + public long getWriteThrottle() { + if (globalTrafficShapingHandler != null) { + return globalTrafficShapingHandler.getWriteLimit(); + } else { + return 0; + } + } + + public int getMaxInitialLineLength() { + return maxInitialLineLength; + } + + public int getMaxHeaderSize() { + return maxHeaderSize; + } + + public int getMaxChunkSize() { + return maxChunkSize; + } + + /** + * Gets the shared ServerConnectionPool for this server. Creates the pool if it doesn't exist. + * + * @return the shared pool, or null if the pool is disabled + */ + @Nullable + public ServerConnectionPool getServerConnectionPool() { + if (!useSharedServerConnectionPool) { + return null; + } + ServerConnectionPool pool = serverConnectionPool; + if (pool == null) { + synchronized (this) { + pool = serverConnectionPool; + if (pool == null) { + pool = createServerConnectionPool(); + serverConnectionPool = pool; + } + } + } + return pool; + } + + private ServerConnectionPool createServerConnectionPool() { + ServerConnectionPoolType poolType = serverConnectionPoolConfig.getPoolType(); + Duration idleTimeout = serverConnectionPoolConfig.getIdleTimeout(); + int maxConnPerHost = serverConnectionPoolConfig.getMaxConnectionsPerHost(); + int maxConn = serverConnectionPoolConfig.getMaxConnections(); + + switch (poolType) { + case CONCURRENT_MAP: + default: + ConcurrentMapServerConnectionPool concurrentMapPool = + new ConcurrentMapServerConnectionPool( + this, globalTrafficShapingHandler, maxConnPerHost, maxConn); + concurrentMapPool.setIdleTimeout(idleTimeout); + return concurrentMapPool; + } + } + + public boolean isPoolSharedMitmConnections() { + return serverConnectionPoolConfig.isPoolSharedMitmConnections(); + } + + public boolean isPoolPerRequestInMitm() { + return serverConnectionPoolConfig.isPoolPerRequestInMitm(); + } + + public boolean isAllowRequestsToOriginServer() { + return allowRequestsToOriginServer; + } + + public boolean isAcceptProxyProtocol() { + return acceptProxyProtocol; + } + + public boolean isSendProxyProtocol() { + return sendProxyProtocol; + } + + @Override + public HttpProxyServerBootstrap clone() { + InetSocketAddress clonedAddress = + new InetSocketAddress( + requestedAddress.getAddress(), + requestedAddress.getPort() == 0 ? 0 : requestedAddress.getPort() + 1); + + ServerConnectionPoolConfig poolConfig = + new ServerConnectionPoolConfig() + .setEnabled(useSharedServerConnectionPool) + .setPoolType(serverConnectionPoolType) + .setMaxConnectionsPerHost(maxConnectionsPerHost) + .setMaxConnections(maxConnections) + .setIdleTimeout(serverConnectionPoolConfig.getIdleTimeout()) + .setPoolSharedMitmConnections(serverConnectionPoolConfig.isPoolSharedMitmConnections()) + .setPoolPerRequestInMitm(serverConnectionPoolConfig.isPoolPerRequestInMitm()); + + DefaultHttpProxyServerConfig serverConfig = + new DefaultHttpProxyServerConfig() + .setTransportProtocol(transportProtocol) + .setRequestedAddress(clonedAddress) + .setSslEngineSource(sslEngineSource) + .setAuthenticateSslClients(authenticateSslClients) + .setProxyAuthenticator(proxyAuthenticator) + .setChainProxyManager(chainProxyManager) + .setMitmManager(mitmManager) + .setFiltersSource(filtersSource) + .setTransparent(transparent) + .setIdleConnectionTimeout(idleConnectionTimeout) + .setActivityTrackers(activityTrackers) + .setConnectTimeout(connectTimeout) + .setServerResolver(serverResolver) + .setReadThrottleBytesPerSecond( + globalTrafficShapingHandler != null + ? globalTrafficShapingHandler.getReadLimit() + : 0) + .setWriteThrottleBytesPerSecond( + globalTrafficShapingHandler != null + ? globalTrafficShapingHandler.getWriteLimit() + : 0) + .setLocalAddress(localAddress) + .setProxyAlias(proxyAlias) + .setMaxInitialLineLength(maxInitialLineLength) + .setMaxHeaderSize(maxHeaderSize) + .setMaxChunkSize(maxChunkSize) + .setAllowRequestsToOriginServer(allowRequestsToOriginServer) + .setAcceptProxyProtocol(acceptProxyProtocol) + .setSendProxyProtocol(sendProxyProtocol) + .setServerConnectionPoolConfig(poolConfig); + + return new org.littleshoot.proxy.impl.DefaultHttpProxyServerBootstrap( + serverGroup, serverConfig); + } + + @Override + public void stop() { + doStop(true); + } + + @Override + public void abort() { + doStop(false); + } + + /** + * Performs cleanup necessary to stop the server. Closes all channels opened by the server and + * unregisters this server from the server group. + * + * @param graceful when true, waits for requests to terminate before stopping the server + */ + protected void doStop(boolean graceful) { + // only stop the server if it hasn't already been stopped + if (stopped.compareAndSet(false, true)) { + if (graceful) { + LOG.info("Shutting down proxy server gracefully"); + } else { + LOG.info("Shutting down proxy server immediately (non-graceful)"); + } + + // Close the shared server connection pool + if (serverConnectionPool != null) { + serverConnectionPool.closeAll(); + } + + closeAllChannels(graceful); + + serverGroup.unregisterProxyServer(this, graceful); + + // remove the shutdown hook that was added when the proxy was started, since it + // has now been stopped + try { + Runtime.getRuntime().removeShutdownHook(jvmShutdownHook); + } catch (IllegalStateException e) { + // ignore -- IllegalStateException means the VM is already shutting down + } + + LOG.info("Done shutting down proxy server"); + } + } + + /** Register a new {@link Channel} with this server, for later closing. */ + protected void registerChannel(Channel channel) { + allChannels.add(channel); + } + + protected void unregisterChannel(Channel channel) { + if (channel.isOpen()) { + // Unlikely to happen, but just in case... + channel.close(); + } + allChannels.remove(channel); + } + + /** + * Closes all channels opened by this proxy server. + * + * @param graceful when false, attempts to shut down all channels immediately and ignores any + * channel-closing exceptions + */ + protected void closeAllChannels(boolean graceful) { + LOG.info("Closing all channels {}", graceful ? "(graceful)" : "(non-graceful)"); + + ChannelGroupFuture future = allChannels.close(); + + // if this is a graceful shutdown, log any channel closing failures. if this + // isn't a graceful shutdown, ignore them. + if (graceful) { + try { + future.await(10, TimeUnit.SECONDS); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + + LOG.warn("Interrupted while waiting for channels to shut down gracefully."); + } + + if (!future.isSuccess()) { + for (ChannelFuture cf : future) { + if (!cf.isSuccess()) { + LOG.info( + "Unable to close channel. Cause of failure for {} is {}", + cf.channel(), + String.valueOf(cf.cause())); + } + } + } + } + } + + HttpProxyServer start() { + if (!serverGroup.isStopped()) { + LOG.info("Starting proxy at address: {}", requestedAddress); + + serverGroup.registerProxyServer(this); + + doStart(); + } else { + throw new IllegalStateException( + "Attempted to start proxy, but proxy's server group is already stopped"); + } + + return this; + } + + private void doStart() { + ServerBootstrap serverBootstrap = + new ServerBootstrap() + .group( serverGroup.getClientToProxyAcceptorPoolForTransport(transportProtocol), serverGroup.getClientToProxyWorkerPoolForTransport(transportProtocol)); - ChannelInitializer initializer = new ChannelInitializer() { - protected void initChannel(Channel ch) { - new ClientToProxyConnection( - DefaultHttpProxyServer.this, - sslEngineSource, - authenticateSslClients, - ch.pipeline(), - globalTrafficShapingHandler); - } + ChannelInitializer initializer = + new ChannelInitializer<>() { + protected void initChannel(Channel ch) { + new ClientToProxyConnection( + DefaultHttpProxyServer.this, + sslEngineSource, + authenticateSslClients, + ch.pipeline(), + globalTrafficShapingHandler); + } }; - switch (transportProtocol) { - case TCP: - LOG.info("Proxy listening with TCP transport"); - serverBootstrap.channelFactory(NioServerSocketChannel::new); - break; - case UDT: - LOG.info("Proxy listening with UDT transport"); - serverBootstrap.channelFactory(NioUdtProvider.BYTE_ACCEPTOR) - .option(ChannelOption.SO_BACKLOG, 10) - .option(ChannelOption.SO_REUSEADDR, true); - break; - default: - throw new UnknownTransportProtocolException(transportProtocol); - } - serverBootstrap.childHandler(initializer); - ChannelFuture future = serverBootstrap.bind(requestedAddress) - .addListener((ChannelFutureListener) future1 -> { - if (future1.isSuccess()) { - registerChannel(future1.channel()); - } - }).awaitUninterruptibly(); - - Throwable cause = future.cause(); - if (cause != null) { - throw new RuntimeException(cause); - } - - this.boundAddress = ((InetSocketAddress) future.channel().localAddress()); - LOG.info("Proxy started at address: " + this.boundAddress); - - Runtime.getRuntime().addShutdownHook(jvmShutdownHook); - } - - protected ChainedProxyManager getChainProxyManager() { - return chainProxyManager; - } - - protected MitmManager getMitmManager() { - return mitmManager; - } - - protected SslEngineSource getSslEngineSource() { - return sslEngineSource; - } - - protected ProxyAuthenticator getProxyAuthenticator() { - return proxyAuthenticator; - } - - public HttpFiltersSource getFiltersSource() { - return filtersSource; - } - - protected Collection getActivityTrackers() { - return activityTrackers; - } - - public String getProxyAlias() { - return proxyAlias; - } - - - protected EventLoopGroup getProxyToServerWorkerFor(TransportProtocol transportProtocol) { - return serverGroup.getProxyToServerWorkerPoolForTransport(transportProtocol); - } - - // TODO: refactor bootstrap into a separate class - private static class DefaultHttpProxyServerBootstrap implements HttpProxyServerBootstrap { - private String name = "LittleProxy"; - private ServerGroup serverGroup = null; - private TransportProtocol transportProtocol = TransportProtocol.TCP; - private InetSocketAddress requestedAddress; - private int port = 8080; - private boolean allowLocalOnly = true; - private SslEngineSource sslEngineSource = null; - private boolean authenticateSslClients = true; - private ProxyAuthenticator proxyAuthenticator = null; - private ChainedProxyManager chainProxyManager = null; - private MitmManager mitmManager = null; - private HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter(); - private boolean transparent = false; - private int idleConnectionTimeout = 70; - private Collection activityTrackers = new ConcurrentLinkedQueue<>(); - private int connectTimeout = 40000; - private HostResolver serverResolver = new DefaultHostResolver(); - private long readThrottleBytesPerSecond; - private long writeThrottleBytesPerSecond; - private InetSocketAddress localAddress; - private String proxyAlias; - private int clientToProxyAcceptorThreads = ServerGroup.DEFAULT_INCOMING_ACCEPTOR_THREADS; - private int clientToProxyWorkerThreads = ServerGroup.DEFAULT_INCOMING_WORKER_THREADS; - private int proxyToServerWorkerThreads = ServerGroup.DEFAULT_OUTGOING_WORKER_THREADS; - private int maxInitialLineLength = MAX_INITIAL_LINE_LENGTH_DEFAULT; - private int maxHeaderSize = MAX_HEADER_SIZE_DEFAULT; - private int maxChunkSize = MAX_CHUNK_SIZE_DEFAULT; - private boolean allowRequestToOriginServer = false; - private boolean acceptProxyProtocol = false; - private boolean sendProxyProtocol = false; - - private DefaultHttpProxyServerBootstrap() { - } - - private DefaultHttpProxyServerBootstrap( - ServerGroup serverGroup, - TransportProtocol transportProtocol, - InetSocketAddress requestedAddress, - SslEngineSource sslEngineSource, - boolean authenticateSslClients, - ProxyAuthenticator proxyAuthenticator, - ChainedProxyManager chainProxyManager, - MitmManager mitmManager, - HttpFiltersSource filtersSource, - boolean transparent, int idleConnectionTimeout, - Collection activityTrackers, - int connectTimeout, HostResolver serverResolver, - long readThrottleBytesPerSecond, - long writeThrottleBytesPerSecond, - InetSocketAddress localAddress, - String proxyAlias, - int maxInitialLineLength, - int maxHeaderSize, - int maxChunkSize, - boolean allowRequestToOriginServer) { - this.serverGroup = serverGroup; - this.transportProtocol = transportProtocol; - this.requestedAddress = requestedAddress; - this.port = requestedAddress.getPort(); - this.sslEngineSource = sslEngineSource; - this.authenticateSslClients = authenticateSslClients; - this.proxyAuthenticator = proxyAuthenticator; - this.chainProxyManager = chainProxyManager; - this.mitmManager = mitmManager; - this.filtersSource = filtersSource; - this.transparent = transparent; - this.idleConnectionTimeout = idleConnectionTimeout; - if (activityTrackers != null) { - this.activityTrackers.addAll(activityTrackers); - } - this.connectTimeout = connectTimeout; - this.serverResolver = serverResolver; - this.readThrottleBytesPerSecond = readThrottleBytesPerSecond; - this.writeThrottleBytesPerSecond = writeThrottleBytesPerSecond; - this.localAddress = localAddress; - this.proxyAlias = proxyAlias; - this.maxInitialLineLength = maxInitialLineLength; - this.maxHeaderSize = maxHeaderSize; - this.maxChunkSize = maxChunkSize; - this.allowRequestToOriginServer = allowRequestToOriginServer; - } - - private DefaultHttpProxyServerBootstrap(Properties props) { - this.withUseDnsSec(ProxyUtils.extractBooleanDefaultFalse( - props, "dnssec")); - this.transparent = ProxyUtils.extractBooleanDefaultFalse( - props, "transparent"); - this.idleConnectionTimeout = ProxyUtils.extractInt(props, - "idle_connection_timeout"); - this.connectTimeout = ProxyUtils.extractInt(props, - "connect_timeout", 0); - this.maxInitialLineLength = ProxyUtils.extractInt(props, - "max_initial_line_length", MAX_INITIAL_LINE_LENGTH_DEFAULT); - this.maxHeaderSize = ProxyUtils.extractInt(props, - "max_header_size", MAX_HEADER_SIZE_DEFAULT); - this.maxChunkSize = ProxyUtils.extractInt(props, - "max_chunk_size", MAX_CHUNK_SIZE_DEFAULT); - } - - @Override - public HttpProxyServerBootstrap withName(String name) { - this.name = name; - return this; - } - - @Override - public HttpProxyServerBootstrap withTransportProtocol( - TransportProtocol transportProtocol) { - this.transportProtocol = transportProtocol; - return this; - } - - @Override - public HttpProxyServerBootstrap withAddress(InetSocketAddress address) { - this.requestedAddress = address; - return this; - } - - @Override - public HttpProxyServerBootstrap withPort(int port) { - this.requestedAddress = null; - this.port = port; - return this; - } - - @Override - public HttpProxyServerBootstrap withNetworkInterface(InetSocketAddress inetSocketAddress) { - this.localAddress = inetSocketAddress; - return this; - } - - @Override - public HttpProxyServerBootstrap withProxyAlias(String alias) { - this.proxyAlias = alias; - return this; - } - - @Override - public HttpProxyServerBootstrap withAllowLocalOnly( - boolean allowLocalOnly) { - this.allowLocalOnly = allowLocalOnly; - return this; - } - - @Override - @Deprecated - public HttpProxyServerBootstrap withListenOnAllAddresses(boolean listenOnAllAddresses) { - LOG.warn("withListenOnAllAddresses() is deprecated and will be removed in a future release. Use withNetworkInterface()."); - return this; - } - - @Override - public HttpProxyServerBootstrap withSslEngineSource( - SslEngineSource sslEngineSource) { - this.sslEngineSource = sslEngineSource; - if (this.mitmManager != null) { - LOG.warn("Enabled encrypted inbound connections with man in the middle. " - + "These are mutually exclusive - man in the middle will be disabled."); - this.mitmManager = null; - } - return this; - } - - @Override - public HttpProxyServerBootstrap withAuthenticateSslClients( - boolean authenticateSslClients) { - this.authenticateSslClients = authenticateSslClients; - return this; - } - - @Override - public HttpProxyServerBootstrap withProxyAuthenticator( - ProxyAuthenticator proxyAuthenticator) { - this.proxyAuthenticator = proxyAuthenticator; - return this; - } - - @Override - public HttpProxyServerBootstrap withChainProxyManager( - ChainedProxyManager chainProxyManager) { - this.chainProxyManager = chainProxyManager; - return this; - } - - @Override - public HttpProxyServerBootstrap withManInTheMiddle( - MitmManager mitmManager) { - this.mitmManager = mitmManager; - if (this.sslEngineSource != null) { - LOG.warn("Enabled man in the middle with encrypted inbound connections. " - + "These are mutually exclusive - encrypted inbound connections will be disabled."); - this.sslEngineSource = null; - } - return this; - } - - @Override - public HttpProxyServerBootstrap withFiltersSource( - HttpFiltersSource filtersSource) { - this.filtersSource = filtersSource; - return this; - } - - @Override - public HttpProxyServerBootstrap withUseDnsSec(boolean useDnsSec) { - if (useDnsSec) { - this.serverResolver = new DnsSecServerResolver(); - } else { - this.serverResolver = new DefaultHostResolver(); - } - return this; - } - - @Override - public HttpProxyServerBootstrap withTransparent( - boolean transparent) { - this.transparent = transparent; - return this; - } - - @Override - public HttpProxyServerBootstrap withIdleConnectionTimeout( - int idleConnectionTimeout) { - this.idleConnectionTimeout = idleConnectionTimeout; - return this; - } - - @Override - public HttpProxyServerBootstrap withConnectTimeout( - int connectTimeout) { - this.connectTimeout = connectTimeout; - return this; - } - - @Override - public HttpProxyServerBootstrap withServerResolver( - HostResolver serverResolver) { - this.serverResolver = serverResolver; - return this; - } - @Override - public HttpProxyServerBootstrap withServerGroup( - ServerGroup group) { - this.serverGroup = group; - return this; - } - - @Override - public HttpProxyServerBootstrap plusActivityTracker( - ActivityTracker activityTracker) { - activityTrackers.add(activityTracker); - return this; - } - - @Override - public HttpProxyServerBootstrap withThrottling(long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond) { - this.readThrottleBytesPerSecond = readThrottleBytesPerSecond; - this.writeThrottleBytesPerSecond = writeThrottleBytesPerSecond; - return this; - } - - @Override - public HttpProxyServerBootstrap withMaxInitialLineLength(int maxInitialLineLength){ - this.maxInitialLineLength = maxInitialLineLength; - return this; - } - - @Override - public HttpProxyServerBootstrap withMaxHeaderSize(int maxHeaderSize){ - this.maxHeaderSize = maxHeaderSize; - return this; - } - - @Override - public HttpProxyServerBootstrap withMaxChunkSize(int maxChunkSize){ - this.maxChunkSize = maxChunkSize; - return this; - } - - @Override - public HttpProxyServerBootstrap withAllowRequestToOriginServer(boolean allowRequestToOriginServer) { - this.allowRequestToOriginServer = allowRequestToOriginServer; - return this; - } - - @Override - public HttpProxyServerBootstrap withAcceptProxyProtocol(boolean acceptProxyProtocol) { - this.acceptProxyProtocol = acceptProxyProtocol; - return this; - } - - @Override - public HttpProxyServerBootstrap withSendProxyProtocol(boolean sendProxyProtocol) { - this.sendProxyProtocol = sendProxyProtocol; - return this; - } - - @Override - public HttpProxyServer start() { - return build().start(); - } - - @Override - public HttpProxyServerBootstrap withThreadPoolConfiguration(ThreadPoolConfiguration configuration) { - this.clientToProxyAcceptorThreads = configuration.getAcceptorThreads(); - this.clientToProxyWorkerThreads = configuration.getClientToProxyWorkerThreads(); - this.proxyToServerWorkerThreads = configuration.getProxyToServerWorkerThreads(); - return this; - } - - private DefaultHttpProxyServer build() { - final ServerGroup serverGroup; - - if (this.serverGroup != null) { - serverGroup = this.serverGroup; - } - else { - serverGroup = new ServerGroup(name, clientToProxyAcceptorThreads, clientToProxyWorkerThreads, proxyToServerWorkerThreads); - } - - return new DefaultHttpProxyServer(serverGroup, - transportProtocol, determineListenAddress(), - sslEngineSource, authenticateSslClients, - proxyAuthenticator, chainProxyManager, mitmManager, - filtersSource, transparent, - idleConnectionTimeout, activityTrackers, connectTimeout, - serverResolver, readThrottleBytesPerSecond, writeThrottleBytesPerSecond, - localAddress, proxyAlias, maxInitialLineLength, maxHeaderSize, maxChunkSize, - allowRequestToOriginServer, acceptProxyProtocol, sendProxyProtocol); - } - - private InetSocketAddress determineListenAddress() { - if (requestedAddress != null) { - return requestedAddress; - } else { - // Binding only to localhost can significantly improve the - // security of the proxy. - if (allowLocalOnly) { - return new InetSocketAddress("127.0.0.1", port); - } else { - return new InetSocketAddress(port); - } - } - } - } + switch (transportProtocol) { + case TCP: + LOG.info("Proxy listening with TCP transport"); + serverBootstrap.channelFactory(NioServerSocketChannel::new); + break; + default: + throw new UnknownTransportProtocolException(transportProtocol); + } + serverBootstrap.childHandler(initializer); + ChannelFuture future = serverBootstrap.bind(requestedAddress).awaitUninterruptibly(); + + Throwable cause = future.cause(); + if (cause != null) { + abort(); + throw new RuntimeException(cause); + } + + Channel serverChannel = future.channel(); + registerChannel(serverChannel); + boundAddress = (InetSocketAddress) serverChannel.localAddress(); + LOG.info("Proxy started at address: {}", boundAddress); + + Runtime.getRuntime().addShutdownHook(jvmShutdownHook); + } + + protected ChainedProxyManager getChainProxyManager() { + return chainProxyManager; + } + + protected MitmManager getMitmManager() { + return mitmManager; + } + + protected SslEngineSource getSslEngineSource() { + return sslEngineSource; + } + + protected ProxyAuthenticator getProxyAuthenticator() { + return proxyAuthenticator; + } + + public HttpFiltersSource getFiltersSource() { + return filtersSource; + } + + protected Collection getActivityTrackers() { + return activityTrackers; + } + + public String getProxyAlias() { + return proxyAlias; + } + + protected EventLoopGroup getProxyToServerWorkerFor(TransportProtocol transportProtocol) { + return serverGroup.getProxyToServerWorkerPoolForTransport(transportProtocol); + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerBootstrap.java b/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerBootstrap.java new file mode 100644 index 00000000..3a505ddc --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerBootstrap.java @@ -0,0 +1,567 @@ +package org.littleshoot.proxy.impl; + +import static java.util.Objects.requireNonNullElseGet; + +import java.net.InetSocketAddress; +import java.time.Duration; +import java.util.Collection; +import java.util.Properties; +import java.util.concurrent.ConcurrentLinkedQueue; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.ActivityTracker; +import org.littleshoot.proxy.ChainedProxyManager; +import org.littleshoot.proxy.DefaultHostResolver; +import org.littleshoot.proxy.DnsSecServerResolver; +import org.littleshoot.proxy.HostResolver; +import org.littleshoot.proxy.HttpFiltersSource; +import org.littleshoot.proxy.HttpFiltersSourceAdapter; +import org.littleshoot.proxy.HttpProxyServer; +import org.littleshoot.proxy.HttpProxyServerBootstrap; +import org.littleshoot.proxy.Launcher; +import org.littleshoot.proxy.MitmManager; +import org.littleshoot.proxy.ProxyAuthenticator; +import org.littleshoot.proxy.ServerConnectionPoolType; +import org.littleshoot.proxy.SslEngineSource; +import org.littleshoot.proxy.TransportProtocol; +import org.littleshoot.proxy.extras.ActivityLogger; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +class DefaultHttpProxyServerBootstrap implements HttpProxyServerBootstrap { + private static final Logger LOG = LoggerFactory.getLogger(DefaultHttpProxyServerBootstrap.class); + + private static final String DEFAULT_LITTLE_PROXY_NAME = "LittleProxy"; + private static final int MAX_INITIAL_LINE_LENGTH_DEFAULT = 8192; + private static final int MAX_HEADER_SIZE_DEFAULT = 8192 * 2; + private static final int MAX_CHUNK_SIZE_DEFAULT = 8192 * 2; + private static final String DEFAULT_JKS_KEYSTORE_PATH = "littleproxy_keystore.jks"; + + private String name = DEFAULT_LITTLE_PROXY_NAME; + @Nullable private ServerGroup serverGroup; + private TransportProtocol transportProtocol = TransportProtocol.TCP; + @Nullable private InetSocketAddress requestedAddress; + private int port = DefaultHttpProxyServer.DEFAULT_PORT; + private boolean allowLocalOnly = true; + @Nullable private SslEngineSource sslEngineSource; + private boolean authenticateSslClients = true; + @Nullable private ProxyAuthenticator proxyAuthenticator; + @Nullable private ChainedProxyManager chainProxyManager; + @Nullable private MitmManager mitmManager; + private HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter(); + private boolean transparent; + private Duration idleConnectionTimeout = Duration.ofSeconds(70); + private final Collection activityTrackers = new ConcurrentLinkedQueue<>(); + private int connectTimeout = 40000; + private HostResolver serverResolver = new DefaultHostResolver(); + private long readThrottleBytesPerSecond; + private long writeThrottleBytesPerSecond; + @Nullable private InetSocketAddress localAddress; + @Nullable private String proxyAlias; + private int clientToProxyAcceptorThreads = ServerGroup.DEFAULT_INCOMING_ACCEPTOR_THREADS; + private int clientToProxyWorkerThreads = ServerGroup.DEFAULT_INCOMING_WORKER_THREADS; + private int proxyToServerWorkerThreads = ServerGroup.DEFAULT_OUTGOING_WORKER_THREADS; + private int maxInitialLineLength = MAX_INITIAL_LINE_LENGTH_DEFAULT; + private int maxHeaderSize = MAX_HEADER_SIZE_DEFAULT; + private int maxChunkSize = MAX_CHUNK_SIZE_DEFAULT; + private boolean allowRequestToOriginServer; + private boolean acceptProxyProtocol; + private boolean sendProxyProtocol; + private boolean useSharedServerConnectionPool = false; + private int maxConnectionsPerHost = 10; + private int maxConnections = ConcurrentMapServerConnectionPool.DEFAULT_MAX_TOTAL_CONNECTIONS; + private ServerConnectionPoolType serverConnectionPoolType = + ServerConnectionPoolType.CONCURRENT_MAP; + @Nullable private Duration poolIdleTimeout; + private boolean poolSharedMitmConnections = false; + private boolean poolPerRequestInMitm = false; + + DefaultHttpProxyServerBootstrap() {} + + DefaultHttpProxyServerBootstrap(Properties props) { + withUseDnsSec(ProxyUtils.extractBooleanDefaultFalse(props, "dnssec")); + transparent = ProxyUtils.extractBooleanDefaultFalse(props, DefaultHttpProxyServer.TRANSPARENT); + idleConnectionTimeout = + Duration.ofSeconds(ProxyUtils.extractInt(props, "idle_connection_timeout")); + connectTimeout = ProxyUtils.extractInt(props, "connect_timeout", 0); + maxInitialLineLength = + ProxyUtils.extractInt(props, "max_initial_line_length", MAX_INITIAL_LINE_LENGTH_DEFAULT); + maxHeaderSize = ProxyUtils.extractInt(props, "max_header_size", MAX_HEADER_SIZE_DEFAULT); + maxChunkSize = ProxyUtils.extractInt(props, "max_chunk_size", MAX_CHUNK_SIZE_DEFAULT); + if (props.containsKey(DefaultHttpProxyServer.NAME)) { + name = props.getProperty(DefaultHttpProxyServer.NAME, DEFAULT_LITTLE_PROXY_NAME); + } + if (props.containsKey(DefaultHttpProxyServer.ADDRESS)) { + requestedAddress = + ProxyUtils.resolveSocketAddress(props.getProperty(DefaultHttpProxyServer.ADDRESS)); + } + if (props.containsKey(DefaultHttpProxyServer.PORT)) { + port = ProxyUtils.extractInt(props, DefaultHttpProxyServer.PORT, Launcher.DEFAULT_PORT); + } + if (props.containsKey(DefaultHttpProxyServer.NIC)) { + localAddress = + new InetSocketAddress( + props.getProperty( + DefaultHttpProxyServer.NIC, DefaultHttpProxyServer.DEFAULT_NIC_VALUE), + 0); + } + if (props.containsKey(DefaultHttpProxyServer.PROXY_ALIAS)) { + proxyAlias = props.getProperty(DefaultHttpProxyServer.PROXY_ALIAS); + } + if (props.containsKey(DefaultHttpProxyServer.ALLOW_LOCAL_ONLY)) { + allowLocalOnly = + ProxyUtils.extractBooleanDefaultFalse(props, DefaultHttpProxyServer.ALLOW_LOCAL_ONLY); + } + if (props.containsKey(DefaultHttpProxyServer.AUTHENTICATE_SSL_CLIENTS)) { + authenticateSslClients = + ProxyUtils.extractBooleanDefaultFalse( + props, DefaultHttpProxyServer.AUTHENTICATE_SSL_CLIENTS); + boolean trustAllServers = + ProxyUtils.extractBooleanDefaultFalse( + props, DefaultHttpProxyServer.SSL_CLIENTS_TRUST_ALL_SERVERS); + boolean sendCerts = + ProxyUtils.extractBooleanDefaultFalse( + props, DefaultHttpProxyServer.SSL_CLIENTS_SEND_CERTS); + + if (authenticateSslClients + && props.containsKey(DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_PATH)) { + String keyStorePath = props.getProperty(DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_PATH); + if (props.containsKey(DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_PASSWORD)) { + String keyStoreAlias = + props.getProperty(DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_ALIAS, ""); + String keyStorePassword = + props.getProperty(DefaultHttpProxyServer.SSL_CLIENTS_KEYSTORE_PASSWORD, ""); + sslEngineSource = + new SelfSignedSslEngineSource( + keyStorePath, trustAllServers, sendCerts, keyStoreAlias, keyStorePassword); + } else { + sslEngineSource = new SelfSignedSslEngineSource(keyStorePath, trustAllServers, sendCerts); + } + } else { + sslEngineSource = + new SelfSignedSslEngineSource(DEFAULT_JKS_KEYSTORE_PATH, trustAllServers, sendCerts); + } + } + if (props.containsKey(DefaultHttpProxyServer.TRANSPARENT)) { + transparent = + ProxyUtils.extractBooleanDefaultFalse(props, DefaultHttpProxyServer.TRANSPARENT); + } + + if (props.containsKey(DefaultHttpProxyServer.THROTTLE_READ_BYTES_PER_SECOND)) { + readThrottleBytesPerSecond = + ProxyUtils.extractLong(props, DefaultHttpProxyServer.THROTTLE_READ_BYTES_PER_SECOND, 0L); + } + if (props.containsKey(DefaultHttpProxyServer.THROTTLE_WRITE_BYTES_PER_SECOND)) { + writeThrottleBytesPerSecond = + ProxyUtils.extractLong(props, DefaultHttpProxyServer.THROTTLE_WRITE_BYTES_PER_SECOND, 0L); + } + + if (props.containsKey(DefaultHttpProxyServer.ALLOW_REQUESTS_TO_ORIGIN_SERVER)) { + allowRequestToOriginServer = + ProxyUtils.extractBooleanDefaultFalse( + props, DefaultHttpProxyServer.ALLOW_REQUESTS_TO_ORIGIN_SERVER); + } + if (props.containsKey(DefaultHttpProxyServer.ALLOW_PROXY_PROTOCOL)) { + acceptProxyProtocol = + ProxyUtils.extractBooleanDefaultFalse(props, DefaultHttpProxyServer.ALLOW_PROXY_PROTOCOL); + } + if (props.containsKey(DefaultHttpProxyServer.SEND_PROXY_PROTOCOL)) { + sendProxyProtocol = + ProxyUtils.extractBooleanDefaultFalse(props, DefaultHttpProxyServer.SEND_PROXY_PROTOCOL); + } + if (props.containsKey(DefaultHttpProxyServer.SERVER_CONNECTION_POOL_TYPE)) { + String poolTypeValue = + props.getProperty(DefaultHttpProxyServer.SERVER_CONNECTION_POOL_TYPE, "CONCURRENT_MAP"); + try { + serverConnectionPoolType = + ServerConnectionPoolType.valueOf(poolTypeValue.trim().toUpperCase()); + } catch (IllegalArgumentException e) { + LOG.warn("Unknown server connection pool type: {}", poolTypeValue); + } + } + if (props.containsKey(DefaultHttpProxyServer.USE_SHARED_SERVER_CONNECTION_POOL)) { + useSharedServerConnectionPool = + Boolean.parseBoolean( + props.getProperty(DefaultHttpProxyServer.USE_SHARED_SERVER_CONNECTION_POOL).trim()); + } + if (props.containsKey(DefaultHttpProxyServer.MAX_CONNECTIONS_PER_HOST)) { + maxConnectionsPerHost = + ProxyUtils.extractInt( + props, DefaultHttpProxyServer.MAX_CONNECTIONS_PER_HOST, maxConnectionsPerHost); + } + if (props.containsKey(DefaultHttpProxyServer.MAX_TOTAL_CONNECTIONS)) { + maxConnections = + ProxyUtils.extractInt( + props, DefaultHttpProxyServer.MAX_TOTAL_CONNECTIONS, maxConnections); + } + if (props.containsKey(DefaultHttpProxyServer.POOL_SHARED_MITM_CONNECTIONS)) { + poolSharedMitmConnections = + Boolean.parseBoolean( + props.getProperty(DefaultHttpProxyServer.POOL_SHARED_MITM_CONNECTIONS).trim()); + } + if (props.containsKey(DefaultHttpProxyServer.POOL_PER_REQUEST_IN_MITM)) { + poolPerRequestInMitm = + Boolean.parseBoolean( + props.getProperty(DefaultHttpProxyServer.POOL_PER_REQUEST_IN_MITM).trim()); + } + if (props.containsKey(DefaultHttpProxyServer.CLIENT_TO_PROXY_WORKER_THREADS)) { + clientToProxyWorkerThreads = + ProxyUtils.extractInt(props, DefaultHttpProxyServer.CLIENT_TO_PROXY_WORKER_THREADS, 0); + } + if (props.containsKey(DefaultHttpProxyServer.PROXY_TO_SERVER_WORKER_THREADS)) { + proxyToServerWorkerThreads = + ProxyUtils.extractInt(props, DefaultHttpProxyServer.PROXY_TO_SERVER_WORKER_THREADS, 0); + } + if (props.containsKey(DefaultHttpProxyServer.ACCEPTOR_THREADS)) { + clientToProxyAcceptorThreads = + ProxyUtils.extractInt(props, DefaultHttpProxyServer.ACCEPTOR_THREADS, 0); + } + if (props.containsKey(DefaultHttpProxyServer.ACTIVITY_LOG_FORMAT)) { + String format = props.getProperty(DefaultHttpProxyServer.ACTIVITY_LOG_FORMAT); + try { + org.littleshoot.proxy.extras.LogFormat logFormat = + org.littleshoot.proxy.extras.LogFormat.valueOf(format.toUpperCase()); + plusActivityTracker(new ActivityLogger(logFormat)); + } catch (IllegalArgumentException e) { + LOG.warn("Unknown activity log format requested in properties: {}", format); + } + } + } + + DefaultHttpProxyServerBootstrap(ServerGroup serverGroup, DefaultHttpProxyServerConfig config) { + this.serverGroup = serverGroup; + this.transportProtocol = config.getTransportProtocol(); + this.requestedAddress = config.getRequestedAddress(); + this.port = config.getRequestedAddress().getPort(); + this.sslEngineSource = config.getSslEngineSource(); + this.authenticateSslClients = config.isAuthenticateSslClients(); + this.proxyAuthenticator = config.getProxyAuthenticator(); + this.chainProxyManager = config.getChainProxyManager(); + this.mitmManager = config.getMitmManager(); + this.filtersSource = config.getFiltersSource(); + this.transparent = config.isTransparent(); + this.idleConnectionTimeout = config.getIdleConnectionTimeout(); + this.activityTrackers.addAll(config.getActivityTrackers()); + this.connectTimeout = config.getConnectTimeout(); + this.serverResolver = config.getServerResolver(); + this.readThrottleBytesPerSecond = config.getReadThrottleBytesPerSecond(); + this.writeThrottleBytesPerSecond = config.getWriteThrottleBytesPerSecond(); + this.localAddress = config.getLocalAddress(); + this.proxyAlias = config.getProxyAlias(); + this.maxInitialLineLength = config.getMaxInitialLineLength(); + this.maxHeaderSize = config.getMaxHeaderSize(); + this.maxChunkSize = config.getMaxChunkSize(); + this.allowRequestToOriginServer = config.isAllowRequestsToOriginServer(); + this.acceptProxyProtocol = config.isAcceptProxyProtocol(); + this.sendProxyProtocol = config.isSendProxyProtocol(); + ServerConnectionPoolConfig poolConfig = config.getServerConnectionPoolConfig(); + this.useSharedServerConnectionPool = poolConfig.isEnabled(); + this.maxConnectionsPerHost = poolConfig.getMaxConnectionsPerHost(); + this.serverConnectionPoolType = poolConfig.getPoolType(); + this.maxConnections = poolConfig.getMaxConnections(); + this.poolIdleTimeout = poolConfig.getIdleTimeout(); + this.poolSharedMitmConnections = poolConfig.isPoolSharedMitmConnections(); + this.poolPerRequestInMitm = poolConfig.isPoolPerRequestInMitm(); + } + + @Override + public HttpProxyServerBootstrap withName(String name) { + this.name = name; + return this; + } + + @Override + public HttpProxyServerBootstrap withAddress(InetSocketAddress address) { + requestedAddress = address; + return this; + } + + @Override + public HttpProxyServerBootstrap withPort(int port) { + requestedAddress = null; + this.port = port; + return this; + } + + @Override + public HttpProxyServerBootstrap withNetworkInterface(InetSocketAddress inetSocketAddress) { + localAddress = inetSocketAddress; + return this; + } + + @Override + public HttpProxyServerBootstrap withProxyAlias(String alias) { + proxyAlias = alias; + return this; + } + + @Override + public HttpProxyServerBootstrap withAllowLocalOnly(boolean allowLocalOnly) { + this.allowLocalOnly = allowLocalOnly; + return this; + } + + @Override + public HttpProxyServerBootstrap withSslEngineSource(SslEngineSource sslEngineSource) { + this.sslEngineSource = sslEngineSource; + if (mitmManager != null) { + LOG.warn( + "Enabled encrypted inbound connections with man in the middle. " + + "These are mutually exclusive - man in the middle will be disabled."); + mitmManager = null; + } + return this; + } + + @Override + public HttpProxyServerBootstrap withAuthenticateSslClients(boolean authenticateSslClients) { + this.authenticateSslClients = authenticateSslClients; + return this; + } + + @Override + public HttpProxyServerBootstrap withProxyAuthenticator(ProxyAuthenticator proxyAuthenticator) { + this.proxyAuthenticator = proxyAuthenticator; + return this; + } + + @Override + public HttpProxyServerBootstrap withChainProxyManager(ChainedProxyManager chainProxyManager) { + this.chainProxyManager = chainProxyManager; + return this; + } + + @Override + public HttpProxyServerBootstrap withManInTheMiddle(MitmManager mitmManager) { + this.mitmManager = mitmManager; + if (sslEngineSource != null) { + LOG.warn( + "Enabled man in the middle with encrypted inbound connections. " + + "These are mutually exclusive - encrypted inbound connections will be disabled."); + sslEngineSource = null; + } + return this; + } + + @Override + public HttpProxyServerBootstrap withFiltersSource(HttpFiltersSource filtersSource) { + this.filtersSource = filtersSource; + return this; + } + + @Override + public HttpProxyServerBootstrap withUseDnsSec(boolean useDnsSec) { + if (useDnsSec) { + serverResolver = new DnsSecServerResolver(); + } else { + serverResolver = new DefaultHostResolver(); + } + return this; + } + + @Override + public HttpProxyServerBootstrap withTransparent(boolean transparent) { + this.transparent = transparent; + return this; + } + + @Override + public HttpProxyServerBootstrap withIdleConnectionTimeout(int idleConnectionTimeoutInSeconds) { + this.idleConnectionTimeout = Duration.ofSeconds(idleConnectionTimeoutInSeconds); + return this; + } + + @Override + public HttpProxyServerBootstrap withIdleConnectionTimeout(Duration idleConnectionTimeout) { + this.idleConnectionTimeout = idleConnectionTimeout; + return this; + } + + @Override + public HttpProxyServerBootstrap withConnectTimeout(int connectTimeout) { + this.connectTimeout = connectTimeout; + return this; + } + + @Override + public HttpProxyServerBootstrap withServerResolver(HostResolver serverResolver) { + this.serverResolver = serverResolver; + return this; + } + + @Override + public HttpProxyServerBootstrap withServerGroup(ServerGroup group) { + serverGroup = group; + return this; + } + + @Override + public HttpProxyServerBootstrap plusActivityTracker(ActivityTracker activityTracker) { + activityTrackers.add(activityTracker); + return this; + } + + @Override + public HttpProxyServerBootstrap withThrottling( + long readThrottleBytesPerSecond, long writeThrottleBytesPerSecond) { + this.readThrottleBytesPerSecond = readThrottleBytesPerSecond; + this.writeThrottleBytesPerSecond = writeThrottleBytesPerSecond; + return this; + } + + @Override + public HttpProxyServerBootstrap withMaxInitialLineLength(int maxInitialLineLength) { + this.maxInitialLineLength = maxInitialLineLength; + return this; + } + + @Override + public HttpProxyServerBootstrap withMaxHeaderSize(int maxHeaderSize) { + this.maxHeaderSize = maxHeaderSize; + return this; + } + + @Override + public HttpProxyServerBootstrap withMaxChunkSize(int maxChunkSize) { + this.maxChunkSize = maxChunkSize; + return this; + } + + @Override + public HttpProxyServerBootstrap withAllowRequestToOriginServer( + boolean allowRequestToOriginServer) { + this.allowRequestToOriginServer = allowRequestToOriginServer; + return this; + } + + @Override + public HttpProxyServerBootstrap withAcceptProxyProtocol(boolean acceptProxyProtocol) { + this.acceptProxyProtocol = acceptProxyProtocol; + return this; + } + + @Override + public HttpProxyServerBootstrap withSendProxyProtocol(boolean sendProxyProtocol) { + this.sendProxyProtocol = sendProxyProtocol; + return this; + } + + @Override + public HttpProxyServerBootstrap withServerConnectionPoolType(ServerConnectionPoolType poolType) { + this.serverConnectionPoolType = + poolType != null ? poolType : ServerConnectionPoolType.CONCURRENT_MAP; + return this; + } + + @Override + public HttpProxyServerBootstrap withSharedServerConnectionPool( + boolean useSharedServerConnectionPool) { + this.useSharedServerConnectionPool = useSharedServerConnectionPool; + return this; + } + + @Override + public HttpProxyServerBootstrap withMaxConnectionsPerHost(int maxConnectionsPerHost) { + this.maxConnectionsPerHost = maxConnectionsPerHost; + return this; + } + + @Override + public HttpProxyServerBootstrap withMaxConnections(int maxConnections) { + this.maxConnections = maxConnections; + return this; + } + + @Override + public HttpProxyServerBootstrap withPoolIdleTimeout(Duration idleTimeout) { + this.poolIdleTimeout = idleTimeout; + return this; + } + + @Override + public HttpProxyServerBootstrap withPoolSharedMitmConnections(boolean poolSharedMitmConnections) { + this.poolSharedMitmConnections = poolSharedMitmConnections; + return this; + } + + @Override + public HttpProxyServerBootstrap withPoolPerRequestInMitm(boolean poolPerRequestInMitm) { + this.poolPerRequestInMitm = poolPerRequestInMitm; + return this; + } + + @Override + public HttpProxyServer start() { + return build().start(); + } + + @Override + public HttpProxyServerBootstrap withThreadPoolConfiguration( + ThreadPoolConfiguration configuration) { + clientToProxyAcceptorThreads = configuration.getAcceptorThreads(); + clientToProxyWorkerThreads = configuration.getClientToProxyWorkerThreads(); + proxyToServerWorkerThreads = configuration.getProxyToServerWorkerThreads(); + return this; + } + + private DefaultHttpProxyServer build() { + final ServerGroup selectedServerGroup = + requireNonNullElseGet( + this.serverGroup, + () -> + new ServerGroup( + name, + clientToProxyAcceptorThreads, + clientToProxyWorkerThreads, + proxyToServerWorkerThreads)); + + ServerConnectionPoolConfig poolConfig = + new ServerConnectionPoolConfig() + .setEnabled(useSharedServerConnectionPool) + .setPoolType(serverConnectionPoolType) + .setMaxConnectionsPerHost(maxConnectionsPerHost) + .setMaxConnections(maxConnections) + .setIdleTimeout(poolIdleTimeout) + .setPoolSharedMitmConnections(poolSharedMitmConnections) + .setPoolPerRequestInMitm(poolPerRequestInMitm); + + DefaultHttpProxyServerConfig serverConfig = + new DefaultHttpProxyServerConfig() + .setTransportProtocol(transportProtocol) + .setRequestedAddress(determineListenAddress()) + .setSslEngineSource(sslEngineSource) + .setAuthenticateSslClients(authenticateSslClients) + .setProxyAuthenticator(proxyAuthenticator) + .setChainProxyManager(chainProxyManager) + .setMitmManager(mitmManager) + .setFiltersSource(filtersSource) + .setTransparent(transparent) + .setIdleConnectionTimeout(idleConnectionTimeout) + .setActivityTrackers(activityTrackers) + .setConnectTimeout(connectTimeout) + .setServerResolver(serverResolver) + .setReadThrottleBytesPerSecond(readThrottleBytesPerSecond) + .setWriteThrottleBytesPerSecond(writeThrottleBytesPerSecond) + .setLocalAddress(localAddress) + .setProxyAlias(proxyAlias) + .setMaxInitialLineLength(maxInitialLineLength) + .setMaxHeaderSize(maxHeaderSize) + .setMaxChunkSize(maxChunkSize) + .setAllowRequestsToOriginServer(allowRequestToOriginServer) + .setAcceptProxyProtocol(acceptProxyProtocol) + .setSendProxyProtocol(sendProxyProtocol) + .setServerConnectionPoolConfig(poolConfig); + + return new DefaultHttpProxyServer(selectedServerGroup, serverConfig); + } + + private InetSocketAddress determineListenAddress() { + if (requestedAddress != null) { + return requestedAddress; + } + if (allowLocalOnly) { + return new InetSocketAddress(DefaultHttpProxyServer.LOCAL_ADDRESS, port); + } + return new InetSocketAddress(port); + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerConfig.java b/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerConfig.java new file mode 100644 index 00000000..b34d6aa8 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerConfig.java @@ -0,0 +1,276 @@ +package org.littleshoot.proxy.impl; + +import java.net.InetSocketAddress; +import java.time.Duration; +import java.util.Collection; +import java.util.concurrent.ConcurrentLinkedQueue; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.ActivityTracker; +import org.littleshoot.proxy.ChainedProxyManager; +import org.littleshoot.proxy.HostResolver; +import org.littleshoot.proxy.HttpFiltersSource; +import org.littleshoot.proxy.MitmManager; +import org.littleshoot.proxy.ProxyAuthenticator; +import org.littleshoot.proxy.SslEngineSource; +import org.littleshoot.proxy.TransportProtocol; + +public class DefaultHttpProxyServerConfig { + private TransportProtocol transportProtocol; + private InetSocketAddress requestedAddress; + @Nullable private SslEngineSource sslEngineSource; + private boolean authenticateSslClients; + @Nullable private ProxyAuthenticator proxyAuthenticator; + @Nullable private ChainedProxyManager chainProxyManager; + @Nullable private MitmManager mitmManager; + private HttpFiltersSource filtersSource; + private boolean transparent; + private Duration idleConnectionTimeout; + private final Collection activityTrackers = new ConcurrentLinkedQueue<>(); + private int connectTimeout; + private HostResolver serverResolver; + private long readThrottleBytesPerSecond; + private long writeThrottleBytesPerSecond; + @Nullable private InetSocketAddress localAddress; + @Nullable private String proxyAlias; + private int maxInitialLineLength; + private int maxHeaderSize; + private int maxChunkSize; + private boolean allowRequestsToOriginServer; + private boolean acceptProxyProtocol; + private boolean sendProxyProtocol; + private ServerConnectionPoolConfig serverConnectionPoolConfig = new ServerConnectionPoolConfig(); + + public TransportProtocol getTransportProtocol() { + return transportProtocol; + } + + public DefaultHttpProxyServerConfig setTransportProtocol(TransportProtocol transportProtocol) { + this.transportProtocol = transportProtocol; + return this; + } + + public InetSocketAddress getRequestedAddress() { + return requestedAddress; + } + + public DefaultHttpProxyServerConfig setRequestedAddress(InetSocketAddress requestedAddress) { + this.requestedAddress = requestedAddress; + return this; + } + + @Nullable + public SslEngineSource getSslEngineSource() { + return sslEngineSource; + } + + public DefaultHttpProxyServerConfig setSslEngineSource( + @Nullable SslEngineSource sslEngineSource) { + this.sslEngineSource = sslEngineSource; + return this; + } + + public boolean isAuthenticateSslClients() { + return authenticateSslClients; + } + + public DefaultHttpProxyServerConfig setAuthenticateSslClients(boolean authenticateSslClients) { + this.authenticateSslClients = authenticateSslClients; + return this; + } + + @Nullable + public ProxyAuthenticator getProxyAuthenticator() { + return proxyAuthenticator; + } + + public DefaultHttpProxyServerConfig setProxyAuthenticator( + @Nullable ProxyAuthenticator proxyAuthenticator) { + this.proxyAuthenticator = proxyAuthenticator; + return this; + } + + @Nullable + public ChainedProxyManager getChainProxyManager() { + return chainProxyManager; + } + + public DefaultHttpProxyServerConfig setChainProxyManager( + @Nullable ChainedProxyManager chainProxyManager) { + this.chainProxyManager = chainProxyManager; + return this; + } + + @Nullable + public MitmManager getMitmManager() { + return mitmManager; + } + + public DefaultHttpProxyServerConfig setMitmManager(@Nullable MitmManager mitmManager) { + this.mitmManager = mitmManager; + return this; + } + + public HttpFiltersSource getFiltersSource() { + return filtersSource; + } + + public DefaultHttpProxyServerConfig setFiltersSource(HttpFiltersSource filtersSource) { + this.filtersSource = filtersSource; + return this; + } + + public boolean isTransparent() { + return transparent; + } + + public DefaultHttpProxyServerConfig setTransparent(boolean transparent) { + this.transparent = transparent; + return this; + } + + public Duration getIdleConnectionTimeout() { + return idleConnectionTimeout; + } + + public DefaultHttpProxyServerConfig setIdleConnectionTimeout(Duration idleConnectionTimeout) { + this.idleConnectionTimeout = idleConnectionTimeout; + return this; + } + + public Collection getActivityTrackers() { + return activityTrackers; + } + + public DefaultHttpProxyServerConfig setActivityTrackers( + Collection activityTrackers) { + this.activityTrackers.clear(); + this.activityTrackers.addAll(activityTrackers); + return this; + } + + public int getConnectTimeout() { + return connectTimeout; + } + + public DefaultHttpProxyServerConfig setConnectTimeout(int connectTimeout) { + this.connectTimeout = connectTimeout; + return this; + } + + public HostResolver getServerResolver() { + return serverResolver; + } + + public DefaultHttpProxyServerConfig setServerResolver(HostResolver serverResolver) { + this.serverResolver = serverResolver; + return this; + } + + public long getReadThrottleBytesPerSecond() { + return readThrottleBytesPerSecond; + } + + public DefaultHttpProxyServerConfig setReadThrottleBytesPerSecond( + long readThrottleBytesPerSecond) { + this.readThrottleBytesPerSecond = readThrottleBytesPerSecond; + return this; + } + + public long getWriteThrottleBytesPerSecond() { + return writeThrottleBytesPerSecond; + } + + public DefaultHttpProxyServerConfig setWriteThrottleBytesPerSecond( + long writeThrottleBytesPerSecond) { + this.writeThrottleBytesPerSecond = writeThrottleBytesPerSecond; + return this; + } + + @Nullable + public InetSocketAddress getLocalAddress() { + return localAddress; + } + + public DefaultHttpProxyServerConfig setLocalAddress(@Nullable InetSocketAddress localAddress) { + this.localAddress = localAddress; + return this; + } + + @Nullable + public String getProxyAlias() { + return proxyAlias; + } + + public DefaultHttpProxyServerConfig setProxyAlias(@Nullable String proxyAlias) { + this.proxyAlias = proxyAlias; + return this; + } + + public int getMaxInitialLineLength() { + return maxInitialLineLength; + } + + public DefaultHttpProxyServerConfig setMaxInitialLineLength(int maxInitialLineLength) { + this.maxInitialLineLength = maxInitialLineLength; + return this; + } + + public int getMaxHeaderSize() { + return maxHeaderSize; + } + + public DefaultHttpProxyServerConfig setMaxHeaderSize(int maxHeaderSize) { + this.maxHeaderSize = maxHeaderSize; + return this; + } + + public int getMaxChunkSize() { + return maxChunkSize; + } + + public DefaultHttpProxyServerConfig setMaxChunkSize(int maxChunkSize) { + this.maxChunkSize = maxChunkSize; + return this; + } + + public boolean isAllowRequestsToOriginServer() { + return allowRequestsToOriginServer; + } + + public DefaultHttpProxyServerConfig setAllowRequestsToOriginServer( + boolean allowRequestsToOriginServer) { + this.allowRequestsToOriginServer = allowRequestsToOriginServer; + return this; + } + + public boolean isAcceptProxyProtocol() { + return acceptProxyProtocol; + } + + public DefaultHttpProxyServerConfig setAcceptProxyProtocol(boolean acceptProxyProtocol) { + this.acceptProxyProtocol = acceptProxyProtocol; + return this; + } + + public boolean isSendProxyProtocol() { + return sendProxyProtocol; + } + + public DefaultHttpProxyServerConfig setSendProxyProtocol(boolean sendProxyProtocol) { + this.sendProxyProtocol = sendProxyProtocol; + return this; + } + + public ServerConnectionPoolConfig getServerConnectionPoolConfig() { + return serverConnectionPoolConfig; + } + + public DefaultHttpProxyServerConfig setServerConnectionPoolConfig( + ServerConnectionPoolConfig serverConnectionPoolConfig) { + this.serverConnectionPoolConfig = + serverConnectionPoolConfig != null + ? serverConnectionPoolConfig + : new ServerConnectionPoolConfig(); + return this; + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/Hostname.java b/src/main/java/org/littleshoot/proxy/impl/Hostname.java new file mode 100644 index 00000000..ab8d8f52 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/Hostname.java @@ -0,0 +1,87 @@ +package org.littleshoot.proxy.impl; + +import static java.lang.System.nanoTime; +import static java.util.concurrent.TimeUnit.NANOSECONDS; +import static java.util.concurrent.TimeUnit.SECONDS; + +import java.io.BufferedReader; +import java.io.IOException; +import java.io.InputStreamReader; +import java.net.InetAddress; +import java.net.UnknownHostException; +import java.util.stream.Stream; +import org.jspecify.annotations.Nullable; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +class Hostname { + private static final Logger LOG = LoggerFactory.getLogger(Hostname.class); + + private static volatile String hostname; + + @Nullable + static String getHostName() { + if (hostname == null) { + hostname = resolveHostName(); + } + return hostname; + } + + @Nullable + private static String resolveHostName() { + long startTime = nanoTime(); + String hostName = + byAllMeans( + env("HOSTNAME"), // Most OSs + env("COMPUTERNAME"), // Windows + Hostname::executeHostname, + Hostname::getLocalHost); + long duration = NANOSECONDS.toMillis(nanoTime() - startTime); + LOG.info("Resolved local machine's hostname \"{}\" in {} ms.", hostName, duration); + return hostName; + } + + @Nullable + @SafeVarargs + private static String byAllMeans(SupplierEx... means) { + return Stream.of(means) + .map(mean -> getOrNull(mean)) + .filter(host -> host != null) + .findFirst() + .orElse(null); + } + + @Nullable + private static String getOrNull(SupplierEx s) { + try { + return s.get(); + } catch (Exception e) { + LOG.info("Failed to resolve local machine's hostname", e); + return null; + } + } + + private static SupplierEx env(String name) { + return () -> System.getenv(name); + } + + /** + * "hostname" command works on Windows, Mac, and Linux. Usually much faster than {@link + * InetAddress#getLocalHost()}. + */ + private static String executeHostname() throws IOException, InterruptedException { + Process p = new ProcessBuilder("hostname").start(); + + try (BufferedReader reader = new BufferedReader(new InputStreamReader(p.getInputStream()))) { + String line = reader.readLine(); + if (p.waitFor(5, SECONDS) && line != null) { + return line.trim(); + } + } + return null; + } + + private static String getLocalHost() throws UnknownHostException { + return InetAddress.getLocalHost().getHostName(); + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/NetworkUtils.java b/src/main/java/org/littleshoot/proxy/impl/NetworkUtils.java deleted file mode 100644 index 811dbabf..00000000 --- a/src/main/java/org/littleshoot/proxy/impl/NetworkUtils.java +++ /dev/null @@ -1,47 +0,0 @@ -package org.littleshoot.proxy.impl; - -import java.net.*; -import java.util.Enumeration; - -/** - * @deprecated This class is no longer used by LittleProxy and may be removed in a future release. - */ -@Deprecated -public class NetworkUtils { - /** - * @deprecated This method is no longer used by LittleProxy and may be removed in a future release. - */ - @Deprecated - public static InetAddress getLocalHost() throws UnknownHostException { - return InetAddress.getLocalHost(); - } - - /** - * @deprecated This method is no longer used by LittleProxy and may be removed in a future release. - */ - @Deprecated - public static InetAddress firstLocalNonLoopbackIpv4Address() { - try { - Enumeration networkInterfaces = NetworkInterface - .getNetworkInterfaces(); - while (networkInterfaces.hasMoreElements()) { - NetworkInterface networkInterface = networkInterfaces - .nextElement(); - if (networkInterface.isUp()) { - for (InterfaceAddress ifAddress : networkInterface - .getInterfaceAddresses()) { - if (ifAddress.getNetworkPrefixLength() > 0 - && ifAddress.getNetworkPrefixLength() <= 32 - && !ifAddress.getAddress().isLoopbackAddress()) { - return ifAddress.getAddress(); - } - } - } - } - return null; - } catch (SocketException se) { - return null; - } - } - -} diff --git a/src/main/java/org/littleshoot/proxy/impl/PendingRequest.java b/src/main/java/org/littleshoot/proxy/impl/PendingRequest.java new file mode 100644 index 00000000..9ddfdb87 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/PendingRequest.java @@ -0,0 +1,38 @@ +package org.littleshoot.proxy.impl; + +import io.netty.handler.codec.http.HttpRequest; +import org.littleshoot.proxy.HttpFilters; + +/** + * Tracks a pending request and its associated client connection and filters for HTTP pipelining. + */ +public class PendingRequest { + private final ClientToProxyConnection clientConnection; + private final HttpRequest request; + private final HttpFilters filters; + private final long timestamp; + + public PendingRequest( + ClientToProxyConnection clientConnection, HttpRequest request, HttpFilters filters) { + this.clientConnection = clientConnection; + this.request = request; + this.filters = filters; + this.timestamp = System.currentTimeMillis(); + } + + public ClientToProxyConnection getClientConnection() { + return clientConnection; + } + + public HttpRequest getRequest() { + return request; + } + + public HttpFilters getFilters() { + return filters; + } + + public long getTimestamp() { + return timestamp; + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/PoolMetrics.java b/src/main/java/org/littleshoot/proxy/impl/PoolMetrics.java new file mode 100644 index 00000000..9ebbb314 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/PoolMetrics.java @@ -0,0 +1,77 @@ +package org.littleshoot.proxy.impl; + +/** Pool metrics statistics. */ +public class PoolMetrics { + private final int totalConnections; + private final int activeConnections; + private final int idleConnections; + private final long borrowCount; + private final long returnCount; + private final long evictionCount; + private final long validationFailureCount; + + public PoolMetrics( + int totalConnections, + int activeConnections, + int idleConnections, + long borrowCount, + long returnCount, + long evictionCount, + long validationFailureCount) { + this.totalConnections = totalConnections; + this.activeConnections = activeConnections; + this.idleConnections = idleConnections; + this.borrowCount = borrowCount; + this.returnCount = returnCount; + this.evictionCount = evictionCount; + this.validationFailureCount = validationFailureCount; + } + + public int getTotalConnections() { + return totalConnections; + } + + public int getActiveConnections() { + return activeConnections; + } + + public int getIdleConnections() { + return idleConnections; + } + + public long getBorrowCount() { + return borrowCount; + } + + public long getReturnCount() { + return returnCount; + } + + public long getEvictionCount() { + return evictionCount; + } + + public long getValidationFailureCount() { + return validationFailureCount; + } + + @Override + public String toString() { + return "PoolMetrics{" + + "total=" + + totalConnections + + ", active=" + + activeConnections + + ", idle=" + + idleConnections + + ", borrowCount=" + + borrowCount + + ", returnCount=" + + returnCount + + ", evictionCount=" + + evictionCount + + ", validationFailureCount=" + + validationFailureCount + + '}'; + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/ProxyConnection.java b/src/main/java/org/littleshoot/proxy/impl/ProxyConnection.java index 782537dc..f924f830 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ProxyConnection.java +++ b/src/main/java/org/littleshoot/proxy/impl/ProxyConnection.java @@ -1,5 +1,7 @@ package org.littleshoot.proxy.impl; +import static org.littleshoot.proxy.impl.ConnectionState.*; + import io.netty.buffer.ByteBuf; import io.netty.buffer.Unpooled; import io.netty.channel.*; @@ -10,792 +12,743 @@ import io.netty.util.ReferenceCounted; import io.netty.util.concurrent.Future; import io.netty.util.concurrent.Promise; -import org.littleshoot.proxy.HttpFilters; - +import java.util.concurrent.atomic.AtomicLong; import javax.net.ssl.SSLEngine; - -import static org.littleshoot.proxy.impl.ConnectionState.*; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.HttpFilters; /** - *

* Base class for objects that represent a connection to/from our proxy. - *

- *

- * A ProxyConnection models a bidirectional message flow on top of a Netty - * {@link Channel}. - *

- *

- * The {@link #read(Object)} method is called whenever a new message arrives on - * the underlying socket. - *

- *

- * The {@link #write(Object)} method can be called by anyone wanting to write - * data out of the connection. - *

- *

- * ProxyConnection has a lifecycle and its current state within that lifecycle - * is recorded as a {@link ConnectionState}. The allowed states and transitions - * vary a little depending on the concrete implementation of ProxyConnection. - * However, all ProxyConnections share the following lifecycle events: - *

- * + * + *

A ProxyConnection models a bidirectional message flow on top of a Netty {@link Channel}. + * + *

The {@link #read(Object)} method is called whenever a new message arrives on the underlying + * socket. + * + *

The {@link #write(Object)} method can be called by anyone wanting to write data out of the + * connection. + * + *

ProxyConnection has a lifecycle and its current state within that lifecycle is recorded as a + * {@link ConnectionState}. The allowed states and transitions vary a little depending on the + * concrete implementation of ProxyConnection. However, all ProxyConnections share the following + * lifecycle events: + * *

    - *
  • {@link #connected()} - Once the underlying channel is active, the - * ProxyConnection is considered connected and moves into - * {@link ConnectionState#AWAITING_INITIAL}. The Channel is recorded at this - * time for later referencing.
  • - *
  • {@link #disconnected()} - When the underlying channel goes inactive, the - * ProxyConnection moves into {@link ConnectionState#DISCONNECTED}
  • - *
  • {@link #becameWritable()} - When the underlying channel becomes - * writeable, this callback is invoked.
  • + *
  • {@link #connected()} - Once the underlying channel is active, the ProxyConnection is + * considered connected and moves into {@link ConnectionState#AWAITING_INITIAL}. The Channel + * is recorded at this time for later referencing. + *
  • {@link #disconnected()} - When the underlying channel goes inactive, the ProxyConnection + * moves into {@link ConnectionState#DISCONNECTED} + *
  • {@link #becameWritable()} - When the underlying channel becomes writeable, this callback is + * invoked. *
- * - *

- * By default, incoming data on the underlying channel is automatically read and - * passed to the {@link #read(Object)} method. Reading can be stopped and - * resumed using {@link #stopReading()} and {@link #resumeReading()}. - *

- * - * @param - * the type of "initial" message. This will be either - * {@link HttpResponse} or {@link HttpRequest}. + * + *

By default, incoming data on the underlying channel is automatically read and passed to the + * {@link #read(Object)} method. Reading can be stopped and resumed using {@link #stopReading()} and + * {@link #resumeReading()}. + * + * @param the type of "initial" message. This will be either {@link HttpResponse} or {@link + * HttpRequest}. */ -abstract class ProxyConnection extends - SimpleChannelInboundHandler { - protected final ProxyConnectionLogger LOG = new ProxyConnectionLogger(this); - - protected final DefaultHttpProxyServer proxyServer; - protected final boolean runsAsSslClient; - - protected volatile ChannelHandlerContext ctx; - protected volatile Channel channel; - - private volatile ConnectionState currentState; - private volatile boolean tunneling = false; - protected volatile long lastReadTime = 0; - - /** - * If using encryption, this holds our {@link SSLEngine}. - */ - protected volatile SSLEngine sslEngine; - - /** - * Construct a new ProxyConnection. - * - * @param initialState - * the state in which this connection starts out - * @param proxyServer - * the {@link DefaultHttpProxyServer} in which we're running - * @param runsAsSslClient - * determines whether this connection acts as an SSL client or - * server (determines who does the handshake) - */ - protected ProxyConnection(ConnectionState initialState, - DefaultHttpProxyServer proxyServer, - boolean runsAsSslClient) { - become(initialState); - this.proxyServer = proxyServer; - this.runsAsSslClient = runsAsSslClient; - } - - /* ************************************************************************* - * Reading - **************************************************************************/ - - /** - * Read is invoked automatically by Netty as messages arrive on the socket. - */ - protected void read(Object msg) { - LOG.debug("Reading: {}", msg); - - lastReadTime = System.currentTimeMillis(); - - if (tunneling) { - // In tunneling mode, this connection is simply shoveling bytes - readRaw((ByteBuf) msg); - } else if ( msg instanceof HAProxyMessage) { - readHAProxyMessage((HAProxyMessage)msg); +@NullMarked +abstract class ProxyConnection extends SimpleChannelInboundHandler { + protected final ProxyConnectionLogger LOG = new ProxyConnectionLogger(this); + + protected final DefaultHttpProxyServer proxyServer; + protected final boolean runsAsSslClient; + + @Nullable protected volatile ChannelHandlerContext ctx; + + @Nullable protected volatile Channel channel; + + private volatile ConnectionState currentState; + protected volatile boolean tunneling; + protected volatile long lastReadTime; + + /** If using encryption, this holds our {@link SSLEngine}. */ + @Nullable protected volatile SSLEngine sslEngine; + + private static final AtomicLong CONNECTION_ID_GENERATOR = new AtomicLong(); + private final long connectionId; + + /** + * Construct a new ProxyConnection. + * + * @param initialState the state in which this connection starts out + * @param proxyServer the {@link DefaultHttpProxyServer} in which we're running + * @param runsAsSslClient determines whether this connection acts as an SSL client or server + * (determines who does the handshake) + */ + protected ProxyConnection( + ConnectionState initialState, DefaultHttpProxyServer proxyServer, boolean runsAsSslClient) { + this.connectionId = CONNECTION_ID_GENERATOR.incrementAndGet(); + become(initialState); + this.proxyServer = proxyServer; + this.runsAsSslClient = runsAsSslClient; + } + + public long getId() { + return connectionId; + } + + /* + * ************************************************************************* + * Reading + **************************************************************************/ + + /** Read is invoked automatically by Netty as messages arrive at the socket. */ + protected void read(Object msg) { + LOG.debug("Reading: {}", msg); + + lastReadTime = System.currentTimeMillis(); + + if (tunneling) { + // In tunneling mode, this connection is simply shoveling bytes + readRaw((ByteBuf) msg); + } else if (msg instanceof HAProxyMessage) { + readHAProxyMessage((HAProxyMessage) msg); + } else if (msg instanceof HttpObject) { + // If not tunneling, then we are always dealing with HttpObjects. + readHTTP((HttpObject) msg); + } else if (msg instanceof ByteBuf) { + readRaw((ByteBuf) msg); + } else { + throw new UnsupportedOperationException( + "Unsupported message type: " + msg.getClass().getName()); + } + } + + /** + * Read an {@link HAProxyMessage} + * + * @param msg {@link HAProxyMessage} + */ + protected abstract void readHAProxyMessage(HAProxyMessage msg); + + /** Handles reading {@link HttpObject}s. */ + @SuppressWarnings("unchecked") + private void readHTTP(HttpObject httpObject) { + ConnectionState nextState = getCurrentState(); + switch (getCurrentState()) { + case AWAITING_INITIAL: + if (httpObject instanceof HttpMessage) { + nextState = readHTTPInitial((I) httpObject); } else { - // If not tunneling, then we are always dealing with HttpObjects. - readHTTP((HttpObject) msg); - } - } - - /** - * Read an {@link HAProxyMessage} - * @param msg {@link HAProxyMessage} - */ - protected abstract void readHAProxyMessage(HAProxyMessage msg); - - - /** - * Handles reading {@link HttpObject}s. - */ - @SuppressWarnings("unchecked") - private void readHTTP(HttpObject httpObject) { - ConnectionState nextState = getCurrentState(); - switch (getCurrentState()) { - case AWAITING_INITIAL: - if (httpObject instanceof HttpMessage) { - nextState = readHTTPInitial((I) httpObject); - } else { - // Similar to the AWAITING_PROXY_AUTHENTICATION case below, we may enter an AWAITING_INITIAL - // state if the proxy responded to an earlier request with a 502 or 504 response, or a short-circuit - // response from a filter. The client may have sent some chunked HttpContent associated with the request - // after the short-circuit response was sent. We can safely drop them. - LOG.debug("Dropping message because HTTP object was not an HttpMessage. HTTP object may be orphaned content from a short-circuited response. Message: {}", httpObject); - } - break; - case AWAITING_CHUNK: - HttpContent chunk = (HttpContent) httpObject; - readHTTPChunk(chunk); - nextState = ProxyUtils.isLastChunk(chunk) ? AWAITING_INITIAL - : AWAITING_CHUNK; - break; - case AWAITING_PROXY_AUTHENTICATION: - if (httpObject instanceof HttpRequest) { - // Once we get an HttpRequest, try to process it as usual - nextState = readHTTPInitial((I) httpObject); - } else { - // Anything that's not an HttpRequest that came in while - // we're pending authentication gets dropped on the floor. This - // can happen if the connected host already sent us some chunks - // (e.g. from a POST) after an initial request that turned out - // to require authentication. - } - break; - case CONNECTING: - LOG.warn("Attempted to read from connection that's in the process of connecting. This shouldn't happen."); - break; - case NEGOTIATING_CONNECT: - LOG.debug("Attempted to read from connection that's in the process of negotiating an HTTP CONNECT. This is probably the LastHttpContent of a chunked CONNECT."); - break; - case AWAITING_CONNECT_OK: - LOG.warn("AWAITING_CONNECT_OK should have been handled by ProxyToServerConnection.read()"); - break; - case HANDSHAKING: - LOG.warn( - "Attempted to read from connection that's in the process of handshaking. This shouldn't happen.", - channel); - break; - case DISCONNECT_REQUESTED: - case DISCONNECTED: - LOG.info("Ignoring message since the connection is closed or about to close"); - break; + // Similar to the AWAITING_PROXY_AUTHENTICATION case below, we may enter an + // AWAITING_INITIAL + // state if the proxy responded to an earlier request with a 502 or 504 + // response, or a short-circuit + // response from a filter. The client may have sent some chunked HttpContent + // associated with the request + // after the short-circuit response was sent. We can safely drop them. + LOG.debug( + "Dropping message because HTTP object was not an HttpMessage. HTTP object may be orphaned content from a short-circuited response. Message: {}", + httpObject); } - become(nextState); - } - - /** - * Implement this to handle reading the initial object (e.g. - * {@link HttpRequest} or {@link HttpResponse}). - */ - protected abstract ConnectionState readHTTPInitial(I httpObject); - - /** - * Implement this to handle reading a chunk in a chunked transfer. - */ - protected abstract void readHTTPChunk(HttpContent chunk); - - /** - * Implement this to handle reading a raw buffer as they are used in HTTP - * tunneling. - */ - protected abstract void readRaw(ByteBuf buf); - - /* ************************************************************************* - * Writing - **************************************************************************/ - - /** - * This method is called by users of the ProxyConnection to send stuff out - * over the socket. - */ - void write(Object msg) { - if (msg instanceof ReferenceCounted) { - LOG.debug("Retaining reference counted message"); - ((ReferenceCounted) msg).retain(); - } - - doWrite(msg); - } - - void doWrite(Object msg) { - LOG.debug("Writing: {}", msg); - - try { - if (msg instanceof HttpObject) { - writeHttp((HttpObject) msg); - } else { - writeRaw((ByteBuf) msg); - } - } finally { - LOG.debug("Wrote: {}", msg); - } - } - - /** - * Writes HttpObjects to the connection asynchronously. - */ - protected void writeHttp(HttpObject httpObject) { - if (ProxyUtils.isLastChunk(httpObject)) { - channel.write(httpObject); - LOG.debug("Writing an empty buffer to signal the end of our chunked transfer"); - writeToChannel(Unpooled.EMPTY_BUFFER); + break; + case AWAITING_CHUNK: + HttpContent chunk = (HttpContent) httpObject; + readHTTPChunk(chunk); + nextState = ProxyUtils.isLastChunk(chunk) ? AWAITING_INITIAL : AWAITING_CHUNK; + break; + case AWAITING_PROXY_AUTHENTICATION: + if (httpObject instanceof HttpRequest) { + // Once we get an HttpRequest, try to process it as usual + nextState = readHTTPInitial((I) httpObject); } else { - writeToChannel(httpObject); + // Anything that's not an HttpRequest that came in while + // we're pending authentication gets dropped on the floor. This + // can happen if the connected host already sent us some chunks + // (e.g. from a POST) after an initial request that turned out + // to require authentication. } - } - - /** - * Writes raw buffers to the connection. - */ - protected void writeRaw(ByteBuf buf) { - writeToChannel(buf); - } - - protected ChannelFuture writeToChannel(final Object msg) { - return channel.writeAndFlush(msg); - } - - /* ************************************************************************* - * Lifecycle - **************************************************************************/ - - /** - * This method is called as soon as the underlying {@link Channel} is - * connected. Note that for proxies with complex {@link ConnectionFlow}s - * that include SSL handshaking and other such things, just because the - * {@link Channel} is connected doesn't mean that our connection is fully - * established. - */ - protected void connected() { - LOG.debug("Connected"); - } - - /** - * This method is called as soon as the underlying {@link Channel} becomes - * disconnected. - */ - protected void disconnected() { - become(DISCONNECTED); - LOG.debug("Disconnected"); - } - - /** - * This method is called when the underlying {@link Channel} times out due - * to an idle timeout. - */ - protected void timedOut() { - disconnect(); - } - - /** - *

- * Enables tunneling on this connection by dropping the HTTP related - * encoders and decoders, as well as idle timers. - *

- * - *

- * Note - the work is done on the {@link ChannelHandlerContext}'s executor - * because {@link ChannelPipeline#remove(String)} can deadlock if called - * directly. - *

- */ - protected ConnectionFlowStep StartTunneling = new ConnectionFlowStep( - this, NEGOTIATING_CONNECT) { + break; + case CONNECTING: + LOG.warn( + "Attempted to read from connection that's in the process of connecting. This shouldn't happen."); + break; + case NEGOTIATING_CONNECT: + LOG.debug( + "Attempted to read from connection that's in the process of negotiating an HTTP CONNECT. This is probably the LastHttpContent of a chunked CONNECT."); + break; + case AWAITING_CONNECT_OK: + LOG.warn("AWAITING_CONNECT_OK should have been handled by ProxyToServerConnection.read()"); + break; + case HANDSHAKING: + LOG.warn( + "Attempted to read from connection that's in the process of handshaking. This shouldn't happen.", + channel); + break; + case DISCONNECT_REQUESTED: + case DISCONNECTED: + LOG.info("Ignoring message since the connection is closed or about to close"); + break; + } + become(nextState); + } + + /** + * Implement this to handle reading the initial object (e.g. {@link HttpRequest} or {@link + * HttpResponse}). + */ + abstract ConnectionState readHTTPInitial(I httpObject); + + /** Implement this to handle reading a chunk in a chunked transfer. */ + protected abstract void readHTTPChunk(HttpContent chunk); + + /** Implement this to handle reading a raw buffer as they are used in HTTP tunneling. */ + protected abstract void readRaw(ByteBuf buf); + + /* + * ************************************************************************* + * Writing + **************************************************************************/ + + /** This method is called by users of the ProxyConnection to send stuff out over the socket. */ + ChannelFuture write(Object msg) { + if (msg instanceof ReferenceCounted) { + LOG.debug("Retaining reference counted message"); + ((ReferenceCounted) msg).retain(); + } + + return doWrite(msg); + } + + ChannelFuture doWrite(Object msg) { + LOG.debug("Writing: {}", msg); + + try { + if (msg instanceof HttpObject) { + return writeHttp((HttpObject) msg); + } else { + return writeRaw((ByteBuf) msg); + } + } finally { + LOG.debug("Wrote: {}", msg); + } + } + + /** Writes HttpObjects to the connection asynchronously. */ + protected ChannelFuture writeHttp(HttpObject httpObject) { + if (ProxyUtils.isLastChunk(httpObject)) { + channel.write(httpObject); + LOG.debug("Writing an empty buffer to signal the end of our chunked transfer"); + return writeToChannel(Unpooled.EMPTY_BUFFER); + } else { + return writeToChannel(httpObject); + } + } + + /** Writes raw buffers to the connection. */ + protected ChannelFuture writeRaw(ByteBuf buf) { + return writeToChannel(buf); + } + + protected ChannelFuture writeToChannel(final Object msg) { + return channel + .writeAndFlush(msg) + .addListener( + l -> { + if (!l.isSuccess()) { + LOG.debug("writeToChannel failed sending message {}", msg, l.cause()); + } + }); + } + + /* + * ************************************************************************* + * Lifecycle + **************************************************************************/ + + /** + * This method is called as soon as the underlying {@link Channel} is connected. Note that for + * proxies with complex {@link ConnectionFlow}s that include SSL handshaking and other such + * things, just because the {@link Channel} is connected doesn't mean that our connection is fully + * established. + */ + protected void connected() { + LOG.debug("Connected"); + } + + /** This method is called as soon as the underlying {@link Channel} becomes disconnected. */ + protected void disconnected() { + become(DISCONNECTED); + LOG.debug("Disconnected"); + } + + /** This method is called when the underlying {@link Channel} times out due to an idle timeout. */ + protected void timedOut() { + disconnect(); + } + + /** + * Enables tunneling on this connection by dropping the HTTP related encoders and decoders, as + * well as idle timers. + * + *

Note - the work is done on the {@link ChannelHandlerContext}'s executor because {@link + * ChannelPipeline#remove(String)} can deadlock if called directly. + */ + protected final ConnectionFlowStep StartTunneling = + new ConnectionFlowStep<>(this, NEGOTIATING_CONNECT) { @Override boolean shouldSuppressInitialRequest() { - return true; + return true; } - protected Future execute() { - try { - ChannelPipeline pipeline = ctx.pipeline(); - if (pipeline.get("encoder") != null) { - pipeline.remove("encoder"); - } - if (pipeline.get("responseWrittenMonitor") != null) { - pipeline.remove("responseWrittenMonitor"); - } - if (pipeline.get("decoder") != null) { - pipeline.remove("decoder"); - } - if (pipeline.get("requestReadMonitor") != null) { - pipeline.remove("requestReadMonitor"); - } - tunneling = true; - return channel.newSucceededFuture(); - } catch (Throwable t) { - return channel.newFailedFuture(t); - } + protected ChannelFuture execute() { + try { + ChannelPipeline pipeline = ctx.pipeline(); + removeHandlerIfPresent(pipeline, "encoder"); + removeHandlerIfPresent(pipeline, "responseWrittenMonitor"); + removeHandlerIfPresent(pipeline, "decoder"); + removeHandlerIfPresent(pipeline, "requestReadMonitor"); + tunneling = true; + return channel.newSucceededFuture(); + } catch (Throwable t) { + return channel.newFailedFuture(t); + } } + }; + + /** + * Encrypts traffic on this connection with SSL/TLS. + * + * @param sslEngine the {@link SSLEngine} for doing the encryption + * @param authenticateClients when {@code true}, client authentication is required; when {@code + * false}, the engine's client-auth configuration is left as-is (see {@link + * #encrypt(ChannelPipeline, SSLEngine, boolean)}) + * @return a Future for when the SSL handshake has completed + */ + protected Future encrypt(SSLEngine sslEngine, boolean authenticateClients) { + return encrypt(ctx.pipeline(), sslEngine, authenticateClients); + } + + /** + * Encrypts traffic on this connection with SSL/TLS. + * + * @param pipeline the ChannelPipeline on which to enable encryption + * @param sslEngine the {@link SSLEngine} for doing the encryption + * @param authenticateClients when {@code true}, client authentication is required ({@link + * SSLEngine#setNeedClientAuth(boolean)}). When {@code false}, the client-authentication + * configuration of the supplied {@link SSLEngine} is left untouched, so callers can opt into + * requesting (but not requiring) a client certificate via their {@link + * org.littleshoot.proxy.SslEngineSource} (for example {@link + * SSLEngine#setWantClientAuth(boolean)} or Netty's {@code ClientAuth.OPTIONAL}). + * @return a Future for when the SSL handshake has completed + */ + protected Future encrypt( + ChannelPipeline pipeline, SSLEngine sslEngine, boolean authenticateClients) { + LOG.debug("Enabling encryption with SSLEngine: {}", sslEngine); + this.sslEngine = sslEngine; + sslEngine.setUseClientMode(runsAsSslClient); + if (authenticateClients) { + sslEngine.setNeedClientAuth(true); + } + if (null != channel) { + channel.config().setAutoRead(true); + } + SslHandler handler = new SslHandler(sslEngine); + if (pipeline.get("ssl") == null) { + pipeline.addFirst("ssl", handler); + } else { + // The second SSL handler is added to handle the case + // where the proxy (running as MITM) has to chain with + // another SSL enabled proxy. The second SSL handler + // is to perform SSL with the server. + pipeline.addAfter("ssl", "sslWithServer", handler); + } + return handler.handshakeFuture(); + } + + /** + * Encrypts the channel using the provided {@link SSLEngine}. + * + * @param sslEngine the {@link SSLEngine} for doing the encryption + */ + protected ConnectionFlowStep EncryptChannel(final SSLEngine sslEngine) { + return new ConnectionFlowStep<>(this, HANDSHAKING) { + @Override + boolean shouldExecuteOnEventLoop() { + return false; + } + + @Override + protected Future execute() { + return encrypt(sslEngine, !runsAsSslClient); + } }; - - /** - * Encrypts traffic on this connection with SSL/TLS. - * - * @param sslEngine - * the {@link SSLEngine} for doing the encryption - * @param authenticateClients - * determines whether to authenticate clients or not - * @return a Future for when the SSL handshake has completed - */ - protected Future encrypt(SSLEngine sslEngine, - boolean authenticateClients) { - return encrypt(ctx.pipeline(), sslEngine, authenticateClients); - } - - /** - * Encrypts traffic on this connection with SSL/TLS. - * - * @param pipeline - * the ChannelPipeline on which to enable encryption - * @param sslEngine - * the {@link SSLEngine} for doing the encryption - * @param authenticateClients - * determines whether to authenticate clients or not - * @return a Future for when the SSL handshake has completed - */ - protected Future encrypt(ChannelPipeline pipeline, - SSLEngine sslEngine, - boolean authenticateClients) { - LOG.debug("Enabling encryption with SSLEngine: {}", - sslEngine); - this.sslEngine = sslEngine; - sslEngine.setUseClientMode(runsAsSslClient); - sslEngine.setNeedClientAuth(authenticateClients); - if (null != channel) { - channel.config().setAutoRead(true); - } - SslHandler handler = new SslHandler(sslEngine); - if(pipeline.get("ssl") == null) { - pipeline.addFirst("ssl", handler); - } else { - // The second SSL handler is added to handle the case - // where the proxy (running as MITM) has to chain with - // another SSL enabled proxy. The second SSL handler - // is to perform SSL with the server. - pipeline.addAfter("ssl", "sslWithServer", handler); - } - return handler.handshakeFuture(); - } - - /** - * Encrypts the channel using the provided {@link SSLEngine}. - * - * @param sslEngine - * the {@link SSLEngine} for doing the encryption - */ - protected ConnectionFlowStep EncryptChannel(final SSLEngine sslEngine) { - return new ConnectionFlowStep(this, HANDSHAKING) { - @Override - boolean shouldExecuteOnEventLoop() { - return false; - } - - @Override - protected Future execute() { - return encrypt(sslEngine, !runsAsSslClient); - } - }; - } - - /** - * Enables decompression and aggregation of content, which is useful for - * certain types of filtering activity. - */ - protected void aggregateContentForFiltering(ChannelPipeline pipeline, - int numberOfBytesToBuffer) { - pipeline.addLast("inflater", new HttpContentDecompressor()); - pipeline.addLast("aggregator", new HttpObjectAggregator( - numberOfBytesToBuffer)); - } - - /** - * Callback that's invoked if this connection becomes saturated. - */ - protected void becameSaturated() { - LOG.debug("Became saturated"); - } - - /** - * Callback that's invoked when this connection becomes writeable again. - */ - protected void becameWritable() { - LOG.debug("Became writeable"); - } - - /** - * Override this to handle exceptions that occurred during asynchronous - * processing on the {@link Channel}. - */ - protected void exceptionCaught(Throwable cause) { - } - - /* ************************************************************************* - * State/Management - **************************************************************************/ - /** - * Disconnects. This will wait for pending writes to be flushed before - * disconnecting. - * - * @return {@code Future} for when we're done disconnecting. If we weren't - * connected, this returns null. - */ - Future disconnect() { - if (channel == null) { - return null; - } else { - final Promise promise = channel.newPromise(); - writeToChannel(Unpooled.EMPTY_BUFFER).addListener( - future -> closeChannel(promise)); - return promise; - } - } - - private void closeChannel(final Promise promise) { - channel.close().addListener( - future -> { - if (future - .isSuccess()) { - promise.setSuccess(null); - } else { - promise.setFailure(future - .cause()); - } - }); - } - - /** - * Indicates whether or not this connection is saturated (i.e. not - * writeable). - */ - protected boolean isSaturated() { - return !this.channel.isWritable(); - } - - /** - * Utility for checking current state. - */ - protected boolean is(ConnectionState state) { - return currentState == state; - } - - /** - * If this connection is currently in the process of going through a - * {@link ConnectionFlow}, this will return true. - */ - protected boolean isConnecting() { - return currentState.isPartOfConnectionFlow(); - } - - /** - * Updates the current state to the given value. - */ - protected void become(ConnectionState state) { - this.currentState = state; - } - - protected ConnectionState getCurrentState() { - return currentState; - } - - public boolean isTunneling() { - return tunneling; - } - - public SSLEngine getSslEngine() { - return sslEngine; - } - - /** - * Call this to stop reading. - */ - protected void stopReading() { - LOG.debug("Stopped reading"); - this.channel.config().setAutoRead(false); - } - - /** - * Call this to resume reading. - */ - protected void resumeReading() { - LOG.debug("Resumed reading"); - this.channel.config().setAutoRead(true); - } - - /** - * Request the ProxyServer for Filters. - * - * By default, no-op filters are returned by DefaultHttpProxyServer. - * Subclasses of ProxyConnection can change this behaviour. - * - * @param httpRequest - * Filter attached to the give HttpRequest (if any) - */ - protected HttpFilters getHttpFiltersFromProxyServer(HttpRequest httpRequest) { - return proxyServer.getFiltersSource().filterRequest(httpRequest, ctx); - } - - ProxyConnectionLogger getLOG() { - return LOG; - } - - /* ************************************************************************* - * Adapting the Netty API - **************************************************************************/ + } + + /** + * Enables decompression and aggregation of content, which is useful for certain types of + * filtering activity. + */ + protected void aggregateContentForFiltering(ChannelPipeline pipeline, int numberOfBytesToBuffer) { + pipeline.addLast("inflater", new HttpContentDecompressor(false, 0)); + pipeline.addLast("aggregator", new HttpObjectAggregator(numberOfBytesToBuffer)); + } + + /** Callback that's invoked if this connection becomes saturated. */ + protected void becameSaturated() { + LOG.debug("Became saturated"); + } + + /** Callback that's invoked when this connection becomes writeable again. */ + protected void becameWritable() { + LOG.debug("Became writeable"); + } + + /** + * Override this to handle exceptions that occurred during asynchronous processing on the {@link + * Channel}. + */ + protected void exceptionCaught(Throwable cause) {} + + /** + * Removes the handler with the given name if it is present in the pipeline. + * + * @param pipeline the pipeline from which to remove the handler. + * @param handlerName the name of the handler to remove. + */ + protected void removeHandlerIfPresent(ChannelPipeline pipeline, String handlerName) { + if (pipeline.get(handlerName) != null) { + pipeline.remove(handlerName); + } + } + + /* + * ************************************************************************* + * State/Management + **************************************************************************/ + /** + * Disconnects. This will wait for pending writes to be flushed before disconnecting. + * + * @return {@code Future} for when we're done disconnecting. If we weren't connected, this + * returns null. + */ + @Nullable Future disconnect() { + if (channel == null) { + return null; + } else { + final Promise promise = channel.newPromise(); + writeToChannel(Unpooled.EMPTY_BUFFER).addListener(future -> closeChannel(promise)); + return promise; + } + } + + private void closeChannel(final Promise promise) { + channel + .close() + .addListener( + future -> { + if (future.isSuccess()) { + promise.setSuccess(null); + } else { + promise.setFailure(future.cause()); + } + }); + } + + /** Indicates whether this connection is saturated (i.e. not writeable). */ + protected boolean isSaturated() { + return !channel.isWritable(); + } + + /** Utility for checking current state. */ + protected boolean is(ConnectionState state) { + return currentState == state; + } + + /** + * If this connection is currently in the process of going through a {@link ConnectionFlow}, this + * will return true. + */ + protected boolean isConnecting() { + return currentState.isPartOfConnectionFlow(); + } + + /** Updates the current state to the given value. */ + protected void become(ConnectionState state) { + currentState = state; + } + + protected ConnectionState getCurrentState() { + return currentState; + } + + public boolean isTunneling() { + return tunneling; + } + + @Nullable + public SSLEngine getSslEngine() { + return sslEngine; + } + + /** Call this to stop reading. */ + protected void stopReading() { + LOG.debug("Stopped reading"); + channel.config().setAutoRead(false); + } + + /** Call this to resume reading. */ + protected void resumeReading() { + LOG.debug("Resumed reading"); + channel.config().setAutoRead(true); + } + + /** + * Request the ProxyServer for Filters. + * + *

By default, no-op filters are returned by DefaultHttpProxyServer. Subclasses of + * ProxyConnection can change this behaviour. + * + * @param httpRequest Filter attached to the give HttpRequest (if any) + */ + @Nullable + protected HttpFilters getHttpFiltersFromProxyServer(HttpRequest httpRequest) { + return proxyServer.getFiltersSource().filterRequest(httpRequest, ctx); + } + + ProxyConnectionLogger getLOG() { + return LOG; + } + + /* + * ************************************************************************* + * Adapting the Netty API + **************************************************************************/ + @Override + protected final void channelRead0(ChannelHandlerContext ctx, Object msg) { + read(msg); + } + + @Override + public void handlerAdded(ChannelHandlerContext ctx) throws Exception { + this.ctx = ctx; + channel = ctx.channel(); + super.handlerAdded(ctx); + } + + @Override + public void channelRegistered(ChannelHandlerContext ctx) throws Exception { + try { + this.ctx = ctx; + channel = ctx.channel(); + proxyServer.registerChannel(ctx.channel()); + } finally { + super.channelRegistered(ctx); + } + } + + @Override + public void channelUnregistered(ChannelHandlerContext ctx) throws Exception { + proxyServer.unregisterChannel(ctx.channel()); + super.channelUnregistered(ctx); + } + + /** Only once the Netty Channel is active to we recognize the ProxyConnection as connected. */ + @Override + public final void channelActive(ChannelHandlerContext ctx) throws Exception { + try { + connected(); + } finally { + super.channelActive(ctx); + } + } + + /** As soon as the Netty Channel is inactive, we recognize the ProxyConnection as disconnected. */ + @Override + public void channelInactive(ChannelHandlerContext ctx) throws Exception { + try { + disconnected(); + } finally { + super.channelInactive(ctx); + } + } + + @Override + public final void channelWritabilityChanged(ChannelHandlerContext ctx) throws Exception { + LOG.debug("Writability changed. Is writable: {}", channel.isWritable()); + try { + if (channel.isWritable()) { + becameWritable(); + } else { + becameSaturated(); + } + } finally { + super.channelWritabilityChanged(ctx); + } + } + + @Override + public final void exceptionCaught(ChannelHandlerContext ctx, Throwable cause) { + exceptionCaught(cause); + } + + /** + * We're looking for {@link IdleStateEvent}s to see if we need to disconnect. + * + *

Note - we don't care what kind of IdleState we got. Thanks to bast for pointing this out. + */ + @Override + public final void userEventTriggered(ChannelHandlerContext ctx, Object evt) throws Exception { + try { + if (evt instanceof IdleStateEvent) { + LOG.debug("Got idle"); + timedOut(); + } + } finally { + super.userEventTriggered(ctx, evt); + } + } + + /* + * ************************************************************************* + * Activity Tracking/Statistics + **************************************************************************/ + + /** Utility handler for monitoring bytes read on this connection. */ + @Sharable + protected abstract class BytesReadMonitor extends ChannelInboundHandlerAdapter { @Override - protected final void channelRead0(ChannelHandlerContext ctx, Object msg) { - read(msg); - } - - @Override - public void channelRegistered(ChannelHandlerContext ctx) throws Exception { - try { - this.ctx = ctx; - this.channel = ctx.channel(); - this.proxyServer.registerChannel(ctx.channel()); - } finally { - super.channelRegistered(ctx); + public void channelRead(ChannelHandlerContext ctx, Object msg) throws Exception { + try { + if (msg instanceof ByteBuf) { + bytesRead(((ByteBuf) msg).readableBytes()); } + } catch (Throwable t) { + LOG.warn("Unable to record bytesRead", t); + } finally { + super.channelRead(ctx, msg); + } } - /** - * Only once the Netty Channel is active to we recognize the ProxyConnection - * as connected. - */ - @Override - public final void channelActive(ChannelHandlerContext ctx) throws Exception { - try { - connected(); - } finally { - super.channelActive(ctx); - } - } + protected abstract void bytesRead(int numberOfBytes); + } - /** - * As soon as the Netty Channel is inactive, we recognize the - * ProxyConnection as disconnected. - */ + /** Utility handler for monitoring requests read on this connection. */ + @Sharable + protected abstract class RequestReadMonitor extends ChannelInboundHandlerAdapter { @Override - public void channelInactive(ChannelHandlerContext ctx) throws Exception { - try { - disconnected(); - } finally { - super.channelInactive(ctx); + public void channelRead(ChannelHandlerContext ctx, Object msg) throws Exception { + try { + if (msg instanceof HttpRequest) { + requestRead((HttpRequest) msg); } + } catch (Throwable t) { + LOG.warn("Unable to record bytesRead", t); + } finally { + super.channelRead(ctx, msg); + } } - @Override - public final void channelWritabilityChanged(ChannelHandlerContext ctx) - throws Exception { - LOG.debug("Writability changed. Is writable: {}", channel.isWritable()); - try { - if (this.channel.isWritable()) { - becameWritable(); - } else { - becameSaturated(); - } - } finally { - super.channelWritabilityChanged(ctx); - } - } - - @Override - public final void exceptionCaught(ChannelHandlerContext ctx, Throwable cause) { - exceptionCaught(cause); - } + protected abstract void requestRead(HttpRequest httpRequest); + } - /** - *

- * We're looking for {@link IdleStateEvent}s to see if we need to - * disconnect. - *

- * - *

- * Note - we don't care what kind of IdleState we got. Thanks to qbast for pointing this out. - *

- */ + /** Utility handler for monitoring responses read on this connection. */ + @Sharable + protected abstract class ResponseReadMonitor extends ChannelInboundHandlerAdapter { @Override - public final void userEventTriggered(ChannelHandlerContext ctx, Object evt) - throws Exception { - try { - if (evt instanceof IdleStateEvent) { - LOG.debug("Got idle"); - timedOut(); - } - } finally { - super.userEventTriggered(ctx, evt); + public void channelRead(ChannelHandlerContext ctx, Object msg) throws Exception { + try { + if (msg instanceof HttpResponse) { + responseRead((HttpResponse) msg); } + } catch (Throwable t) { + LOG.warn("Unable to record bytesRead", t); + } finally { + super.channelRead(ctx, msg); + } } - /* ************************************************************************* - * Activity Tracking/Statistics - **************************************************************************/ + protected abstract void responseRead(HttpResponse httpResponse); + } - /** - * Utility handler for monitoring bytes read on this connection. - */ - @Sharable - protected abstract class BytesReadMonitor extends - ChannelInboundHandlerAdapter { - @Override - public void channelRead(ChannelHandlerContext ctx, Object msg) - throws Exception { - try { - if (msg instanceof ByteBuf) { - bytesRead(((ByteBuf) msg).readableBytes()); - } - } catch (Throwable t) { - LOG.warn("Unable to record bytesRead", t); - } finally { - super.channelRead(ctx, msg); - } + /** Utility handler for monitoring bytes written on this connection. */ + @Sharable + protected abstract class BytesWrittenMonitor extends ChannelOutboundHandlerAdapter { + @Override + public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise promise) + throws Exception { + try { + if (msg instanceof ByteBuf) { + bytesWritten(((ByteBuf) msg).readableBytes()); } - - protected abstract void bytesRead(int numberOfBytes); + } catch (Throwable t) { + LOG.warn("Unable to record bytesRead", t); + } finally { + super.write(ctx, msg, promise); + } } - /** - * Utility handler for monitoring requests read on this connection. - */ - @Sharable - protected abstract class RequestReadMonitor extends - ChannelInboundHandlerAdapter { - @Override - public void channelRead(ChannelHandlerContext ctx, Object msg) - throws Exception { - try { - if (msg instanceof HttpRequest) { - requestRead((HttpRequest) msg); - } - } catch (Throwable t) { - LOG.warn("Unable to record bytesRead", t); - } finally { - super.channelRead(ctx, msg); - } - } + protected abstract void bytesWritten(int numberOfBytes); + } - protected abstract void requestRead(HttpRequest httpRequest); - } + /** Utility handler for monitoring requests written on this connection. */ + @Sharable + protected abstract static class RequestWrittenMonitor extends ChannelOutboundHandlerAdapter { + @Override + public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise promise) + throws Exception { + HttpRequest originalRequest = null; + if (msg instanceof HttpRequest) { + originalRequest = (HttpRequest) msg; + } - /** - * Utility handler for monitoring responses read on this connection. - */ - @Sharable - protected abstract class ResponseReadMonitor extends - ChannelInboundHandlerAdapter { - @Override - public void channelRead(ChannelHandlerContext ctx, Object msg) - throws Exception { - try { - if (msg instanceof HttpResponse) { - responseRead((HttpResponse) msg); - } - } catch (Throwable t) { - LOG.warn("Unable to record bytesRead", t); - } finally { - super.channelRead(ctx, msg); - } - } + if (null != originalRequest) { + requestWriting(originalRequest); + } - protected abstract void responseRead(HttpResponse httpResponse); - } + super.write(ctx, msg, promise); - /** - * Utility handler for monitoring bytes written on this connection. - */ - @Sharable - protected abstract class BytesWrittenMonitor extends - ChannelOutboundHandlerAdapter { - @Override - public void write(ChannelHandlerContext ctx, - Object msg, ChannelPromise promise) - throws Exception { - try { - if (msg instanceof ByteBuf) { - bytesWritten(((ByteBuf) msg).readableBytes()); - } - } catch (Throwable t) { - LOG.warn("Unable to record bytesRead", t); - } finally { - super.write(ctx, msg, promise); - } - } + if (null != originalRequest) { + requestWritten(originalRequest); + } - protected abstract void bytesWritten(int numberOfBytes); + if (msg instanceof HttpContent) { + contentWritten((HttpContent) msg); + } } - /** - * Utility handler for monitoring requests written on this connection. - */ - @Sharable - protected abstract class RequestWrittenMonitor extends - ChannelOutboundHandlerAdapter { - @Override - public void write(ChannelHandlerContext ctx, - Object msg, ChannelPromise promise) - throws Exception { - HttpRequest originalRequest = null; - if (msg instanceof HttpRequest) { - originalRequest = (HttpRequest) msg; - } - - if (null != originalRequest) { - requestWriting(originalRequest); - } - - super.write(ctx, msg, promise); - - if (null != originalRequest) { - requestWritten(originalRequest); - } - - if (msg instanceof HttpContent) { - contentWritten((HttpContent) msg); - } - } - - /** - * Invoked immediately before an HttpRequest is written. - */ - protected abstract void requestWriting(HttpRequest httpRequest); + /** Invoked immediately before an HttpRequest is written. */ + protected abstract void requestWriting(HttpRequest httpRequest); - /** - * Invoked immediately after an HttpRequest has been sent. - */ - protected abstract void requestWritten(HttpRequest httpRequest); + /** Invoked immediately after an HttpRequest has been sent. */ + protected abstract void requestWritten(HttpRequest httpRequest); - /** - * Invoked immediately after an HttpContent has been sent. - */ - protected abstract void contentWritten(HttpContent httpContent); - } + /** Invoked immediately after an HttpContent has been sent. */ + protected abstract void contentWritten(HttpContent httpContent); + } - /** - * Utility handler for monitoring responses written on this connection. - */ - @Sharable - protected abstract class ResponseWrittenMonitor extends - ChannelOutboundHandlerAdapter { - @Override - public void write(ChannelHandlerContext ctx, - Object msg, ChannelPromise promise) - throws Exception { - try { - if (msg instanceof HttpResponse) { - responseWritten(((HttpResponse) msg)); - } - } catch (Throwable t) { - LOG.warn("Error while invoking responseWritten callback", t); - } finally { - super.write(ctx, msg, promise); - } + /** Utility handler for monitoring responses written on this connection. */ + @Sharable + protected abstract class ResponseWrittenMonitor extends ChannelOutboundHandlerAdapter { + @Override + public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise promise) + throws Exception { + try { + if (msg instanceof HttpResponse) { + responseWritten(((HttpResponse) msg)); } - - protected abstract void responseWritten(HttpResponse httpResponse); - } - + } catch (Throwable t) { + LOG.warn("Error while invoking responseWritten callback", t); + } finally { + super.write(ctx, msg, promise); + } + } + + protected abstract void responseWritten(HttpResponse httpResponse); + } + + /** + * Gets the channel handler context + * + * @return the channel handler context + */ + public ChannelHandlerContext getContext() { + return ctx; + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ProxyConnectionLogger.java b/src/main/java/org/littleshoot/proxy/impl/ProxyConnectionLogger.java index e6eac83e..cb387cf6 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ProxyConnectionLogger.java +++ b/src/main/java/org/littleshoot/proxy/impl/ProxyConnectionLogger.java @@ -1,184 +1,166 @@ package org.littleshoot.proxy.impl; +import java.util.Arrays; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.slf4j.helpers.MessageFormatter; import org.slf4j.spi.LocationAwareLogger; -import java.util.Arrays; - /** - *

- * A helper class that logs messages for ProxyConnections. All it does is make - * sure that the Channel and current state are always included in the log - * messages (if available). - *

+ * A helper class that logs messages for ProxyConnections. All it does is make sure that the Channel + * and current state are always included in the log messages (if available). * - *

- * Note that this depends on us using a LocationAwareLogger so that we can - * report the line numbers of the caller rather than this helper class. - * If the SLF4J binding does not provide a LocationAwareLogger, then a fallback - * to Logger is provided. - *

+ *

Note that this depends on us using a LocationAwareLogger so that we can report the line + * numbers of the caller rather than this helper class. If the SLF4J binding does not provide a + * LocationAwareLogger, then a fallback to Logger is provided. */ class ProxyConnectionLogger { - private final ProxyConnection connection; - private final LogDispatch dispatch; - private final Logger logger; - private final String fqcn = this.getClass().getCanonicalName(); - - public ProxyConnectionLogger(ProxyConnection connection) { - this.connection = connection; - final Logger lg = LoggerFactory.getLogger(connection - .getClass()); - if (lg instanceof LocationAwareLogger) { - dispatch = new LocationAwareLogggerDispatch((LocationAwareLogger) lg); - } - else { - dispatch = new LoggerDispatch(); - } - logger = lg; + private final ProxyConnection connection; + private final LogDispatch dispatch; + private final Logger logger; + private final String fqcn = getClass().getCanonicalName(); + + public ProxyConnectionLogger(ProxyConnection connection) { + this.connection = connection; + logger = LoggerFactory.getLogger(connection.getClass()); + dispatch = + logger instanceof LocationAwareLogger + ? new LocationAwareLoggerDispatch((LocationAwareLogger) logger) + : new LoggerDispatch(); + } + + protected void error(String message, Object... params) { + if (logger.isErrorEnabled()) { + dispatch.doLog(LocationAwareLogger.ERROR_INT, message, params, null); } + } - protected void error(String message, Object... params) { - if (logger.isErrorEnabled()) { - dispatch.doLog(LocationAwareLogger.ERROR_INT, message, params, null); - } + protected void error(String message, Throwable t) { + if (logger.isErrorEnabled()) { + dispatch.doLog(LocationAwareLogger.ERROR_INT, message, null, t); } + } - protected void error(String message, Throwable t) { - if (logger.isErrorEnabled()) { - dispatch.doLog(LocationAwareLogger.ERROR_INT, message, null, t); - } + protected void warn(String message, Object... params) { + if (logger.isWarnEnabled()) { + dispatch.doLog(LocationAwareLogger.WARN_INT, message, params, null); } + } - protected void warn(String message, Object... params) { - if (logger.isWarnEnabled()) { - dispatch.doLog(LocationAwareLogger.WARN_INT, message, params, null); - } + protected void warn(String message, Throwable t) { + if (logger.isWarnEnabled()) { + dispatch.doLog(LocationAwareLogger.WARN_INT, message, null, t); } + } - protected void warn(String message, Throwable t) { - if (logger.isWarnEnabled()) { - dispatch.doLog(LocationAwareLogger.WARN_INT, message, null, t); - } + protected void info(String message, Object... params) { + if (logger.isInfoEnabled()) { + dispatch.doLog(LocationAwareLogger.INFO_INT, message, params, null); } + } - protected void info(String message, Object... params) { - if (logger.isInfoEnabled()) { - dispatch.doLog(LocationAwareLogger.INFO_INT, message, params, null); - } + protected void info(String message, Throwable t) { + if (logger.isInfoEnabled()) { + dispatch.doLog(LocationAwareLogger.INFO_INT, message, null, t); } + } - protected void info(String message, Throwable t) { - if (logger.isInfoEnabled()) { - dispatch.doLog(LocationAwareLogger.INFO_INT, message, null, t); - } + protected void debug(String message, Object... params) { + if (logger.isDebugEnabled()) { + dispatch.doLog(LocationAwareLogger.DEBUG_INT, message, params, null); } + } - protected void debug(String message, Object... params) { - if (logger.isDebugEnabled()) { - dispatch.doLog(LocationAwareLogger.DEBUG_INT, message, params, null); - } + protected void debug(String message, Throwable t) { + if (logger.isDebugEnabled()) { + dispatch.doLog(LocationAwareLogger.DEBUG_INT, message, null, t); } + } - protected void debug(String message, Throwable t) { - if (logger.isDebugEnabled()) { - dispatch.doLog(LocationAwareLogger.DEBUG_INT, message, null, t); - } + protected void log(int level, String message, Object... params) { + if (level != LocationAwareLogger.DEBUG_INT || logger.isDebugEnabled()) { + dispatch.doLog(level, message, params, null); } + } - protected void log(int level, String message, Object... params) { - if (level != LocationAwareLogger.DEBUG_INT || logger.isDebugEnabled()) { - dispatch.doLog(level, message, params, null); - } + protected void log(int level, String message, Throwable t) { + if (level != LocationAwareLogger.DEBUG_INT || logger.isDebugEnabled()) { + dispatch.doLog(level, message, null, t); } + } - protected void log(int level, String message, Throwable t) { - if (level != LocationAwareLogger.DEBUG_INT || logger.isDebugEnabled()) { - dispatch.doLog(level, message, null, t); - } - } + private interface LogDispatch { + void doLog(int level, String message, Object[] params, Throwable t); + } - private interface LogDispatch { - void doLog(int level, String message, Object[] params, Throwable t); + private String fullMessage(String message) { + String stateMessage = connection.getCurrentState().toString(); + if (connection.isTunneling()) { + stateMessage += " {tunneling}"; } - - private String fullMessage(String message) { - String stateMessage = connection.getCurrentState().toString(); - if (connection.isTunneling()) { - stateMessage += " {tunneling}"; - } - String messagePrefix = "(" + stateMessage + ")"; - if (connection.channel != null) { - messagePrefix = messagePrefix + " " + connection.channel; - } - return messagePrefix + ": " + message; + String messagePrefix = "(" + stateMessage + ")"; + if (connection.channel != null) { + messagePrefix = messagePrefix + " " + connection.channel; } - - /** - * Fallback dispatch if a LocationAwareLogger is not available from - * the SLF4J LoggerFactory. - */ - private class LoggerDispatch implements LogDispatch { - @Override - public void doLog(int level, String message, Object[] params, Throwable t) { - String formattedMessage = fullMessage(message); - - final Object[] paramsWithThrowable; - - if (t != null) { - if (params == null) { - paramsWithThrowable = new Object[1]; - paramsWithThrowable[0] = t; - } else { - paramsWithThrowable = Arrays.copyOf(params, params.length + 1); - paramsWithThrowable[params.length] = t; - } - } - else { - paramsWithThrowable = params; - } - switch (level) { - case LocationAwareLogger.TRACE_INT: - logger.trace(formattedMessage, paramsWithThrowable); - break; - case LocationAwareLogger.DEBUG_INT: - logger.debug(formattedMessage, paramsWithThrowable); - break; - case LocationAwareLogger.INFO_INT: - logger.info(formattedMessage, paramsWithThrowable); - break; - case LocationAwareLogger.WARN_INT: - logger.warn(formattedMessage, paramsWithThrowable); - break; - case LocationAwareLogger.ERROR_INT: - default: - logger.error(formattedMessage, paramsWithThrowable); - break; - } - } + return messagePrefix + ": " + message; + } + + /** Fallback dispatch if a LocationAwareLogger is not available from the SLF4J LoggerFactory. */ + private class LoggerDispatch implements LogDispatch { + @Override + public void doLog(int level, String message, Object[] params, Throwable t) { + String formattedMessage = fullMessage(message); + + final Object[] paramsWithThrowable; + + if (t != null) { + if (params == null) { + paramsWithThrowable = new Object[1]; + paramsWithThrowable[0] = t; + } else { + paramsWithThrowable = Arrays.copyOf(params, params.length + 1); + paramsWithThrowable[params.length] = t; + } + } else { + paramsWithThrowable = params; + } + switch (level) { + case LocationAwareLogger.TRACE_INT: + logger.trace(formattedMessage, paramsWithThrowable); + break; + case LocationAwareLogger.DEBUG_INT: + logger.debug(formattedMessage, paramsWithThrowable); + break; + case LocationAwareLogger.INFO_INT: + logger.info(formattedMessage, paramsWithThrowable); + break; + case LocationAwareLogger.WARN_INT: + logger.warn(formattedMessage, paramsWithThrowable); + break; + case LocationAwareLogger.ERROR_INT: + default: + logger.error(formattedMessage, paramsWithThrowable); + break; + } } + } - /** - * Dispatcher for a LocationAwareLogger. - */ - private class LocationAwareLogggerDispatch implements LogDispatch { + /** Dispatcher for a LocationAwareLogger. */ + private class LocationAwareLoggerDispatch implements LogDispatch { - private LocationAwareLogger log; + private final LocationAwareLogger log; - public LocationAwareLogggerDispatch(LocationAwareLogger log) { - this.log = log; - } + public LocationAwareLoggerDispatch(LocationAwareLogger log) { + this.log = log; + } - @Override - public void doLog(int level, String message, Object[] params, Throwable t) { - String formattedMessage = fullMessage(message); - if (params != null && params.length > 0) { - formattedMessage = MessageFormatter.arrayFormat(formattedMessage, - params).getMessage(); - } - log.log(null, fqcn, level, formattedMessage, null, t); - } + @Override + public void doLog(int level, String message, Object[] params, Throwable t) { + String formattedMessage = fullMessage(message); + if (params != null && params.length > 0) { + formattedMessage = MessageFormatter.arrayFormat(formattedMessage, params).getMessage(); + } + log.log(null, fqcn, level, formattedMessage, null, t); } -} \ No newline at end of file + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/ProxyThreadPools.java b/src/main/java/org/littleshoot/proxy/impl/ProxyThreadPools.java index 4466b9d5..de550e3a 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ProxyThreadPools.java +++ b/src/main/java/org/littleshoot/proxy/impl/ProxyThreadPools.java @@ -1,64 +1,78 @@ package org.littleshoot.proxy.impl; -import com.google.common.collect.ImmutableList; import io.netty.channel.EventLoopGroup; import io.netty.channel.nio.NioEventLoopGroup; - import java.nio.channels.spi.SelectorProvider; import java.util.List; /** - * Encapsulates the thread pools used by the proxy. Contains the acceptor thread pool as well as the client-to-proxy and - * proxy-to-server thread pools. + * Encapsulates the thread pools used by the proxy. Contains the acceptor thread pool as well as the + * client-to-proxy and proxy-to-server thread pools. */ public class ProxyThreadPools { - /** - * These {@link EventLoopGroup}s accept incoming connections to the - * proxies. A different EventLoopGroup is used for each - * TransportProtocol, since these have to be configured differently. - */ - private final NioEventLoopGroup clientToProxyAcceptorPool; + /** + * These {@link EventLoopGroup}s accept incoming connections to the proxies. A different + * EventLoopGroup is used for each TransportProtocol, since these have to be configured + * differently. + */ + private final NioEventLoopGroup clientToProxyAcceptorPool; - /** - * These {@link EventLoopGroup}s process incoming requests to the - * proxies. A different EventLoopGroup is used for each - * TransportProtocol, since these have to be configured differently. - */ - private final NioEventLoopGroup clientToProxyWorkerPool; + /** + * These {@link EventLoopGroup}s process incoming requests to the proxies. A different + * EventLoopGroup is used for each TransportProtocol, since these have to be configured + * differently. + */ + private final NioEventLoopGroup clientToProxyWorkerPool; - /** - * These {@link EventLoopGroup}s are used for making outgoing - * connections to servers. A different EventLoopGroup is used for each - * TransportProtocol, since these have to be configured differently. - */ - private final NioEventLoopGroup proxyToServerWorkerPool; + /** + * These {@link EventLoopGroup}s are used for making outgoing connections to servers. A different + * EventLoopGroup is used for each TransportProtocol, since these have to be configured + * differently. + */ + private final NioEventLoopGroup proxyToServerWorkerPool; - public ProxyThreadPools(SelectorProvider selectorProvider, int incomingAcceptorThreads, int incomingWorkerThreads, int outgoingWorkerThreads, String serverGroupName, int serverGroupId) { - clientToProxyAcceptorPool = new NioEventLoopGroup(incomingAcceptorThreads, new CategorizedThreadFactory(serverGroupName, "ClientToProxyAcceptor", serverGroupId), selectorProvider); + public ProxyThreadPools( + SelectorProvider selectorProvider, + int incomingAcceptorThreads, + int incomingWorkerThreads, + int outgoingWorkerThreads, + String serverGroupName, + int serverGroupId) { + clientToProxyAcceptorPool = + new NioEventLoopGroup( + incomingAcceptorThreads, + new CategorizedThreadFactory(serverGroupName, "ClientToProxyAcceptor", serverGroupId), + selectorProvider); - clientToProxyWorkerPool = new NioEventLoopGroup(incomingWorkerThreads, new CategorizedThreadFactory(serverGroupName, "ClientToProxyWorker", serverGroupId), selectorProvider); - clientToProxyWorkerPool.setIoRatio(90); + clientToProxyWorkerPool = + new NioEventLoopGroup( + incomingWorkerThreads, + new CategorizedThreadFactory(serverGroupName, "ClientToProxyWorker", serverGroupId), + selectorProvider); + clientToProxyWorkerPool.setIoRatio(90); - proxyToServerWorkerPool = new NioEventLoopGroup(outgoingWorkerThreads, new CategorizedThreadFactory(serverGroupName, "ProxyToServerWorker", serverGroupId), selectorProvider); - proxyToServerWorkerPool.setIoRatio(90); - } + proxyToServerWorkerPool = + new NioEventLoopGroup( + outgoingWorkerThreads, + new CategorizedThreadFactory(serverGroupName, "ProxyToServerWorker", serverGroupId), + selectorProvider); + proxyToServerWorkerPool.setIoRatio(90); + } - /** - * Returns all event loops (acceptor and worker thread pools) in this pool. - */ - public List getAllEventLoops() { - return ImmutableList.of(clientToProxyAcceptorPool, clientToProxyWorkerPool, proxyToServerWorkerPool); - } + /** Returns all event loops (acceptor and worker thread pools) in this pool. */ + public List getAllEventLoops() { + return List.of(clientToProxyAcceptorPool, clientToProxyWorkerPool, proxyToServerWorkerPool); + } - public NioEventLoopGroup getClientToProxyAcceptorPool() { - return clientToProxyAcceptorPool; - } + public NioEventLoopGroup getClientToProxyAcceptorPool() { + return clientToProxyAcceptorPool; + } - public NioEventLoopGroup getClientToProxyWorkerPool() { - return clientToProxyWorkerPool; - } + public NioEventLoopGroup getClientToProxyWorkerPool() { + return clientToProxyWorkerPool; + } - public NioEventLoopGroup getProxyToServerWorkerPool() { - return proxyToServerWorkerPool; - } + public NioEventLoopGroup getProxyToServerWorkerPool() { + return proxyToServerWorkerPool; + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ProxyToServerConnection.java b/src/main/java/org/littleshoot/proxy/impl/ProxyToServerConnection.java index 24ec47e1..ab2f08c5 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ProxyToServerConnection.java +++ b/src/main/java/org/littleshoot/proxy/impl/ProxyToServerConnection.java @@ -1,6 +1,16 @@ package org.littleshoot.proxy.impl; +import static java.util.Locale.ROOT; +import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_CHUNK; +import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_CONNECT_OK; +import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_INITIAL; +import static org.littleshoot.proxy.impl.ConnectionState.CONNECTING; +import static org.littleshoot.proxy.impl.ConnectionState.DISCONNECTED; +import static org.littleshoot.proxy.impl.ConnectionState.DISCONNECT_REQUESTED; +import static org.littleshoot.proxy.impl.ConnectionState.HANDSHAKING; + import com.google.common.net.HostAndPort; +import com.google.errorprone.annotations.CheckReturnValue; import io.netty.bootstrap.Bootstrap; import io.netty.buffer.ByteBuf; import io.netty.channel.Channel; @@ -12,8 +22,10 @@ import io.netty.channel.ChannelOption; import io.netty.channel.ChannelPipeline; import io.netty.channel.socket.nio.NioSocketChannel; -import io.netty.channel.udt.nio.NioUdtProvider; +import io.netty.handler.codec.haproxy.HAProxyCommand; import io.netty.handler.codec.haproxy.HAProxyMessage; +import io.netty.handler.codec.haproxy.HAProxyProtocolVersion; +import io.netty.handler.codec.haproxy.HAProxyProxiedProtocol; import io.netty.handler.codec.http.FullHttpResponse; import io.netty.handler.codec.http.HttpContent; import io.netty.handler.codec.http.HttpMessage; @@ -56,1214 +68,1961 @@ import io.netty.resolver.DefaultAddressResolverGroup; import io.netty.util.ReferenceCounted; import io.netty.util.concurrent.Future; +import java.io.IOException; +import java.net.InetSocketAddress; +import java.net.UnknownHostException; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Queue; +import java.util.concurrent.ConcurrentLinkedQueue; +import java.util.concurrent.RejectedExecutionException; +import javax.net.ssl.SSLEngine; +import javax.net.ssl.SSLHandshakeException; +import javax.net.ssl.SSLProtocolException; +import javax.net.ssl.SSLSession; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; import org.littleshoot.proxy.ActivityTracker; import org.littleshoot.proxy.ChainedProxy; import org.littleshoot.proxy.ChainedProxyAdapter; import org.littleshoot.proxy.ChainedProxyManager; import org.littleshoot.proxy.ChainedProxyType; +import org.littleshoot.proxy.FlowContext; import org.littleshoot.proxy.FullFlowContext; import org.littleshoot.proxy.HttpFilters; import org.littleshoot.proxy.MitmManager; import org.littleshoot.proxy.TransportProtocol; import org.littleshoot.proxy.UnknownTransportProtocolException; import org.littleshoot.proxy.extras.HAProxyMessageEncoder; - -import javax.net.ssl.SSLProtocolException; -import javax.net.ssl.SSLSession; -import java.io.IOException; -import java.net.InetSocketAddress; -import java.net.UnknownHostException; -import java.util.ArrayList; -import java.util.List; -import java.util.Queue; -import java.util.concurrent.ConcurrentLinkedQueue; -import java.util.concurrent.RejectedExecutionException; - -import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_CHUNK; -import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_CONNECT_OK; -import static org.littleshoot.proxy.impl.ConnectionState.AWAITING_INITIAL; -import static org.littleshoot.proxy.impl.ConnectionState.CONNECTING; -import static org.littleshoot.proxy.impl.ConnectionState.DISCONNECTED; -import static org.littleshoot.proxy.impl.ConnectionState.HANDSHAKING; +import org.littleshoot.proxy.extras.ProxyProtocolMessage; /** - *

- * Represents a connection from our proxy to a server on the web. - * ProxyConnections are reused fairly liberally, and can go from disconnected to - * connected, back to disconnected and so on. - *

- * - *

- * Connecting a {@link ProxyToServerConnection} can involve more than just - * connecting the underlying {@link Channel}. In particular, the connection may - * use encryption (i.e. TLS) and it may also establish an HTTP CONNECT tunnel. - * The various steps involved in fully establishing a connection are - * encapsulated in the property {@link #connectionFlow}, which is initialized in + * Represents a connection from our proxy to a server on the web. ProxyConnections are reused fairly + * liberally, and can go from disconnected to connected, back to disconnected and so on. + * + *

Connecting a {@link ProxyToServerConnection} can involve more than just connecting the + * underlying {@link Channel}. In particular, the connection may use encryption (i.e. TLS) and it + * may also establish an HTTP CONNECT tunnel. The various steps involved in fully establishing a + * connection are encapsulated in the property {@link #connectionFlow}, which is initialized in * {@link #initializeConnectionFlow()}. - *

*/ @Sharable +@NullMarked public class ProxyToServerConnection extends ProxyConnection { - private static final String SOCKS_ENCODER_NAME = "socksEncoder"; - private static final String SOCKS_DECODER_NAME = "socksDecoder"; - private final ClientToProxyConnection clientConnection; - private final ProxyToServerConnection serverConnection = this; - private volatile TransportProtocol transportProtocol; - private volatile ChainedProxyType chainedProxyType; - private volatile InetSocketAddress remoteAddress; - private volatile InetSocketAddress localAddress; - private volatile AddressResolverGroup remoteAddressResolver; - private volatile String username; - private volatile String password; - private final String serverHostAndPort; - private volatile ChainedProxy chainedProxy; - private final Queue availableChainedProxies; - - /** - * The filters to apply to response/chunks received from server. - */ - private volatile HttpFilters currentFilters; - - /** - * Encapsulates the flow for establishing a connection, which can vary - * depending on how things are configured. - */ - private volatile ConnectionFlow connectionFlow; - - /** - * Disables SNI when initializing connection flow in {@link #initializeConnectionFlow()}. This value is set to true - * when retrying a connection without SNI to work around Java's SNI handling issue (see - * {@link #connectionFailed(Throwable)}). - */ - private volatile boolean disableSni = false; - - /** - * While we're in the process of connecting, it's possible that we'll - * receive a new message to write. This lock helps us synchronize and wait - * for the connection to be established before writing the next message. - */ - private final Object connectLock = new Object(); - - /** - * This is the initial request received prior to connecting. We keep track - * of it so that we can process it after connection finishes. - */ - private volatile HttpRequest initialRequest; - - /** - * Keeps track of HttpRequests that have been issued so that we can - * associate them with responses that we get back - */ - private volatile HttpRequest currentHttpRequest; - - /** - * While we're doing a chunked transfer, this keeps track of the initial - * HttpResponse object for our transfer (which is useful for its headers). - */ - private volatile HttpResponse currentHttpResponse; - - /** - * Limits bandwidth when throttling is enabled. - */ - private volatile GlobalTrafficShapingHandler trafficHandler; - - /** - * Minimum size of the adaptive recv buffer when throttling is enabled. - */ - private static final int MINIMUM_RECV_BUFFER_SIZE_BYTES = 64; - - /** - * Create a new ProxyToServerConnection. - */ - static ProxyToServerConnection create(DefaultHttpProxyServer proxyServer, - ClientToProxyConnection clientConnection, - String serverHostAndPort, - HttpFilters initialFilters, - HttpRequest initialHttpRequest, - GlobalTrafficShapingHandler globalTrafficShapingHandler) - throws UnknownHostException { - Queue chainedProxies = new ConcurrentLinkedQueue<>(); - ChainedProxyManager chainedProxyManager = proxyServer - .getChainProxyManager(); - if (chainedProxyManager != null) { - chainedProxyManager.lookupChainedProxies(initialHttpRequest, - chainedProxies, clientConnection.getClientDetails()); - if (chainedProxies.size() == 0) { - // ChainedProxyManager returned no proxies, can't connect - return null; - } - } - return new ProxyToServerConnection(proxyServer, - clientConnection, - serverHostAndPort, - chainedProxies.poll(), - chainedProxies, - initialFilters, - globalTrafficShapingHandler); + // Pipeline handler names: + private static final String HTTP_ENCODER_NAME = "encoder"; + private static final String HTTP_DECODER_NAME = "decoder"; + private static final String HTTP_PROXY_ENCODER_NAME = "proxy-protocol-encoder"; + private static final String HTTP_REQUEST_WRITTEN_MONITOR_NAME = "requestWrittenMonitor"; + private static final String HTTP_RESPONSE_READ_MONITOR_NAME = "responseReadMonitor"; + private static final String SOCKS_ENCODER_NAME = "socksEncoder"; + private static final String SOCKS_DECODER_NAME = "socksDecoder"; + private static final String MAIN_HANDLER_NAME = "handler"; + private final ClientToProxyConnection clientConnection; + @Nullable private final ServerConnectionPool connectionPool; + private final ProxyToServerConnection serverConnection = this; + private volatile TransportProtocol transportProtocol; + private volatile ChainedProxyType chainedProxyType; + private volatile InetSocketAddress remoteAddress; + private volatile InetSocketAddress localAddress; + @Nullable private volatile AddressResolverGroup remoteAddressResolver; + @Nullable private volatile String username; + @Nullable private volatile String password; + private final String serverHostAndPort; + @Nullable private volatile ChainedProxy chainedProxy; + private final Queue availableChainedProxies; + + /** The filters to apply to response/chunks received from server. */ + private volatile HttpFilters currentFilters; + + /** + * Encapsulates the flow for establishing a connection, which can vary depending on how things are + * configured. + */ + @Nullable private volatile ConnectionFlow connectionFlow; + + /** + * Disables SNI when initializing connection flow in {@link #initializeConnectionFlow()}. This + * value is set to true when retrying a connection without SNI to work around Java's SNI handling + * issue (see {@link #connectionFailed(Throwable)}). + */ + private volatile boolean disableSni; + + /** + * Flag to skip SSL when connecting to server. This is set to true when retrying after SSL + * handshake fails with a non-SSL server (see {@link #connectionFailed(Throwable)}). + */ + private volatile boolean disableSslForNonTls; + + /** + * While we're in the process of connecting, it's possible that we'll receive a new message to + * write. This lock helps us synchronize and wait for the connection to be established before + * writing the next message. + */ + private final Object connectLock = new Object(); + + /** + * This is the initial request received prior to connecting. We keep track of it so that we can + * process it after connection finishes. + */ + @Nullable private volatile HttpRequest initialRequest; + + /** + * Keeps track of HttpRequests that have been issued so that we can associate them with responses + * that we get back + */ + @Nullable private volatile HttpRequest currentHttpRequest; + + /** + * While we're doing a chunked transfer, this keeps track of the initial HttpResponse object for + * our transfer (which is useful for its headers). + */ + @Nullable private volatile HttpResponse currentHttpResponse; + + /** + * For pooled connections, tracks the client connection that made the current request. This is set + * before writing a request and cleared after the response is complete. + */ + @Nullable private volatile ClientToProxyConnection currentClientConnectionForRequest; + + /** Limits bandwidth when throttling is enabled. */ + private final GlobalTrafficShapingHandler trafficHandler; + + /** + * When true, this connection is managed by the shared pool and has upstream TLS established + * (MITM). The connection should not be auto-released to pool after individual HTTP responses; it + * stays attached to the client's MITM session until the client disconnects. + */ + private volatile boolean mitmPooled; + + /** + * When true, this connection should be released to the pool immediately after the connection flow + * completes. Used in per-request MITM mode (Phase 2) so the CONNECT-created connection is + * available for subsequent HTTP requests via pool.getOrCreateConnection(). + */ + private volatile boolean releaseToPoolOnConnectComplete; + + /** Create a new ProxyToServerConnection. */ + @Nullable + @CheckReturnValue + static ProxyToServerConnection create( + DefaultHttpProxyServer proxyServer, + ClientToProxyConnection clientConnection, + String serverHostAndPort, + HttpFilters initialFilters, + HttpRequest initialHttpRequest, + GlobalTrafficShapingHandler globalTrafficShapingHandler) + throws UnknownHostException { + Queue chainedProxies = new ConcurrentLinkedQueue<>(); + ChainedProxyManager chainedProxyManager = proxyServer.getChainProxyManager(); + if (chainedProxyManager != null) { + chainedProxyManager.lookupChainedProxies( + initialHttpRequest, chainedProxies, clientConnection.getClientDetails()); + if (chainedProxies.isEmpty()) { + // ChainedProxyManager returned no proxies, can't connect + return null; + } } - - private ProxyToServerConnection( - DefaultHttpProxyServer proxyServer, - ClientToProxyConnection clientConnection, - String serverHostAndPort, - ChainedProxy chainedProxy, - Queue availableChainedProxies, - HttpFilters initialFilters, - GlobalTrafficShapingHandler globalTrafficShapingHandler) - throws UnknownHostException { - super(DISCONNECTED, proxyServer, true); - this.clientConnection = clientConnection; - this.serverHostAndPort = serverHostAndPort; - this.chainedProxy = chainedProxy; - this.availableChainedProxies = availableChainedProxies; - this.trafficHandler = globalTrafficShapingHandler; - this.currentFilters = initialFilters; - - // Report connection status to HttpFilters - currentFilters.proxyToServerConnectionQueued(); - - setupConnectionParameters(); + return new ProxyToServerConnection( + proxyServer, + clientConnection, + serverHostAndPort, + chainedProxies.poll(), + chainedProxies, + initialFilters, + globalTrafficShapingHandler); + } + + private ProxyToServerConnection( + DefaultHttpProxyServer proxyServer, + ClientToProxyConnection clientConnection, + String serverHostAndPort, + ChainedProxy chainedProxy, + Queue availableChainedProxies, + HttpFilters initialFilters, + GlobalTrafficShapingHandler globalTrafficShapingHandler) + throws UnknownHostException { + super(DISCONNECTED, proxyServer, true); + this.clientConnection = clientConnection; + this.connectionPool = null; + this.serverHostAndPort = serverHostAndPort; + this.chainedProxy = chainedProxy; + this.availableChainedProxies = availableChainedProxies; + this.trafficHandler = globalTrafficShapingHandler; + this.currentFilters = initialFilters; + + // Report connection status to HttpFilters + currentFilters.proxyToServerConnectionQueued(); + + setupConnectionParameters(); + } + + /** Create a new ProxyToServerConnection that is managed by a shared pool. */ + @Nullable + @CheckReturnValue + static ProxyToServerConnection createForPool( + DefaultHttpProxyServer proxyServer, + ServerConnectionPool connectionPool, + ClientToProxyConnection clientConnection, + String serverHostAndPort, + @Nullable ChainedProxy chainedProxy, + HttpFilters initialFilters, + HttpRequest initialHttpRequest, + GlobalTrafficShapingHandler globalTrafficShapingHandler) + throws UnknownHostException { + Queue chainedProxies = new ConcurrentLinkedQueue<>(); + ChainedProxy resolvedChainedProxy = chainedProxy; + ChainedProxyManager chainedProxyManager = proxyServer.getChainProxyManager(); + if (chainedProxyManager != null) { + chainedProxyManager.lookupChainedProxies( + initialHttpRequest, chainedProxies, clientConnection.getClientDetails()); + if (chainedProxies.isEmpty() && resolvedChainedProxy == null) { + return null; + } + if (resolvedChainedProxy != null) { + chainedProxies.remove(resolvedChainedProxy); + } else { + resolvedChainedProxy = chainedProxies.poll(); + } + } + return new ProxyToServerConnection( + proxyServer, + connectionPool, + clientConnection, + serverHostAndPort, + resolvedChainedProxy, + chainedProxies, + initialFilters, + globalTrafficShapingHandler); + } + + /** Constructor for pooled connections. */ + private ProxyToServerConnection( + DefaultHttpProxyServer proxyServer, + ServerConnectionPool connectionPool, + ClientToProxyConnection clientConnection, + String serverHostAndPort, + ChainedProxy chainedProxy, + Queue availableChainedProxies, + HttpFilters initialFilters, + GlobalTrafficShapingHandler globalTrafficShapingHandler) + throws UnknownHostException { + super(DISCONNECTED, proxyServer, true); + this.clientConnection = clientConnection; + this.connectionPool = connectionPool; + this.serverHostAndPort = serverHostAndPort; + this.chainedProxy = chainedProxy; + this.availableChainedProxies = availableChainedProxies; + this.trafficHandler = globalTrafficShapingHandler; + this.currentFilters = initialFilters; + + // Report connection status to HttpFilters + currentFilters.proxyToServerConnectionQueued(); + + setupConnectionParameters(); + } + + /** Returns true if this connection is currently connected. */ + public boolean isConnected() { + return is(ConnectionState.AWAITING_CHUNK) + || is(ConnectionState.AWAITING_CONNECT_OK) + || is(ConnectionState.AWAITING_INITIAL) + || is(ConnectionState.NEGOTIATING_CONNECT); + } + + /** + * Returns true if this connection is available to handle a new request. A connection is available + * if it's connected and not currently processing another request. + */ + public boolean isAvailableForNewRequest() { + if (!isConnected()) { + return false; } + // Check if we're currently idle (not processing a request) + return currentHttpRequest == null && currentHttpResponse == null; + } + + public boolean isManagedByPool() { + return connectionPool != null; + } + + void setMitmPooled(boolean mitmPooled) { + this.mitmPooled = mitmPooled; + } + + /** + * Gets the client connection to use for this request. If a connection pool is being used, this + * will resolve to the correct client connection based on the pending request. For non-pooled + * connections (CONNECT requests), this returns the direct clientConnection reference. + */ + ClientToProxyConnection getClientConnection() { + // For pooled connections, use the current client connection if set + // For non-pooled connections (CONNECT), use the direct reference + if (connectionPool != null && currentClientConnectionForRequest != null) { + return currentClientConnectionForRequest; + } + return clientConnection; + } - /* ************************************************************************* - * Reading - **************************************************************************/ + /** Closes this server connection. */ + public void close() { + if (channel != null) { + channel.close(); + } + } + + /* ************************************************************************* + * Lifecycle + **************************************************************************/ + + @Override + protected void connected() { + super.connected(); + recordServerConnected(); + } + + /* ************************************************************************* + * Reading + **************************************************************************/ + + @Override + protected void read(Object msg) { + if (isConnecting()) { + // When in the middle of connecting (e.g., during CONNECT tunnel establishment), + // we need to pass the CONNECT response through the HttpFilters before handling + // it in the connection flow. This fixes issue #77: + // https://github.com/LittleProxy/LittleProxy/issues/77 + if (msg instanceof HttpObject) { + HttpObject httpObject = (HttpObject) msg; + currentFilters.serverToProxyResponseReceiving(); - @Override - protected void read(Object msg) { - if (isConnecting()) { - LOG.debug( - "In the middle of connecting, forwarding message to connection flow: {}", - msg); - this.connectionFlow.read(msg); - } else { - super.read(msg); + HttpObject filteredHttpObject = currentFilters.serverToProxyResponse(httpObject); + if (filteredHttpObject == null) { + LOG.debug("Filter returned null, forcing disconnect"); + become(DISCONNECT_REQUESTED); + return; } - } - @Override - protected void readHAProxyMessage(HAProxyMessage msg) { - // NO-OP, - // We never expect server to send a proxy protocol message. + currentFilters.serverToProxyResponseReceived(); + + // Pass the filtered message to the connection flow + LOG.debug( + "In the middle of connecting, forwarding filtered message to connection flow: {}", + filteredHttpObject); + connectionFlow.read(filteredHttpObject); + } else { + // Raw data (e.g., for tunneling) - pass directly to connection flow + LOG.debug( + "In the middle of connecting, forwarding raw message to connection flow: {}", msg); + connectionFlow.read(msg); + } + } else { + // Check if we need to perform TLS detection in MITM mode + checkAndPerformTlsDetection(msg); + super.read(msg); + } + } + + @Override + protected void readHAProxyMessage(HAProxyMessage msg) { + // NO-OP, + // We never expect server to send a proxy protocol message. + } + + @Override + ConnectionState readHTTPInitial(HttpResponse httpResponse) { + LOG.debug("Received raw response: {}", httpResponse); + + if (httpResponse.decoderResult().isFailure()) { + LOG.debug( + "Could not parse response from server. Decoder result: {}", + httpResponse.decoderResult().toString()); + + // create a "substitute" Bad Gateway response from the server, since we couldn't understand + // what the actual + // response from the server was. set the keep-alive on the substitute response to false so the + // proxy closes + // the connection to the server, since we don't know what state the server thinks the + // connection is in. + FullHttpResponse substituteResponse = + ProxyUtils.createFullHttpResponse( + HttpVersion.HTTP_1_1, + HttpResponseStatus.BAD_GATEWAY, + "Unable to parse response from server"); + HttpUtil.setKeepAlive(substituteResponse, false); + httpResponse = substituteResponse; } - @Override - protected ConnectionState readHTTPInitial(HttpResponse httpResponse) { - LOG.debug("Received raw response: {}", httpResponse); - - if (httpResponse.decoderResult().isFailure()) { - LOG.debug("Could not parse response from server. Decoder result: {}", httpResponse.decoderResult().toString()); - - // create a "substitute" Bad Gateway response from the server, since we couldn't understand what the actual - // response from the server was. set the keep-alive on the substitute response to false so the proxy closes - // the connection to the server, since we don't know what state the server thinks the connection is in. - FullHttpResponse substituteResponse = ProxyUtils.createFullHttpResponse(HttpVersion.HTTP_1_1, - HttpResponseStatus.BAD_GATEWAY, - "Unable to parse response from server"); - HttpUtil.setKeepAlive(substituteResponse, false); - httpResponse = substituteResponse; - } + // For pooled connections, look up the pending request to get the correct client connection + // This supports HTTP pipelining where multiple requests are sent and responses come back in + // order + if (connectionPool != null && channel != null) { + PendingRequest pendingRequest = connectionPool.removePendingRequest(channel); + if (pendingRequest != null) { + this.currentClientConnectionForRequest = pendingRequest.getClientConnection(); + this.currentHttpRequest = pendingRequest.getRequest(); + this.currentFilters = pendingRequest.getFilters(); + } + } - currentFilters.serverToProxyResponseReceiving(); + currentFilters.serverToProxyResponseReceiving(); - rememberCurrentResponse(httpResponse); - respondWith(httpResponse); + rememberCurrentResponse(httpResponse); + respondWith(httpResponse); - if (ProxyUtils.isChunked(httpResponse)) { - return AWAITING_CHUNK; - } else { - currentFilters.serverToProxyResponseReceived(); + if (ProxyUtils.isChunked(httpResponse)) { + return AWAITING_CHUNK; + } else { + currentFilters.serverToProxyResponseReceived(); + markResponseComplete(); - return AWAITING_INITIAL; - } + return AWAITING_INITIAL; } - - @Override - protected void readHTTPChunk(HttpContent chunk) { - respondWith(chunk); + } + + @Override + protected void readHTTPChunk(HttpContent chunk) { + respondWith(chunk); + } + + @Override + protected void readRaw(ByteBuf buf) { + getClientConnection().write(buf); + } + + /** + * Responses to HEAD requests aren't supposed to have content, but Netty doesn't know that any + * given response is to a HEAD request, so it needs to be told that there's no content so that it + * doesn't hang waiting for it. + * + *

See the documentation for {@link HttpResponseDecoder} for information about why HEAD + * requests need special handling. + * + *

Thanks to nataliakoval for pointing out that + * with connections being reused as they are, this needs to be sensitive to the current request. + */ + private class HeadAwareHttpResponseDecoder extends HttpResponseDecoder { + + public HeadAwareHttpResponseDecoder( + int maxInitialLineLength, int maxHeaderSize, int maxChunkSize) { + super(maxInitialLineLength, maxHeaderSize, maxChunkSize); } @Override - protected void readRaw(ByteBuf buf) { - clientConnection.write(buf); + protected boolean isContentAlwaysEmpty(HttpMessage httpMessage) { + // The current HTTP Request can be null when this proxy is + // negotiating a CONNECT request with a chained proxy + // while it is running as a MITM. Since the response to a + // CONNECT request does not have any content, we return true. + if (currentHttpRequest == null) { + return true; + } else { + return ProxyUtils.isHEAD(currentHttpRequest) || super.isContentAlwaysEmpty(httpMessage); + } } + } - /** - *

- * Responses to HEAD requests aren't supposed to have content, but Netty - * doesn't know that any given response is to a HEAD request, so it needs to - * be told that there's no content so that it doesn't hang waiting for it. - *

- * - *

- * See the documentation for {@link HttpResponseDecoder} for information - * about why HEAD requests need special handling. - *

- * - *

- * Thanks to nataliakoval for - * pointing out that with connections being reused as they are, this needs - * to be sensitive to the current request. - *

- */ - private class HeadAwareHttpResponseDecoder extends HttpResponseDecoder { - - public HeadAwareHttpResponseDecoder(int maxInitialLineLength, - int maxHeaderSize, int maxChunkSize) { - super(maxInitialLineLength, maxHeaderSize, maxChunkSize); - } + /* ************************************************************************* + * Writing + **************************************************************************/ - @Override - protected boolean isContentAlwaysEmpty(HttpMessage httpMessage) { - // The current HTTP Request can be null when this proxy is - // negotiating a CONNECT request with a chained proxy - // while it is running as a MITM. Since the response to a - // CONNECT request does not have any content, we return true. - if(currentHttpRequest == null) { - return true; - } else { - return ProxyUtils.isHEAD(currentHttpRequest) || super.isContentAlwaysEmpty(httpMessage); - } - } - } + /** Like {@link #write(Object)} and also sets the current filters to the given value. */ + void write(Object msg, HttpFilters filters) { + currentFilters = filters; + write(msg); + } - /* ************************************************************************* - * Writing - **************************************************************************/ - - /** - * Like {@link #write(Object)} and also sets the current filters to the - * given value. - */ - void write(Object msg, HttpFilters filters) { - this.currentFilters = filters; - write(msg); - } + @Override + ChannelFuture write(Object msg) { + LOG.debug("Requested write of {}", msg); - @Override - void write(Object msg) { - LOG.debug("Requested write of {}", msg); + if (msg instanceof ReferenceCounted) { + LOG.debug("Retaining reference counted message"); + ((ReferenceCounted) msg).retain(); + } + boolean needsReuseFlow = + msg instanceof HttpRequest + && ProxyUtils.isCONNECT((HttpRequest) msg) + && !is(DISCONNECTED) + && connectionPool != null; + + if ((is(DISCONNECTED) || needsReuseFlow) && msg instanceof HttpRequest) { + LOG.debug("Currently disconnected, connect and then write the message"); + connectAndWrite((HttpRequest) msg); + return getClientConnection().channel.newSucceededFuture(); + } else { + if (isConnecting()) { + synchronized (connectLock) { + if (isConnecting()) { + LOG.debug( + "Attempted to write while still in the process of connecting, waiting for connection."); + getClientConnection().stopReading(); + try { + connectLock.wait(30000); + } catch (InterruptedException ie) { + LOG.warn("Interrupted while waiting for connect monitor"); + } + } + } + } + + // only write this message if a connection was established and is not in the process of + // disconnecting or + // already disconnected + if (isConnecting() || getCurrentState().isDisconnectingOrDisconnected()) { + LOG.debug( + "Connection failed or timed out while waiting to write message to server. Message will be discarded: {}", + msg); if (msg instanceof ReferenceCounted) { - LOG.debug("Retaining reference counted message"); - ((ReferenceCounted) msg).retain(); + // fix: when connection was disconnecting or disconnected, retain() the refCnt = 2, can + // not release the msg , final leading to OutOfDirectMemoryError. + ((ReferenceCounted) msg).release(); } + return channel.newFailedFuture( + new Exception( + "Connection failed or timed out while waiting to write message to server. Message will be discarded.")); + } - if (is(DISCONNECTED) && msg instanceof HttpRequest) { - LOG.debug("Currently disconnected, connect and then write the message"); - connectAndWrite((HttpRequest) msg); - } else { - if (isConnecting()) { - synchronized (connectLock) { - if (isConnecting()) { - LOG.debug("Attempted to write while still in the process of connecting, waiting for connection."); - clientConnection.stopReading(); - try { - connectLock.wait(30000); - } catch (InterruptedException ie) { - LOG.warn("Interrupted while waiting for connect monitor"); - } - } - } - } - - // only write this message if a connection was established and is not in the process of disconnecting or - // already disconnected - if (isConnecting() || getCurrentState().isDisconnectingOrDisconnected()) { - LOG.debug("Connection failed or timed out while waiting to write message to server. Message will be discarded: {}", msg); - return; - } - - LOG.debug("Using existing connection to: {}", remoteAddress); - doWrite(msg); - } + LOG.debug("Using existing connection to: {}", remoteAddress); + return doWrite(msg); } + } - @Override - protected void writeHttp(HttpObject httpObject) { - if (chainedProxy != null) { - chainedProxy.filterRequest(httpObject); - } - if (httpObject instanceof HttpRequest) { - // Remember that we issued this HttpRequest for later + @Override + protected ChannelFuture writeHttp(HttpObject httpObject) { + if (chainedProxy != null) { + chainedProxy.filterRequest(httpObject); + } + if (httpObject instanceof HttpRequest) { + if (connectionPool == null) { + // Remember that we issued this HttpRequest for later + currentHttpRequest = (HttpRequest) httpObject; + } + + // For pooled connections, register pending request to support HTTP pipelining + if (connectionPool != null && channel != null) { + ClientToProxyConnection clientConn = getClientConnection(); + if (clientConn != null) { + connectionPool.registerPendingRequest( + channel, clientConn, (HttpRequest) httpObject, currentFilters); + if (currentHttpRequest == null) { currentHttpRequest = (HttpRequest) httpObject; + } } - super.writeHttp(httpObject); + } } - - /* ************************************************************************* - * Lifecycle - **************************************************************************/ - - @Override - protected void become(ConnectionState newState) { - // Report connection status to HttpFilters - if (getCurrentState() == DISCONNECTED && newState == CONNECTING) { - currentFilters.proxyToServerConnectionStarted(); - } else if (getCurrentState() == CONNECTING) { - if (newState == HANDSHAKING) { - currentFilters.proxyToServerConnectionSSLHandshakeStarted(); - } else if (newState == AWAITING_INITIAL) { - currentFilters.proxyToServerConnectionSucceeded(ctx); - } else if (newState == DISCONNECTED) { - currentFilters.proxyToServerConnectionFailed(); - } - } else if (getCurrentState() == HANDSHAKING) { - if (newState == AWAITING_INITIAL) { - currentFilters.proxyToServerConnectionSucceeded(ctx); - } else if (newState == DISCONNECTED) { - currentFilters.proxyToServerConnectionFailed(); - } - } else if (getCurrentState() == AWAITING_CHUNK - && newState != AWAITING_CHUNK) { - currentFilters.serverToProxyResponseReceived(); - } - - super.become(newState); + return super.writeHttp(httpObject); + } + + /* ************************************************************************* + * Lifecycle + **************************************************************************/ + + @Override + protected void become(ConnectionState newState) { + // Report connection status to HttpFilters + if (getCurrentState() == DISCONNECTED && newState == CONNECTING) { + currentFilters.proxyToServerConnectionStarted(); + } else if (getCurrentState() == CONNECTING) { + if (newState == HANDSHAKING) { + currentFilters.proxyToServerConnectionSSLHandshakeStarted(); + } else if (newState == AWAITING_INITIAL) { + currentFilters.proxyToServerConnectionSucceeded(ctx); + } else if (newState == DISCONNECTED) { + currentFilters.proxyToServerConnectionFailed(); + } + } else if (getCurrentState() == HANDSHAKING) { + if (newState == AWAITING_INITIAL) { + currentFilters.proxyToServerConnectionSucceeded(ctx); + } else if (newState == DISCONNECTED) { + currentFilters.proxyToServerConnectionFailed(); + } + } else if (getCurrentState() == AWAITING_CHUNK && newState != AWAITING_CHUNK) { + currentFilters.serverToProxyResponseReceived(); + markResponseComplete(); } - @Override - protected void becameSaturated() { - super.becameSaturated(); - this.clientConnection.serverBecameSaturated(this); + super.become(newState); + } + + @Override + protected void becameSaturated() { + super.becameSaturated(); + recordConnectionSaturated(); + getClientConnection().serverBecameSaturated(this); + } + + @Override + protected void becameWritable() { + super.becameWritable(); + recordConnectionWritable(); + getClientConnection().serverBecameWriteable(this); + } + + @Override + protected void timedOut() { + super.timedOut(); + recordConnectionTimedOut(); + getClientConnection().timedOut(this); + } + + @Override + protected void disconnected() { + super.disconnected(); + recordServerDisconnected(); + if (chainedProxy != null) { + // Let the ChainedProxy know that we disconnected + try { + chainedProxy.disconnected(); + } catch (Exception e) { + LOG.error("Unable to record connectionFailed", e); + } } - - @Override - protected void becameWritable() { - super.becameWritable(); - this.clientConnection.serverBecameWriteable(this); + // Remove from pool if this connection was managed by a pool + if (connectionPool != null) { + connectionPool.removeConnection(this); + if (channel != null) { + connectionPool.drainPendingRequests(channel); + } } - - @Override - protected void timedOut() { - super.timedOut(); - clientConnection.timedOut(this); + ClientToProxyConnection clientConn = getClientConnection(); + if (clientConn != null) { + clientConn.serverDisconnected(this); } - - @Override - protected void disconnected() { - super.disconnected(); - if (this.chainedProxy != null) { - // Let the ChainedProxy know that we disconnected - try { - this.chainedProxy.disconnected(); - } catch (Exception e) { - LOG.error("Unable to record connectionFailed", e); - } - } - clientConnection.serverDisconnected(this); + } + + @Override + protected void exceptionCaught(Throwable cause) { + try { + if (!is(DISCONNECTED)) { + recordConnectionExceptionCaught(cause); + } + if (cause instanceof ProxyConnectException) { + LOG.info( + "A ProxyConnectException occurred on ProxyToServerConnection: " + cause.getMessage()); + connectionFlow.fail(cause); + } else if (cause instanceof IOException) { + // IOExceptions are expected errors, for example when a server drops the connection. rather + // than flood + // the logs with stack traces for these expected exceptions, log the message at the INFO + // level and the + // stack trace at the DEBUG level. + LOG.info("An IOException occurred on ProxyToServerConnection: " + cause.getMessage()); + LOG.debug("An IOException occurred on ProxyToServerConnection", cause); + } else if (cause instanceof RejectedExecutionException) { + LOG.info( + "An executor rejected a read or write operation on the ProxyToServerConnection (this is normal if the proxy is shutting down). Message: " + + cause.getMessage()); + LOG.debug("A RejectedExecutionException occurred on ProxyToServerConnection", cause); + } else { + LOG.error("Caught an exception on ProxyToServerConnection", cause); + } + } finally { + if (!is(DISCONNECTED)) { + LOG.info("Disconnecting open connection to server"); + disconnect(); + getClientConnection().serverConnectionFailed(this, getCurrentState(), cause); + } + } + // This can happen if we couldn't make the initial connection due + // to something like an unresolved address, for example, or a timeout. + // There will not be any requests written on an unopened + // connection, so there should not be any further action to take here. + } + + /* ************************************************************************* + * State Management + **************************************************************************/ + public TransportProtocol getTransportProtocol() { + return transportProtocol; + } + + public ChainedProxyType getChainedProxyType() { + return chainedProxyType; + } + + public InetSocketAddress getRemoteAddress() { + return remoteAddress; + } + + public String getServerHostAndPort() { + return serverHostAndPort; + } + + public boolean hasUpstreamChainedProxy() { + return getChainedProxyAddress() != null; + } + + @Nullable + public InetSocketAddress getChainedProxyAddress() { + return chainedProxy == null ? null : chainedProxy.getChainedProxyAddress(); + } + + @Nullable + public ChainedProxy getChainedProxy() { + return chainedProxy; + } + + @Nullable + public HttpRequest getInitialRequest() { + return initialRequest; + } + + @Override + protected HttpFilters getHttpFiltersFromProxyServer(HttpRequest httpRequest) { + return currentFilters; + } + + /* ************************************************************************* + * Private Implementation + **************************************************************************/ + + /** + * Keeps track of the current HttpResponse so that we can associate its headers with future + * related chunks for this same transfer. + */ + private void rememberCurrentResponse(HttpResponse response) { + LOG.debug("Remembering the current response."); + // We need to make a copy here because the response will be + // modified in various ways before we need to do things like + // analyze response headers for whether to close the + // connection (which may not happen for a while for large, chunked + // responses, for example). + currentHttpResponse = ProxyUtils.copyMutableResponseFields(response); + } + + /** Respond to the client with the given {@link HttpObject}. */ + private void respondWith(HttpObject httpObject) { + // Use the current client connection if set (for pooled connections), + // otherwise use the direct reference + ClientToProxyConnection targetClientConnection = getClientConnection(); + + if (targetClientConnection == null) { + LOG.warn("No client connection available to respond to"); + return; } - @Override - protected void exceptionCaught(Throwable cause) { - try { - if (cause instanceof ProxyConnectException) { - LOG.info("A ProxyConnectException occurred on ProxyToServerConnection: " + cause.getMessage()); - connectionFlow.fail(cause); - } else if (cause instanceof IOException) { - // IOExceptions are expected errors, for example when a server drops the connection. rather than flood - // the logs with stack traces for these expected exceptions, log the message at the INFO level and the - // stack trace at the DEBUG level. - LOG.info("An IOException occurred on ProxyToServerConnection: " + cause.getMessage()); - LOG.debug("An IOException occurred on ProxyToServerConnection", cause); - } else if (cause instanceof RejectedExecutionException) { - LOG.info("An executor rejected a read or write operation on the ProxyToServerConnection (this is normal if the proxy is shutting down). Message: " + cause.getMessage()); - LOG.debug("A RejectedExecutionException occurred on ProxyToServerConnection", cause); - } else { - LOG.error("Caught an exception on ProxyToServerConnection", cause); - } - } finally { - if (!is(DISCONNECTED)) { - LOG.info("Disconnecting open connection to server"); - disconnect(); - } + targetClientConnection.respond( + this, currentFilters, currentHttpRequest, currentHttpResponse, httpObject); + } + + /** + * Sets the client connection that is making the current request. Used for pooled connections to + * track which client to send the response to. + */ + void setCurrentClientConnectionForRequest(@Nullable ClientToProxyConnection clientConnection) { + this.currentClientConnectionForRequest = clientConnection; + } + + void setRemoteAddress(InetSocketAddress remoteAddress) { + this.remoteAddress = remoteAddress; + } + + void releaseToPool() { + this.currentClientConnectionForRequest = null; + this.currentHttpResponse = null; + if (connectionPool != null) { + this.currentHttpRequest = null; + connectionPool.releaseConnection(this); + } + } + + private void markResponseComplete() { + this.currentClientConnectionForRequest = null; + this.currentHttpResponse = null; + if (connectionPool != null) { + if (channel != null) { + PendingRequest nextPending = connectionPool.removePendingRequest(channel); + if (nextPending != null) { + this.currentClientConnectionForRequest = nextPending.getClientConnection(); + this.currentHttpRequest = nextPending.getRequest(); + this.currentFilters = nextPending.getFilters(); + return; } - // This can happen if we couldn't make the initial connection due - // to something like an unresolved address, for example, or a timeout. - // There will not have been be any requests written on an unopened - // connection, so there should not be any further action to take here. + } + this.currentHttpRequest = null; + // MITM pooled connections stay attached to the client session and are not released + // to pool after individual HTTP responses. They are released when the client disconnects. + if (!mitmPooled) { + connectionPool.releaseConnection(this); + } + } else { + this.currentHttpRequest = null; } - - /* ************************************************************************* - * State Management - **************************************************************************/ - public TransportProtocol getTransportProtocol() { - return transportProtocol; + } + + /** + * Configures the connection to the upstream server and begins the {@link ConnectionFlow}. + * + * @param initialRequest the current HTTP request being handled + */ + private void connectAndWrite(HttpRequest initialRequest) { + LOG.debug("Starting new connection to: {}", remoteAddress); + + // Remember our initial request so that we can write it after connecting + this.initialRequest = initialRequest; + initializeConnectionFlow(); + connectionFlow.start(); + } + + /** + * This method initializes our {@link ConnectionFlow} based on however this connection has been + * configured. If the {@link #disableSni} value is true, this method will not pass peer + * information to the MitmManager when handling CONNECTs. + */ + private void initializeConnectionFlow() { + boolean isReused = channel != null && channel.isActive(); + + if (isReused) { + LOG.debug("Reusing existing connection for CONNECT, skipping TCP/TLS setup"); + connectionFlow = new ConnectionFlow(getClientConnection(), this, connectLock); + } else { + connectionFlow = + new ConnectionFlow(getClientConnection(), this, connectLock).then(ConnectChannel); } - public ChainedProxyType getChainedProxyType() { - return chainedProxyType; + boolean sendProxyProtocol = proxyServer.isSendProxyProtocol(); + boolean chained = hasUpstreamChainedProxy(); + boolean chainedSocks = + chained + && (chainedProxyType == ChainedProxyType.SOCKS4 + || chainedProxyType == ChainedProxyType.SOCKS5); + boolean chainedHttp = chained && chainedProxyType == ChainedProxyType.HTTP; + boolean isConnect = ProxyUtils.isCONNECT(initialRequest); + + // Where to write the PROXY header so it reaches the final server: + // - Direct: first, right after connecting (peer is the final server). + // - HTTP CONNECT chain: tunnelled after the CONNECT handshake (see the CONNECT block below); + // the intermediate sees a plain CONNECT. + // - SOCKS chain: skipped (warning). + // - Non-CONNECT HTTP chain: skipped (warning) - no tunnel to the final server. + boolean sendProxyHeaderFirst = sendProxyProtocol && !chained; + boolean tunnelProxyHeaderThroughConnect = sendProxyProtocol && chainedHttp && isConnect; + + if (sendProxyProtocol && chainedSocks) { + LOG.warn( + "PROXY protocol is not compatible with SOCKS upstream proxies ({}). Skipping PROXY header.", + chainedProxyType); + } else if (sendProxyProtocol && chainedHttp && !isConnect) { + LOG.warn( + "PROXY protocol cannot be forwarded through a non-CONNECT HTTP chained proxy. " + + "Skipping PROXY header."); } - public InetSocketAddress getRemoteAddress() { - return remoteAddress; + if (sendProxyHeaderFirst && !isReused) { + connectionFlow.then(SendProxyProtocolHeader); } - public String getServerHostAndPort() { - return serverHostAndPort; + if (chained && !isReused) { + if (chainedProxy.requiresEncryption()) { + connectionFlow.then(serverConnection.EncryptChannel(newChainedProxySslEngine())); + } + switch (chainedProxyType) { + case SOCKS4: + connectionFlow.then(SOCKS4CONNECTWithChainedProxy); + break; + case SOCKS5: + connectionFlow.then(SOCKS5InitialRequest); + break; + default: + break; + } } - public boolean hasUpstreamChainedProxy() { - return getChainedProxyAddress() != null; - } + if (isConnect) { + // If we're chaining to an upstream HTTP proxy, forward the CONNECT request. + // Do not chain the CONNECT request for SOCKS proxies. + if (chainedHttp && !isReused) { + connectionFlow.then(serverConnection.HTTPCONNECTWithChainedProxy); + } + + // Write the PROXY header into the established tunnel (first bytes to the final server), + // before + // StartTunneling/EncryptChannel remove the HAProxyMessageEncoder. + if (tunnelProxyHeaderThroughConnect && !isReused) { + connectionFlow.then(SendProxyProtocolHeader); + } + + MitmManager mitmManager = proxyServer.getMitmManager(); + boolean isMitmEnabled = currentFilters.proxyToServerAllowMitm() && mitmManager != null; + + if (isMitmEnabled && connectionPool != null) { + if (proxyServer.isPoolPerRequestInMitm()) { + releaseToPoolOnConnectComplete = true; + } else { + setMitmPooled(true); + } + } + + if (isMitmEnabled) { + // When MITM is enabled and when chained proxy is set up, remoteAddress + // will be the chained proxy's address. So we use serverHostAndPoint + // which is the end server's address. + HostAndPort parsedHostAndPort = HostAndPort.fromString(serverHostAndPort); + + // Check if we should skip SSL (e.g., after a retry for non-SSL server) + if (!disableSslForNonTls && !isReused) { + // SNI may be disabled for this request due to a previous failed attempt to connect to the + // server + // with SNI enabled. + if (disableSni) { + connectionFlow.then( + serverConnection.EncryptChannel(proxyServer.getMitmManager().serverSslEngine())); + } else { + connectionFlow.then( + serverConnection.EncryptChannel( + proxyServer + .getMitmManager() + .serverSslEngine( + parsedHostAndPort.getHost(), parsedHostAndPort.getPort()))); + } + } - public InetSocketAddress getChainedProxyAddress() { - return chainedProxy == null ? null : chainedProxy - .getChainedProxyAddress(); + if (!disableSslForNonTls) { + connectionFlow + .then(getClientConnection().RespondCONNECTSuccessful) + .then(serverConnection.MitmEncryptClientChannel); + } else { + // For non-SSL servers, just respond CONNECT successful. + // ClientToProxyConnection will call encryptForMitm() if client sends TLS. + connectionFlow.then(getClientConnection().RespondCONNECTSuccessful); + } + } else { + if (isReused) { + connectionFlow.then(getClientConnection().RespondCONNECTSuccessful); + } else { + connectionFlow + .then(serverConnection.StartTunneling) + .then(getClientConnection().RespondCONNECTSuccessful) + .then(getClientConnection().StartTunneling); + } + } } + } + + /** + * A connection flow step that waits for the server's response to the CONNECT request. This is + * used in MITM mode when we don't want to add SSL to the server connection upfront - instead, we + * wait for the server's response and then inspect the first bytes from the client to detect TLS. + */ + private final ConnectionFlowStep MitmWaitForServerConnectResponse = + new ConnectionFlowStep<>(this, ConnectionState.AWAITING_CONNECT_OK) { + @Override + boolean shouldSuppressInitialRequest() { + return true; + } - public ChainedProxy getChainedProxy() { - return chainedProxy; - } + @Override + protected Future execute() { + // This step just marks that we're waiting for the server's response + return channel.newSucceededFuture(); + } - public HttpRequest getInitialRequest() { - return initialRequest; - } + @Override + public void read(ConnectionFlow flow, Object msg) { + // Server responded to CONNECT - this step is complete + LOG.debug("Server responded to CONNECT in MITM mode: {}", msg); + flow.advance(); + } + }; + + /** + * A connection flow step that waits for the first bytes from the client to determine if SSL/TLS + * is needed. This inspects the first byte to detect TLS handshake. + */ + private final ConnectionFlowStep MitmDetectTlsAndEncrypt = + new ConnectionFlowStep<>(this, ConnectionState.NEGOTIATING_CONNECT) { + @Override + boolean shouldSuppressInitialRequest() { + return true; + } - @Override - protected HttpFilters getHttpFiltersFromProxyServer(HttpRequest httpRequest) { - return currentFilters; - } + @Override + protected Future execute() { + // Don't complete yet - wait for first data from client in read() + return channel.newSucceededFuture(); + } - /* ************************************************************************* - * Private Implementation - **************************************************************************/ - - /** - * Keeps track of the current HttpResponse so that we can associate its - * headers with future related chunks for this same transfer. - */ - private void rememberCurrentResponse(HttpResponse response) { - LOG.debug("Remembering the current response."); - // We need to make a copy here because the response will be - // modified in various ways before we need to do things like - // analyze response headers for whether or not to close the - // connection (which may not happen for a while for large, chunked - // responses, for example). - currentHttpResponse = ProxyUtils.copyMutableResponseFields(response); - } + @Override + public void read(ConnectionFlow flow, Object msg) { + // First data from client - inspect for TLS + if (msg instanceof ByteBuf) { + ByteBuf buf = (ByteBuf) msg; + if (buf.readableBytes() > 0) { + byte firstByte = buf.getByte(buf.readerIndex()); + boolean isTlsHandshake = (firstByte & 0xFF) == 0x16; + + LOG.debug( + "Inspecting first byte from client: 0x{} - TLS handshake: {}", + Integer.toHexString(firstByte & 0xFF), + isTlsHandshake); + + if (isTlsHandshake) { + // This is a TLS connection - encrypt both server and client connections + encryptForMitm(); + } + + tlsInspectionDone = true; + flow.advance(); + return; + } + } + // If not a ByteBuf or no readable bytes, just advance + tlsInspectionDone = true; + flow.advance(); + } + }; + + /** Flag to track whether we've inspected the first bytes for TLS detection in MITM mode. */ + private volatile boolean tlsInspectionDone = false; + + /** + * Checks if we need to perform TLS detection for MITM mode and does so if needed. This inspects + * the first bytes from the client to determine if SSL/TLS is needed. + */ + private void checkAndPerformTlsDetection(Object msg) { + MitmManager mitmManager = proxyServer.getMitmManager(); + boolean isMitmEnabled = currentFilters.proxyToServerAllowMitm() && mitmManager != null; + + // Only do TLS detection once, and only when in MITM mode + if (isMitmEnabled && !tlsInspectionDone && msg instanceof ByteBuf) { + ByteBuf buf = (ByteBuf) msg; + if (buf.readableBytes() > 0) { + // Peek at the first byte to determine if this is a TLS handshake + // TLS handshake always starts with 0x16 (decimal 22) + byte firstByte = buf.getByte(buf.readerIndex()); + boolean isTlsHandshake = (firstByte & 0xFF) == 0x16; + + LOG.debug( + "Inspecting first byte from client: 0x{} - TLS handshake: {}", + Integer.toHexString(firstByte & 0xFF), + isTlsHandshake); + + if (isTlsHandshake) { + // This is a TLS connection - encrypt both server and client connections + encryptForMitm(); + } - /** - * Respond to the client with the given {@link HttpObject}. - */ - private void respondWith(HttpObject httpObject) { - clientConnection.respond(this, currentFilters, currentHttpRequest, - currentHttpResponse, httpObject); + tlsInspectionDone = true; + } + } + } + + /** Encrypts both server and client connections for MITM. */ + private void encryptForMitm() { + HostAndPort parsedHostAndPort = HostAndPort.fromString(serverHostAndPort); + int port = parsedHostAndPort.getPort(); + + // Encrypt the server connection (for MITM) + Future serverEncryptFuture; + if (disableSni) { + serverEncryptFuture = encrypt(proxyServer.getMitmManager().serverSslEngine(), true); + } else { + serverEncryptFuture = + encrypt( + proxyServer.getMitmManager().serverSslEngine(parsedHostAndPort.getHost(), port), + true); } - /** - * Configures the connection to the upstream server and begins the {@link ConnectionFlow}. - * - * @param initialRequest the current HTTP request being handled - */ - private void connectAndWrite(HttpRequest initialRequest) { - LOG.debug("Starting new connection to: {}", remoteAddress); - - // Remember our initial request so that we can write it after connecting - this.initialRequest = initialRequest; - initializeConnectionFlow(); - connectionFlow.start(); + // Encrypt the client connection for MITM, but wait for server encryption first + serverEncryptFuture.addListener( + future -> { + if (future.isSuccess()) { + ClientToProxyConnection targetClient = getClientConnection(); + targetClient + .encrypt( + proxyServer + .getMitmManager() + .clientSslEngineFor(initialRequest, sslEngine.getSession()), + false) + .addListener( + clientFuture -> { + if (clientFuture.isSuccess()) { + targetClient.setMitming(true); + } else { + LOG.warn("Failed to encrypt client connection for MITM"); + } + }); + } else { + LOG.warn("Failed to encrypt server connection for MITM"); + } + }); + } + + SSLEngine newChainedProxySslEngine() { + if (remoteAddress != null) { + SSLEngine peerAwareSslEngine = + chainedProxy.newSslEngine(remoteAddress.getHostString(), remoteAddress.getPort()); + if (peerAwareSslEngine != null) { + return peerAwareSslEngine; + } } - /** - * This method initializes our {@link ConnectionFlow} based on however this connection has been configured. If - * the {@link #disableSni} value is true, this method will not pass peer information to the MitmManager when - * handling CONNECTs. - */ - private void initializeConnectionFlow() { - this.connectionFlow = new ConnectionFlow(clientConnection, this, - connectLock) - .then(ConnectChannel); - - if (hasUpstreamChainedProxy()) { - if (chainedProxy.requiresEncryption()) { - connectionFlow.then(serverConnection.EncryptChannel(chainedProxy.newSslEngine())); - } - switch (chainedProxyType) { - case SOCKS4: - connectionFlow.then(SOCKS4CONNECTWithChainedProxy); - break; - case SOCKS5: - connectionFlow.then(SOCKS5InitialRequest); - break; - default: - break; - } - } + return chainedProxy.newSslEngine(); + } - if (ProxyUtils.isCONNECT(initialRequest)) { - // If we're chaining to an upstream HTTP proxy, forward the CONNECT request. - // Do not chain the CONNECT request for SOCKS proxies. - if (hasUpstreamChainedProxy() && (chainedProxyType == ChainedProxyType.HTTP)) { - connectionFlow.then(serverConnection.HTTPCONNECTWithChainedProxy); + final ConnectionFlowStep SendProxyProtocolHeader = + new ConnectionFlowStep<>(this, CONNECTING) { + @Override + protected Future execute() { + HAProxyMessage haProxyMessage = clientConnection.getHaProxyMessage(); + ProxyProtocolMessage proxyProtocolMessage; + if (haProxyMessage != null) { + proxyProtocolMessage = new ProxyProtocolMessage(haProxyMessage); + } else { + InetSocketAddress clientAddr = clientConnection.getClientAddress(); + if (clientAddr == null + || clientAddr.getAddress() == null + || remoteAddress == null + || remoteAddress.getAddress() == null) { + LOG.warn("Cannot send PROXY protocol header: addresses not available"); + return channel.newSucceededFuture(); } - - MitmManager mitmManager = proxyServer.getMitmManager(); - boolean isMitmEnabled = mitmManager != null; - - if (isMitmEnabled) { - // When MITM is enabled and when chained proxy is set up, remoteAddress - // will be the chained proxy's address. So we use serverHostAndPort - // which is the end server's address. - HostAndPort parsedHostAndPort = HostAndPort.fromString(serverHostAndPort); - - // SNI may be disabled for this request due to a previous failed attempt to connect to the server - // with SNI enabled. - if (disableSni) { - connectionFlow.then(serverConnection.EncryptChannel(proxyServer.getMitmManager() - .serverSslEngine())); - } else { - connectionFlow.then(serverConnection.EncryptChannel(proxyServer.getMitmManager() - .serverSslEngine(parsedHostAndPort.getHost(), parsedHostAndPort.getPort()))); - } - - connectionFlow - .then(clientConnection.RespondCONNECTSuccessful) - .then(serverConnection.MitmEncryptClientChannel); + java.net.InetAddress clientInet = clientAddr.getAddress(); + java.net.InetAddress serverInet = remoteAddress.getAddress(); + boolean clientIsV6 = clientInet instanceof java.net.Inet6Address; + boolean serverIsV6 = serverInet instanceof java.net.Inet6Address; + HAProxyProxiedProtocol protocol; + if (clientIsV6 && serverIsV6) { + protocol = HAProxyProxiedProtocol.TCP6; + } else if (!clientIsV6 && !serverIsV6) { + protocol = HAProxyProxiedProtocol.TCP4; } else { - connectionFlow.then(serverConnection.StartTunneling) - .then(clientConnection.RespondCONNECTSuccessful) - .then(clientConnection.StartTunneling); + LOG.warn( + "Cannot send PROXY protocol header: mixed address families (client={}, server={})", + clientInet.getClass().getSimpleName(), + serverInet.getClass().getSimpleName()); + return channel.newSucceededFuture(); } + proxyProtocolMessage = + new ProxyProtocolMessage( + HAProxyProtocolVersion.V1, + HAProxyCommand.PROXY, + protocol, + clientAddr.getAddress().getHostAddress(), + remoteAddress.getAddress().getHostAddress(), + clientAddr.getPort(), + remoteAddress.getPort()); + } + return writeToChannel(proxyProtocolMessage); } + }; + + private void addFirstOrReplaceHandler(String name, ChannelHandler handler) { + if (channel.pipeline().context(name) != null) { + channel.pipeline().replace(name, name, handler); + } else { + channel.pipeline().addFirst(name, handler); } - - private void addFirstOrReplaceHandler(String name, ChannelHandler handler) { - if (channel.pipeline().context(name) != null) { - channel.pipeline().replace(name, name, handler); - } - else { - channel.pipeline().addFirst(name, handler); - } - } - - private void removeHandlerIfPresent(String name) { - if (channel.pipeline().context(name) != null) { - channel.pipeline().remove(name); - } - } + } - /** - * Opens the socket connection. - */ - private ConnectionFlowStep ConnectChannel = new ConnectionFlowStep(this, - CONNECTING) { + private void removeHandlerIfPresent(String name) { + removeHandlerIfPresent(channel.pipeline(), name); + } + + /** Opens the socket connection. */ + private final ConnectionFlowStep ConnectChannel = + new ConnectionFlowStep<>(this, CONNECTING) { @Override boolean shouldExecuteOnEventLoop() { - return false; + return false; } @Override protected Future execute() { - Bootstrap cb = new Bootstrap() - .group(proxyServer.getProxyToServerWorkerFor(transportProtocol)) - .resolver(remoteAddressResolver); + Bootstrap cb = + new Bootstrap() + .group(proxyServer.getProxyToServerWorkerFor(transportProtocol)) + .resolver(remoteAddressResolver); - switch (transportProtocol) { + switch (transportProtocol) { case TCP: - LOG.debug("Connecting to server with TCP"); - cb.channelFactory(NioSocketChannel::new); - break; - case UDT: - LOG.debug("Connecting to server with UDT"); - cb.channelFactory(NioUdtProvider.BYTE_CONNECTOR) - .option(ChannelOption.SO_REUSEADDR, true); - break; + LOG.debug("Connecting to server with TCP"); + cb.channelFactory(NioSocketChannel::new); + break; default: - throw new UnknownTransportProtocolException(transportProtocol); - } + throw new UnknownTransportProtocolException(transportProtocol); + } - cb.handler(new ChannelInitializer() { + cb.handler( + new ChannelInitializer<>() { protected void initChannel(Channel ch) { - initChannelPipeline(ch.pipeline(), initialRequest); + initChannelPipeline(ch.pipeline()); } - }); - cb.option(ChannelOption.CONNECT_TIMEOUT_MILLIS, - proxyServer.getConnectTimeout()); - - if (localAddress != null) { - return cb.connect(remoteAddress, localAddress); - } else { - return cb.connect(remoteAddress); - } + }); + cb.option(ChannelOption.CONNECT_TIMEOUT_MILLIS, proxyServer.getConnectTimeout()); + + if (localAddress != null) { + return cb.connect(remoteAddress, localAddress); + } else { + return cb.connect(remoteAddress); + } } - }; + }; - /** - * Writes the HTTP CONNECT to the server and waits for a 200 response. - */ - private ConnectionFlowStep HTTPCONNECTWithChainedProxy = new ConnectionFlowStep( - this, AWAITING_CONNECT_OK) { + /** Writes the HTTP CONNECT to the server and waits for a 200 response. */ + private final ConnectionFlowStep HTTPCONNECTWithChainedProxy = + new ConnectionFlowStep<>(this, AWAITING_CONNECT_OK) { protected Future execute() { - LOG.debug("Handling CONNECT request through Chained Proxy"); - chainedProxy.filterRequest(initialRequest); - MitmManager mitmManager = proxyServer.getMitmManager(); - boolean isMitmEnabled = mitmManager != null; - /* - * We ignore the LastHttpContent which we read from the client - * connection when we are negotiating connect (see readHttp() - * in ProxyConnection). This cannot be ignored while we are - * doing MITM + Chained Proxy because the HttpRequestEncoder - * of the ProxyToServerConnection will be in an invalid state - * when the next request is written. Writing the EmptyLastContent - * resets its state. - */ - if(isMitmEnabled){ - ChannelFuture future = writeToChannel(initialRequest); - future.addListener((ChannelFutureListener) arg0 -> { - if(arg0.isSuccess()){ + LOG.debug("Handling CONNECT request through Chained Proxy"); + chainedProxy.filterRequest(initialRequest); + MitmManager mitmManager = proxyServer.getMitmManager(); + boolean isMitmEnabled = currentFilters.proxyToServerAllowMitm() && mitmManager != null; + /* + * We ignore the LastHttpContent which we read from the client + * connection when we are negotiating connect (see readHttp() + * in ProxyConnection). This cannot be ignored while we are + * doing MITM + Chained Proxy because the HttpRequestEncoder + * of the ProxyToServerConnection will be in an invalid state + * when the next request is written. Writing the EmptyLastContent + * resets its state. + */ + if (isMitmEnabled) { + ChannelFuture future = writeToChannel(initialRequest); + future.addListener( + (ChannelFutureListener) + arg0 -> { + if (arg0.isSuccess()) { writeToChannel(LastHttpContent.EMPTY_LAST_CONTENT); - } - }); - return future; - } else { - return writeToChannel(initialRequest); - } + } + }); + return future; + } else { + return writeToChannel(initialRequest); + } } void onSuccess(ConnectionFlow flow) { - // Do nothing, since we want to wait for the CONNECT response to - // come back + // Do nothing, since we want to wait for the CONNECT response to + // come back } void read(ConnectionFlow flow, Object msg) { - // Here we're handling the response from a chained proxy to our - // earlier CONNECT request - boolean connectOk = false; - if (msg instanceof HttpResponse) { - HttpResponse httpResponse = (HttpResponse) msg; - int statusCode = httpResponse.status().code(); - if (statusCode >= 200 && statusCode <= 299) { - connectOk = true; - } - } - if (connectOk) { - flow.advance(); - } else { - flow.fail(); + // Here we're handling the response from a chained proxy to our + // earlier CONNECT request + boolean connectOk = false; + if (msg instanceof HttpResponse) { + HttpResponse httpResponse = (HttpResponse) msg; + int statusCode = httpResponse.status().code(); + if (statusCode >= 200 && statusCode <= 299) { + connectOk = true; } + } + if (connectOk) { + flow.advance(); + } else { + flow.fail(); + } } - }; - - /** - * Establishes a SOCKS4 connection. - */ - private ConnectionFlowStep SOCKS4CONNECTWithChainedProxy = new ConnectionFlowStep( - this, AWAITING_CONNECT_OK) { + }; + + /** Establishes a SOCKS4 connection. */ + private final ConnectionFlowStep SOCKS4CONNECTWithChainedProxy = + new ConnectionFlowStep<>(this, AWAITING_CONNECT_OK) { @Override protected Future execute() { - InetSocketAddress destinationAddress; - try { - destinationAddress = addressFor(serverHostAndPort, proxyServer); - } catch (UnknownHostException e) { - return channel.newFailedFuture(e); - } - - DefaultSocks4CommandRequest connectRequest = new DefaultSocks4CommandRequest( - Socks4CommandType.CONNECT, destinationAddress.getHostString(), destinationAddress.getPort()); - - addFirstOrReplaceHandler(SOCKS_ENCODER_NAME, Socks4ClientEncoder.INSTANCE); - addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks4ClientDecoder()); - return writeToChannel(connectRequest); + InetSocketAddress destinationAddress; + try { + destinationAddress = addressFor(serverHostAndPort, proxyServer); + } catch (UnknownHostException e) { + return channel.newFailedFuture(e); + } + + DefaultSocks4CommandRequest connectRequest = + new DefaultSocks4CommandRequest( + Socks4CommandType.CONNECT, + destinationAddress.getHostString(), + destinationAddress.getPort()); + + addFirstOrReplaceHandler(SOCKS_ENCODER_NAME, Socks4ClientEncoder.INSTANCE); + addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks4ClientDecoder()); + return writeToChannel(connectRequest); } @Override void read(ConnectionFlow flow, Object msg) { - removeHandlerIfPresent(SOCKS_ENCODER_NAME); - removeHandlerIfPresent(SOCKS_DECODER_NAME); - if (msg instanceof Socks4CommandResponse) { - if (((Socks4CommandResponse) msg).status() == Socks4CommandStatus.SUCCESS) { - flow.advance(); - return; - } + removeHandlerIfPresent(SOCKS_ENCODER_NAME); + removeHandlerIfPresent(SOCKS_DECODER_NAME); + if (msg instanceof Socks4CommandResponse) { + if (((Socks4CommandResponse) msg).status() == Socks4CommandStatus.SUCCESS) { + flow.advance(); + return; } - flow.fail(); + } + flow.fail(); } - + @Override void onSuccess(ConnectionFlow flow) { - // Do not advance the flow until the SOCKS response has been parsed + // Do not advance the flow until the SOCKS response has been parsed } - }; - - /** - * Initiates a SOCKS5 connection. - */ - private ConnectionFlowStep SOCKS5InitialRequest = new ConnectionFlowStep( - this, AWAITING_CONNECT_OK) { + }; + + /** Initiates a SOCKS5 connection. */ + private final ConnectionFlowStep SOCKS5InitialRequest = + new ConnectionFlowStep<>(this, AWAITING_CONNECT_OK) { @Override protected Future execute() { - List authMethods = new ArrayList<>(2); - authMethods.add(Socks5AuthMethod.NO_AUTH); - if ((username != null) || (password != null)) { - authMethods.add(Socks5AuthMethod.PASSWORD); - } - DefaultSocks5InitialRequest initialRequest = new DefaultSocks5InitialRequest(authMethods); - - addFirstOrReplaceHandler(SOCKS_ENCODER_NAME, Socks5ClientEncoder.DEFAULT); - addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks5InitialResponseDecoder()); - return writeToChannel(initialRequest); + List authMethods = new ArrayList<>(2); + authMethods.add(Socks5AuthMethod.NO_AUTH); + if ((username != null) || (password != null)) { + authMethods.add(Socks5AuthMethod.PASSWORD); + } + DefaultSocks5InitialRequest initialRequest = new DefaultSocks5InitialRequest(authMethods); + + addFirstOrReplaceHandler(SOCKS_ENCODER_NAME, Socks5ClientEncoder.DEFAULT); + addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks5InitialResponseDecoder()); + return writeToChannel(initialRequest); } @Override void read(ConnectionFlow flow, Object msg) { - if (msg instanceof Socks5InitialResponse) { - Socks5AuthMethod selectedAuthMethod = ((Socks5InitialResponse) msg).authMethod(); - - final boolean authSuccess; - if (selectedAuthMethod == Socks5AuthMethod.NO_AUTH) { - // Immediately proceed to SOCKS CONNECT - flow.first(SOCKS5CONNECTRequestWithChainedProxy); - authSuccess = true; - } - else if (selectedAuthMethod == Socks5AuthMethod.PASSWORD) { - // Insert a password negotiation step: - flow.first(SOCKS5SendPasswordCredentials); - authSuccess = true; - } - else { - // Server returned Socks5AuthMethod.UNACCEPTED or a method we do not support - authSuccess = false; - } + if (msg instanceof Socks5InitialResponse) { + Socks5AuthMethod selectedAuthMethod = ((Socks5InitialResponse) msg).authMethod(); + + final boolean authSuccess; + if (selectedAuthMethod == Socks5AuthMethod.NO_AUTH) { + // Immediately proceed to SOCKS CONNECT + flow.first(SOCKS5CONNECTRequestWithChainedProxy); + authSuccess = true; + } else if (selectedAuthMethod == Socks5AuthMethod.PASSWORD) { + // Insert a password negotiation step: + flow.first(SOCKS5SendPasswordCredentials); + authSuccess = true; + } else { + // Server returned Socks5AuthMethod.UNACCEPTED or a method we do not support + authSuccess = false; + } - if (authSuccess) { - flow.advance(); - return; - } + if (authSuccess) { + flow.advance(); + return; } - flow.fail(); + } + flow.fail(); } @Override void onSuccess(ConnectionFlow flow) { - // Do not advance the flow until the SOCKS response has been parsed + // Do not advance the flow until the SOCKS response has been parsed } - }; - - /** - * Sends SOCKS5 password credentials after {@link #SOCKS5InitialRequest} has completed. - */ - private ConnectionFlowStep SOCKS5SendPasswordCredentials = new ConnectionFlowStep( - this, AWAITING_CONNECT_OK) { + }; + + /** Sends SOCKS5 password credentials after {@link #SOCKS5InitialRequest} has completed. */ + private final ConnectionFlowStep SOCKS5SendPasswordCredentials = + new ConnectionFlowStep<>(this, AWAITING_CONNECT_OK) { @Override protected Future execute() { - DefaultSocks5PasswordAuthRequest authRequest = new DefaultSocks5PasswordAuthRequest( - username != null ? username : "", password != null ? password : ""); + DefaultSocks5PasswordAuthRequest authRequest = + new DefaultSocks5PasswordAuthRequest( + username != null ? username : "", password != null ? password : ""); - addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks5PasswordAuthResponseDecoder()); - return writeToChannel(authRequest); + addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks5PasswordAuthResponseDecoder()); + return writeToChannel(authRequest); } @Override void read(ConnectionFlow flow, Object msg) { - if (msg instanceof Socks5PasswordAuthResponse) { - if (((Socks5PasswordAuthResponse) msg).status() == Socks5PasswordAuthStatus.SUCCESS) { - flow.first(SOCKS5CONNECTRequestWithChainedProxy); - flow.advance(); - return; - } + if (msg instanceof Socks5PasswordAuthResponse) { + if (((Socks5PasswordAuthResponse) msg).status() == Socks5PasswordAuthStatus.SUCCESS) { + flow.first(SOCKS5CONNECTRequestWithChainedProxy); + flow.advance(); + return; } - flow.fail(); + } + flow.fail(); } @Override void onSuccess(ConnectionFlow flow) { - // Do not advance the flow until the SOCKS response has been parsed + // Do not advance the flow until the SOCKS response has been parsed } - }; - - /** - * Establishes a SOCKS5 connection after {@link #SOCKS5InitialRequest} and - * (optionally) {@link #SOCKS5SendPasswordCredentials} have completed. - */ - private ConnectionFlowStep SOCKS5CONNECTRequestWithChainedProxy = new ConnectionFlowStep( - this, AWAITING_CONNECT_OK) { + }; + + /** + * Establishes a SOCKS5 connection after {@link #SOCKS5InitialRequest} and (optionally) {@link + * #SOCKS5SendPasswordCredentials} have completed. + */ + private final ConnectionFlowStep SOCKS5CONNECTRequestWithChainedProxy = + new ConnectionFlowStep<>(this, AWAITING_CONNECT_OK) { @Override protected Future execute() { - InetSocketAddress destinationAddress = unresolvedAddressFor(serverHostAndPort); - DefaultSocks5CommandRequest connectRequest = new DefaultSocks5CommandRequest( - Socks5CommandType.CONNECT, Socks5AddressType.DOMAIN, destinationAddress.getHostString(), destinationAddress.getPort()); - - addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks5CommandResponseDecoder()); - return writeToChannel(connectRequest); + InetSocketAddress destinationAddress = unresolvedAddressFor(serverHostAndPort); + DefaultSocks5CommandRequest connectRequest = + new DefaultSocks5CommandRequest( + Socks5CommandType.CONNECT, + Socks5AddressType.DOMAIN, + destinationAddress.getHostString(), + destinationAddress.getPort()); + + addFirstOrReplaceHandler(SOCKS_DECODER_NAME, new Socks5CommandResponseDecoder()); + return writeToChannel(connectRequest); } @Override void read(ConnectionFlow flow, Object msg) { - removeHandlerIfPresent(SOCKS_ENCODER_NAME); - removeHandlerIfPresent(SOCKS_DECODER_NAME); - if (msg instanceof Socks5CommandResponse) { - if (((Socks5CommandResponse) msg).status() == Socks5CommandStatus.SUCCESS) { - flow.advance(); - return; - } + removeHandlerIfPresent(SOCKS_ENCODER_NAME); + removeHandlerIfPresent(SOCKS_DECODER_NAME); + if (msg instanceof Socks5CommandResponse) { + if (((Socks5CommandResponse) msg).status() == Socks5CommandStatus.SUCCESS) { + flow.advance(); + return; } - flow.fail(); + } + flow.fail(); } @Override void onSuccess(ConnectionFlow flow) { - // Do not advance the flow until the SOCKS response has been parsed + // Do not advance the flow until the SOCKS response has been parsed } - }; - - /** - *

- * Encrypts the client channel based on our server {@link SSLSession}. - *

- * - *

- * This does not wait for the handshake to finish so that we can go on and - * respond to the CONNECT request. - *

- */ - private ConnectionFlowStep MitmEncryptClientChannel = new ConnectionFlowStep( - this, HANDSHAKING) { + }; + + /** + * Encrypts the client channel based on our server {@link SSLSession}. + * + *

This does not wait for the handshake to finish so that we can go on and respond to the + * CONNECT request. + */ + private final ConnectionFlowStep MitmEncryptClientChannel = + new ConnectionFlowStep<>(this, HANDSHAKING) { @Override boolean shouldExecuteOnEventLoop() { - return false; + return false; } @Override boolean shouldSuppressInitialRequest() { - return true; + return true; } @Override protected Future execute() { - return clientConnection - .encrypt(proxyServer.getMitmManager() - .clientSslEngineFor(initialRequest, sslEngine.getSession()), false) - .addListener( - future -> { - if (future.isSuccess()) { - clientConnection.setMitming(true); - } - }); - } - }; - - /** - * Called when the connection to the server or upstream chained proxy fails. This method may return true to indicate - * that the connection should be retried. If returning true, this method must set up the connection itself. - * - * @param cause the reason that our attempt to connect failed (can be null) - * @return true if we are trying to fall back to another connection - */ - protected boolean connectionFailed(Throwable cause) - throws UnknownHostException { - // unlike a browser, java throws an exception when receiving an unrecognized_name TLS warning, even if the server - // sends back a valid certificate for the expected host. we can retry the connection without SNI to allow the proxy - // to connect to these misconfigured hosts. we should only retry the connection without SNI if the connection - // failure happened when SNI was enabled, to prevent never-ending connection attempts due to SNI warnings. - if (!disableSni && cause instanceof SSLProtocolException) { - // unfortunately java does not expose the specific TLS alert number (112), so we have to look for the - // unrecognized_name string in the exception's message - if (cause.getMessage() != null && cause.getMessage().contains("unrecognized_name")) { - LOG.debug("Failed to connect to server due to an unrecognized_name SSL warning. Retrying connection without SNI."); - - // disable SNI, re-setup the connection, and restart the connection flow - disableSni = true; - resetConnectionForRetry(); - connectAndWrite(initialRequest); - - return true; - } - } - - // the connection issue wasn't due to an unrecognized_name error, or the connection attempt failed even after - // disabling SNI. before falling back to a chained proxy, re-enable SNI. - disableSni = false; - - if (chainedProxy != null) { - LOG.info("Connection to upstream server via chained proxy failed", cause); - // Let the ChainedProxy know that we were unable to connect - chainedProxy.connectionFailed(cause); - } else { - LOG.info("Connection to upstream server failed", cause); - } - - // attempt to connect using a chained proxy, if available - chainedProxy = availableChainedProxies.poll(); - if (chainedProxy != null) { - LOG.info("Retrying connecting using the next available chained proxy"); - - resetConnectionForRetry(); - - connectAndWrite(initialRequest); - return true; + return getClientConnection() + .encrypt( + proxyServer + .getMitmManager() + .clientSslEngineFor(initialRequest, sslEngine.getSession()), + false) + .addListener( + future -> { + if (future.isSuccess()) { + getClientConnection().setMitming(true); + } + }); } - - // no chained proxy fallback or other retry mechanism available - return false; + }; + + /** + * Called when the connection to the server or upstream chained proxy fails. This method may + * return true to indicate that the connection should be retried. If returning true, this method + * must set up the connection itself. + * + * @param cause the reason that our attempt to connect failed (can be null) + * @return true if we are trying to fall back to another connection + */ + protected boolean connectionFailed(Throwable cause) throws UnknownHostException { + // unlike a browser, java throws an exception when receiving an unrecognized_name TLS warning, + // even if the server + // sends back a valid certificate for the expected host. we can retry the connection without SNI + // to allow the proxy + // to connect to these misconfigured hosts. we should only retry the connection without SNI if + // the + // connection + // failure happened when SNI was enabled, to prevent never-ending connection attempts due to SNI + // warnings. + if (!disableSni && (cause instanceof SSLProtocolException) + || (cause instanceof SSLHandshakeException)) { + // unfortunately java does not expose the specific TLS alert number (112), so we have to look + // for the + // unrecognized_name string in the exception's message + if (cause.getMessage() != null && cause.getMessage().contains("unrecognized_name")) { + LOG.debug( + "Failed to connect to server due to an unrecognized_name SSL warning. Retrying connection without SNI."); + + // disable SNI, re-setup the connection, and restart the connection flow + disableSni = true; + resetConnectionForRetry(); + connectAndWrite(initialRequest); + + return true; + } } - /** - * Convenience method to prepare to retry this connection. Closes the connection's channel and sets up - * the connection again using {@link #setupConnectionParameters()}. - * - * @throws UnknownHostException when {@link #setupConnectionParameters()} is unable to resolve the hostname - */ - private void resetConnectionForRetry() throws UnknownHostException { - // Remove ourselves as handler on the old context - this.ctx.pipeline().remove(this); - this.ctx.close(); - this.ctx = null; - - this.setupConnectionParameters(); + // If SSL handshake fails with a non-SSL server (like for websockets), retry without SSL. + // This handles the case described in https://github.com/LittleProxy/LittleProxy/issues/71 + // where CONNECT requests to non-SSL servers (ws://) fail because the proxy tries to use SSL. + // We detect this by checking for specific error patterns that indicate the server doesn't speak + // SSL. + if (shouldRetryWithoutSsl(cause)) { + LOG.debug( + "SSL handshake failed with non-SSL server. Retrying connection without SSL. Cause: {}", + cause.getMessage()); + + // Set a flag to skip SSL in the next attempt + disableSslForNonTls = true; + resetConnectionForRetry(); + connectAndWrite(initialRequest); + + return true; } - /** - * Set up our connection parameters based on server address and chained - * proxies. - * - * @throws UnknownHostException when unable to resolve the hostname to an IP address - */ - private void setupConnectionParameters() throws UnknownHostException { - if (chainedProxy != null - && chainedProxy != ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION) { - this.transportProtocol = chainedProxy.getTransportProtocol(); - this.chainedProxyType = chainedProxy.getChainedProxyType(); - this.localAddress = chainedProxy.getLocalAddress(); - this.remoteAddress = chainedProxy.getChainedProxyAddress(); - this.remoteAddressResolver = DefaultAddressResolverGroup.INSTANCE; - this.username = chainedProxy.getUsername(); - this.password = chainedProxy.getPassword(); - } else { - this.transportProtocol = TransportProtocol.TCP; - this.chainedProxyType = ChainedProxyType.HTTP; - this.username = null; - this.password = null; - - // Report DNS resolution to HttpFilters - this.remoteAddress = this.currentFilters.proxyToServerResolutionStarted(serverHostAndPort); - - // save the hostname and port of the unresolved address in hostAndPort, in case name resolution fails - String hostAndPort = null; - try { - if (this.remoteAddress == null) { - hostAndPort = serverHostAndPort; - this.remoteAddress = addressFor(serverHostAndPort, proxyServer); - } else if (this.remoteAddress.isUnresolved()) { - // filter returned an unresolved address, so resolve it using the proxy server's resolver - hostAndPort = HostAndPort.fromParts(this.remoteAddress.getHostName(), this.remoteAddress.getPort()).toString(); - this.remoteAddress = proxyServer.getServerResolver().resolve(this.remoteAddress.getHostName(), - this.remoteAddress.getPort()); - } - } catch (UnknownHostException e) { - // unable to resolve the hostname to an IP address. notify the filters of the failure before allowing the - // exception to bubble up. - this.currentFilters.proxyToServerResolutionFailed(hostAndPort); - - throw e; - } + // the connection issue wasn't due to an unrecognized_name error, or the connection attempt + // failed even after + // disabling SNI. before falling back to a chained proxy, re-enable SNI. + disableSni = false; + + if (chainedProxy != null) { + LOG.info("Connection to upstream server via chained proxy failed", cause); + // Let the ChainedProxy know that we were unable to connect + chainedProxy.connectionFailed(cause); + } else { + LOG.info("Connection to upstream server failed", cause); + } - this.currentFilters.proxyToServerResolutionSucceeded(serverHostAndPort, this.remoteAddress); + // attempt to connect using a chained proxy, if available + chainedProxy = availableChainedProxies.poll(); + if (chainedProxy != null) { + LOG.info("Retrying connecting using the next available chained proxy"); + resetConnectionForRetry(); - this.localAddress = proxyServer.getLocalAddress(); - } + connectAndWrite(initialRequest); + return true; } - /** - * Initialize our {@link ChannelPipeline} to connect the upstream server. - * LittleProxy acts as a client here. - * - * A {@link ChannelPipeline} invokes the read (Inbound) handlers in - * ascending ordering of the list and then the write (Outbound) handlers in - * descending ordering. - * - * Regarding the Javadoc of {@link HttpObjectAggregator} it's needed to have - * the {@link HttpResponseEncoder} or {@link HttpRequestEncoder} before the - * {@link HttpObjectAggregator} in the {@link ChannelPipeline}. - */ - private void initChannelPipeline(ChannelPipeline pipeline, HttpRequest httpRequest) { - - if (trafficHandler != null) { - pipeline.addLast("global-traffic-shaping", trafficHandler); - } - - pipeline.addLast("bytesReadMonitor", bytesReadMonitor); - pipeline.addLast("bytesWrittenMonitor", bytesWrittenMonitor); - - if ( proxyServer.isSendProxyProtocol()) { - pipeline.addLast("proxy-protocol-encoder", new HAProxyMessageEncoder()); - } - pipeline.addLast("encoder", new HttpRequestEncoder()); - pipeline.addLast("decoder", new HeadAwareHttpResponseDecoder( - proxyServer.getMaxInitialLineLength(), - proxyServer.getMaxHeaderSize(), - proxyServer.getMaxChunkSize())); - - // Enable aggregation for filtering if necessary - int numberOfBytesToBuffer = proxyServer.getFiltersSource() - .getMaximumResponseBufferSizeInBytes(); - if (numberOfBytesToBuffer > 0) { - aggregateContentForFiltering(pipeline, numberOfBytesToBuffer); - } + resetInitialRequest(); + + // no chained proxy fallback or other retry mechanism available + return false; + } + + /** + * Convenience method to prepare to retry this connection. Closes the connection's channel and + * sets up the connection again using {@link #setupConnectionParameters()}. + * + * @throws UnknownHostException when {@link #setupConnectionParameters()} is unable to resolve the + * hostname + */ + private void resetConnectionForRetry() throws UnknownHostException { + // Clear cached flow context so that setupConnectionParameters() creates a fresh one + clientConnection.clearFlowContextForServerConnection(this); + + // Remove ourselves as handler on the old context + ctx.pipeline().remove(this); + ctx.close(); + ctx = null; + + setupConnectionParameters(); + } + + /** + * Checks if we should retry the connection without SSL. This is used when SSL handshake fails + * because the server doesn't speak SSL (like for websockets on non-SSL ports). + * + * @param cause the cause of the connection failure + * @return true if we should retry without SSL + */ + boolean shouldRetryWithoutSsl(@Nullable Throwable cause) { + if (cause == null) { + return false; + } - pipeline.addLast("responseReadMonitor", responseReadMonitor); - pipeline.addLast("requestWrittenMonitor", requestWrittenMonitor); + // Don't retry if we've already tried without SSL + if (disableSslForNonTls) { + return false; + } - // Set idle timeout - pipeline.addLast( - "idle", - new IdleStateHandler(0, 0, proxyServer - .getIdleConnectionTimeout())); + // Only retry if this is an SSL handshake failure + if (!(cause instanceof javax.net.ssl.SSLException)) { + return false; + } - pipeline.addLast("handler", this); + // Check for patterns that indicate the server doesn't speak SSL + String message = cause.getMessage(); + if (message != null) { + String lowerMessage = message.toLowerCase(ROOT); + // "Remote host terminated the handshake" - server doesn't support SSL + // "end of file" - server closed connection unexpectedly + // "connection reset" - server doesn't speak SSL + // "not an SSL/TLS record" - server sent HTTP response to SSL handshake + return lowerMessage.contains("remote host terminated") + || lowerMessage.contains("end of file") + || lowerMessage.contains("connection reset") + || lowerMessage.contains("not an ssl"); } - /** - *

- * Do all the stuff that needs to be done after our {@link ConnectionFlow} - * has succeeded. - *

- * - * @param shouldForwardInitialRequest - * whether or not we should forward the initial HttpRequest to - * the server after the connection has been established. - */ - void connectionSucceeded(boolean shouldForwardInitialRequest) { - become(AWAITING_INITIAL); - if (this.chainedProxy != null) { - // Notify the ChainedProxy that we successfully connected - try { - this.chainedProxy.connectionSucceeded(); - } catch (Exception e) { - LOG.error("Unable to record connectionSucceeded", e); - } + return false; + } + + /** + * Set up our connection parameters based on server address and chained proxies. + * + * @throws UnknownHostException when unable to resolve the hostname to an IP address + */ + private void setupConnectionParameters() throws UnknownHostException { + if (chainedProxy != null && chainedProxy != ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION) { + transportProtocol = chainedProxy.getTransportProtocol(); + chainedProxyType = chainedProxy.getChainedProxyType(); + localAddress = chainedProxy.getLocalAddress(); + remoteAddress = chainedProxy.getChainedProxyAddress(); + remoteAddressResolver = DefaultAddressResolverGroup.INSTANCE; + username = chainedProxy.getUsername(); + password = chainedProxy.getPassword(); + } else { + transportProtocol = TransportProtocol.TCP; + chainedProxyType = ChainedProxyType.HTTP; + username = null; + password = null; + + // Report DNS resolution to HttpFilters + long dnsStartTime = System.currentTimeMillis(); + clientConnection.flowContext().setTimingData("dns_resolution_start_time_ms", dnsStartTime); + remoteAddress = currentFilters.proxyToServerResolutionStarted(serverHostAndPort); + + // save the hostname and port of the unresolved address in hostAndPort, in case name + // resolution fails + String hostAndPort = null; + try { + if (remoteAddress == null) { + hostAndPort = serverHostAndPort; + remoteAddress = addressFor(serverHostAndPort, proxyServer); + } else if (remoteAddress.isUnresolved()) { + // filter returned an unresolved address, so resolve it using the proxy server's resolver + hostAndPort = + HostAndPort.fromParts(remoteAddress.getHostName(), remoteAddress.getPort()) + .toString(); + remoteAddress = + proxyServer + .getServerResolver() + .resolve(remoteAddress.getHostName(), remoteAddress.getPort()); } - clientConnection.serverConnectionSucceeded(this, - shouldForwardInitialRequest); + } catch (UnknownHostException e) { + // unable to resolve the hostname to an IP address. notify the filters of the failure before + // allowing the + // exception to bubble up. + currentFilters.proxyToServerResolutionFailed(hostAndPort); + + throw e; + } + + currentFilters.proxyToServerResolutionSucceeded(serverHostAndPort, remoteAddress); + long dnsEndTime = System.currentTimeMillis(); + FlowContext clientFlowContext = clientConnection.flowContext(); + // Only cache server flow context AFTER DNS resolution succeeds + FullFlowContext serverFlowContext = clientConnection.flowContextForServerConnection(this); + clientFlowContext.setTimingData("dns_resolution_end_time_ms", dnsEndTime); + serverFlowContext.setTimingData("dns_resolution_start_time_ms", dnsStartTime); + serverFlowContext.setTimingData("dns_resolution_end_time_ms", dnsEndTime); + serverFlowContext.setTimingData("dns_resolution_time_ms", dnsEndTime - dnsStartTime); + + localAddress = proxyServer.getLocalAddress(); + } + } + + /** + * Initialize our {@link ChannelPipeline} to connect the upstream server. LittleProxy acts as a + * client here. + * + *

A {@link ChannelPipeline} invokes the read (Inbound) handlers in ascending ordering of the + * list and then the write (Outbound) handlers in descending ordering. + * + *

Regarding the Javadoc of {@link HttpObjectAggregator} it's needed to have the {@link + * HttpResponseEncoder} or {@link HttpRequestEncoder} before the {@link HttpObjectAggregator} in + * the {@link ChannelPipeline}. + */ + private void initChannelPipeline(ChannelPipeline pipeline) { + + if (trafficHandler != null) { + pipeline.addLast("global-traffic-shaping", trafficHandler); + } - if (shouldForwardInitialRequest) { - LOG.debug("Writing initial request: {}", initialRequest); - write(initialRequest); - } else { - LOG.debug("Dropping initial request: {}", initialRequest); - } + pipeline.addLast("bytesReadMonitor", bytesReadMonitor); + pipeline.addLast("bytesWrittenMonitor", bytesWrittenMonitor); - // we're now done with the initialRequest: it's either been forwarded to the upstream server (HTTP requests), or - // completely dropped (HTTPS CONNECTs). if the initialRequest is reference counted (typically because the HttpObjectAggregator is in - // the pipeline to generate FullHttpRequests), we need to manually release it to avoid a memory leak. - if (initialRequest instanceof ReferenceCounted) { - ((ReferenceCounted)initialRequest).release(); - } + if (proxyServer.isSendProxyProtocol()) { + pipeline.addLast(HTTP_PROXY_ENCODER_NAME, new HAProxyMessageEncoder()); + } + pipeline.addLast(HTTP_ENCODER_NAME, new HttpRequestEncoder()); + pipeline.addLast( + HTTP_DECODER_NAME, + new HeadAwareHttpResponseDecoder( + proxyServer.getMaxInitialLineLength(), + proxyServer.getMaxHeaderSize(), + proxyServer.getMaxChunkSize())); + + // Enable aggregation for filtering if necessary + int numberOfBytesToBuffer = + proxyServer.getFiltersSource().getMaximumResponseBufferSizeInBytes(); + if (numberOfBytesToBuffer > 0) { + aggregateContentForFiltering(pipeline, numberOfBytesToBuffer); } - /** - * Build an {@link InetSocketAddress} for the given hostAndPort. - * - * @param hostAndPort String representation of the host and port - * @param proxyServer the current {@link DefaultHttpProxyServer} - * @return a resolved InetSocketAddress for the specified hostAndPort - * @throws UnknownHostException if hostAndPort could not be resolved, or if the input string could not be parsed into - * a host and port. - */ - public static InetSocketAddress addressFor(String hostAndPort, DefaultHttpProxyServer proxyServer) - throws UnknownHostException { - HostAndPort parsedHostAndPort; - try { - parsedHostAndPort = HostAndPort.fromString(hostAndPort); - } catch (IllegalArgumentException e) { - // we couldn't understand the hostAndPort string, so there is no way we can resolve it. - throw new UnknownHostException(hostAndPort); - } - - String host = parsedHostAndPort.getHost(); - int port = parsedHostAndPort.getPortOrDefault(80); + pipeline.addLast(HTTP_RESPONSE_READ_MONITOR_NAME, responseReadMonitor); + pipeline.addLast(HTTP_REQUEST_WRITTEN_MONITOR_NAME, requestWrittenMonitor); + + // Set idle timeout + pipeline.addLast("idle", new IdleStateHandler(0, 0, proxyServer.getIdleConnectionTimeout())); + + pipeline.addLast(MAIN_HANDLER_NAME, this); + } + + /** + * Do all the stuff that needs to be done after our {@link ConnectionFlow} has succeeded. + * + * @param shouldForwardInitialRequest whether we should forward the initial HttpRequest to the + * server after the connection has been established. + */ + void connectionSucceeded(boolean shouldForwardInitialRequest) { + become(AWAITING_INITIAL); + if (chainedProxy != null) { + // Notify the ChainedProxy that we successfully connected + try { + chainedProxy.connectionSucceeded(); + } catch (Exception e) { + LOG.error("Unable to record connectionSucceeded", e); + } + } + getClientConnection().serverConnectionSucceeded(this, shouldForwardInitialRequest); - return proxyServer.getServerResolver().resolve(host, port); + if (shouldForwardInitialRequest) { + LOG.debug("Writing initial request: {}", initialRequest); + write(initialRequest); + } else { + LOG.debug("Dropping initial request: {}", initialRequest); } - /** - * Similar to {@link #addressFor(String, DefaultHttpProxyServer)} except that it does - * not resolve the address. - * @param hostAndPort the host and port to parse. - * @return an unresolved {@link InetSocketAddress}. - */ - private static InetSocketAddress unresolvedAddressFor(String hostAndPort) { - HostAndPort parsedHostAndPort = HostAndPort.fromString(hostAndPort); - String host = parsedHostAndPort.getHost(); - int port = parsedHostAndPort.getPortOrDefault(80); - return InetSocketAddress.createUnresolved(host, port); + // we're now done with the initialRequest: it's either been forwarded to the upstream server + // (HTTP requests), or + // completely dropped (HTTPS CONNECTs). if the initialRequest is reference counted (typically + // because the HttpObjectAggregator is in + // the pipeline to generate FullHttpRequests), we need to manually release it to avoid a memory + // leak. + resetInitialRequest(); + + // Phase 2 per-request MITM: release the CONNECT-created connection to pool so subsequent + // HTTP requests can borrow it via pool.getOrCreateConnection(). + if (releaseToPoolOnConnectComplete) { + releaseToPoolOnConnectComplete = false; + releaseToPool(); } + } - /* ************************************************************************* - * Activity Tracking/Statistics - * - * We track statistics on bytes, requests and responses by adding handlers - * at the appropriate parts of the pipeline (see initChannelPipeline()). - **************************************************************************/ + private void resetInitialRequest() { + if (initialRequest instanceof ReferenceCounted) { + ((ReferenceCounted) initialRequest).release(); + } + } + + /** + * Build an {@link InetSocketAddress} for the given hostAndPort. + * + * @param hostAndPort String representation of the host and port + * @param proxyServer the current {@link DefaultHttpProxyServer} + * @return a resolved InetSocketAddress for the specified hostAndPort + * @throws UnknownHostException if hostAndPort could not be resolved, or if the input string could + * not be parsed into a host and port. + */ + public static InetSocketAddress addressFor(String hostAndPort, DefaultHttpProxyServer proxyServer) + throws UnknownHostException { + HostAndPort parsedHostAndPort; + try { + parsedHostAndPort = HostAndPort.fromString(hostAndPort); + } catch (IllegalArgumentException e) { + // we couldn't understand the hostAndPort string, so there is no way we can resolve it. + throw new UnknownHostException(hostAndPort); + } - private final BytesReadMonitor bytesReadMonitor = new BytesReadMonitor() { + String host = parsedHostAndPort.getHost(); + int port = parsedHostAndPort.getPortOrDefault(80); + + return proxyServer.getServerResolver().resolve(host, port); + } + + /** + * Similar to {@link #addressFor(String, DefaultHttpProxyServer)} except that it does not resolve + * the address. + * + * @param hostAndPort the host and port to parse. + * @return an unresolved {@link InetSocketAddress}. + */ + private static InetSocketAddress unresolvedAddressFor(String hostAndPort) { + HostAndPort parsedHostAndPort = HostAndPort.fromString(hostAndPort); + String host = parsedHostAndPort.getHost(); + int port = parsedHostAndPort.getPortOrDefault(80); + return InetSocketAddress.createUnresolved(host, port); + } + + void switchToWebSocketProtocol() { + final List orderedHandlersToRemove = + Arrays.asList( + HTTP_REQUEST_WRITTEN_MONITOR_NAME, + HTTP_RESPONSE_READ_MONITOR_NAME, + HTTP_PROXY_ENCODER_NAME, + HTTP_ENCODER_NAME, + HTTP_DECODER_NAME); + if (channel.pipeline().get(MAIN_HANDLER_NAME) != null) { + channel + .pipeline() + .replace( + MAIN_HANDLER_NAME, + "pipe-to-client", + new WebSocketFramePipeHandler(clientConnection, currentFilters, false)); + } + orderedHandlersToRemove.forEach(this::removeHandlerIfPresent); + tunneling = true; + } + + /* ************************************************************************* + * Activity Tracking/Statistics + * + * We track statistics on bytes, requests and responses by adding handlers + * at the appropriate parts of the pipeline (see initChannelPipeline()). + **************************************************************************/ + + private final BytesReadMonitor bytesReadMonitor = + new BytesReadMonitor() { @Override protected void bytesRead(int numberOfBytes) { - FullFlowContext flowContext = new FullFlowContext(clientConnection, - ProxyToServerConnection.this); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.bytesReceivedFromServer(flowContext, numberOfBytes); - } + FullFlowContext flowContext = + getClientConnection().flowContextForServerConnection(ProxyToServerConnection.this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.bytesReceivedFromServer(flowContext, numberOfBytes); + } } - }; + }; - private ResponseReadMonitor responseReadMonitor = new ResponseReadMonitor() { + private final ResponseReadMonitor responseReadMonitor = + new ResponseReadMonitor() { @Override protected void responseRead(HttpResponse httpResponse) { - FullFlowContext flowContext = new FullFlowContext(clientConnection, - ProxyToServerConnection.this); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.responseReceivedFromServer(flowContext, httpResponse); - } + FullFlowContext flowContext = + getClientConnection().flowContextForServerConnection(ProxyToServerConnection.this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.responseReceivedFromServer(flowContext, httpResponse); + } } - }; + }; - private BytesWrittenMonitor bytesWrittenMonitor = new BytesWrittenMonitor() { + private final BytesWrittenMonitor bytesWrittenMonitor = + new BytesWrittenMonitor() { @Override protected void bytesWritten(int numberOfBytes) { - FullFlowContext flowContext = new FullFlowContext(clientConnection, - ProxyToServerConnection.this); - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.bytesSentToServer(flowContext, numberOfBytes); - } + FullFlowContext flowContext = + getClientConnection().flowContextForServerConnection(ProxyToServerConnection.this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.bytesSentToServer(flowContext, numberOfBytes); + } } - }; + }; - private RequestWrittenMonitor requestWrittenMonitor = new RequestWrittenMonitor() { + private final RequestWrittenMonitor requestWrittenMonitor = + new RequestWrittenMonitor() { @Override protected void requestWriting(HttpRequest httpRequest) { - FullFlowContext flowContext = new FullFlowContext(clientConnection, - ProxyToServerConnection.this); - try { - for (ActivityTracker tracker : proxyServer - .getActivityTrackers()) { - tracker.requestSentToServer(flowContext, httpRequest); - } - } catch (Throwable t) { - LOG.warn("Error while invoking ActivityTracker on request", t); + FullFlowContext flowContext = + getClientConnection().flowContextForServerConnection(ProxyToServerConnection.this); + try { + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + tracker.requestSentToServer(flowContext, httpRequest); } + } catch (Throwable t) { + LOG.warn("Error while invoking ActivityTracker on request", t); + } - currentFilters.proxyToServerRequestSending(); + currentFilters.proxyToServerRequestSending(); } @Override - protected void requestWritten(HttpRequest httpRequest) { - } + protected void requestWritten(HttpRequest httpRequest) {} @Override protected void contentWritten(HttpContent httpContent) { - if (httpContent instanceof LastHttpContent) { - currentFilters.proxyToServerRequestSent(); - } + if (httpContent instanceof LastHttpContent) { + currentFilters.proxyToServerRequestSent(); + } } - }; + }; + + void recordServerConnected() { + ClientToProxyConnection clientConn = getClientConnection(); + FullFlowContext flowContext = clientConn.flowContextForServerConnection(this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + try { + tracker.serverConnected(flowContext, remoteAddress); + } catch (Exception e) { + LOG.error("Unable to recordServerConnected", e); + } + } + } + void recordServerDisconnected() { + ClientToProxyConnection clientConn = getClientConnection(); + FullFlowContext flowContext = clientConn.flowContextForServerConnection(this); + try { + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + try { + tracker.serverDisconnected(flowContext, remoteAddress); + } catch (Exception e) { + LOG.error("Unable to recordServerDisconnected", e); + } + } + } finally { + clientConn.clearFlowContextForServerConnection(this); + } + } + + void recordConnectionSaturated() { + ClientToProxyConnection clientConn = getClientConnection(); + FullFlowContext flowContext = clientConn.flowContextForServerConnection(this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + try { + tracker.connectionSaturated(flowContext); + } catch (Exception e) { + LOG.error("Unable to recordConnectionSaturated", e); + } + } + } + + void recordConnectionWritable() { + ClientToProxyConnection clientConn = getClientConnection(); + FullFlowContext flowContext = clientConn.flowContextForServerConnection(this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + try { + tracker.connectionWritable(flowContext); + } catch (Exception e) { + LOG.error("Unable to recordConnectionWritable", e); + } + } + } + + void recordConnectionTimedOut() { + ClientToProxyConnection clientConn = getClientConnection(); + FullFlowContext flowContext = clientConn.flowContextForServerConnection(this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + try { + tracker.connectionTimedOut(flowContext); + } catch (Exception e) { + LOG.error("Unable to recordConnectionTimedOut", e); + } + } + } + + void recordConnectionExceptionCaught(Throwable cause) { + ClientToProxyConnection clientConn = getClientConnection(); + FullFlowContext flowContext = clientConn.flowContextForServerConnection(this); + for (ActivityTracker tracker : proxyServer.getActivityTrackers()) { + try { + tracker.connectionExceptionCaught(flowContext, cause); + } catch (Exception e) { + LOG.error("Unable to recordConnectionExceptionCaught", e); + } + } + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ProxyUtils.java b/src/main/java/org/littleshoot/proxy/impl/ProxyUtils.java index 1c0df1a9..fe402348 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ProxyUtils.java +++ b/src/main/java/org/littleshoot/proxy/impl/ProxyUtils.java @@ -1,630 +1,665 @@ package org.littleshoot.proxy.impl; -import com.google.common.base.Splitter; -import com.google.common.collect.ImmutableList; -import com.google.common.collect.ImmutableSet; +import static java.util.stream.Collectors.toList; + +import com.google.errorprone.annotations.CheckReturnValue; import io.netty.buffer.ByteBuf; import io.netty.buffer.Unpooled; -import io.netty.channel.udt.nio.NioUdtProvider; -import io.netty.handler.codec.http.*; +import io.netty.handler.codec.http.DefaultFullHttpResponse; +import io.netty.handler.codec.http.DefaultHttpResponse; +import io.netty.handler.codec.http.FullHttpResponse; +import io.netty.handler.codec.http.HttpHeaderNames; +import io.netty.handler.codec.http.HttpHeaderValues; +import io.netty.handler.codec.http.HttpHeaders; +import io.netty.handler.codec.http.HttpMessage; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpObject; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpVersion; +import io.netty.handler.codec.http.LastHttpContent; import io.netty.util.AsciiString; +import java.net.InetSocketAddress; +import java.nio.charset.StandardCharsets; +import java.util.ArrayList; +import java.util.Collection; +import java.util.Collections; +import java.util.List; +import java.util.Properties; +import java.util.Set; +import java.util.regex.Pattern; +import java.util.stream.Stream; import org.apache.commons.lang3.StringUtils; import org.apache.commons.lang3.math.NumberUtils; +import org.apache.logging.log4j.util.Strings; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.Nullable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import java.io.IOException; -import java.net.InetAddress; -import java.nio.charset.StandardCharsets; -import java.text.SimpleDateFormat; -import java.util.*; -import java.util.regex.Pattern; - -/** - * Utilities for the proxy. - */ +/** Utilities for the proxy. */ public class ProxyUtils { - /** - * Hop-by-hop headers that should be removed when proxying, as defined by the HTTP 1.1 spec, section 13.5.1 - * (http://www.w3.org/Protocols/rfc2616/rfc2616-sec13.html#sec13.5.1). Transfer-Encoding is NOT included in this list, since LittleProxy - * does not typically modify the transfer encoding. See also {@link #shouldRemoveHopByHopHeader(String)}. - * - * Header names are stored as lowercase to make case-insensitive comparisons easier. - */ - @SuppressWarnings("deprecation") // Don't remove header names from this set until they're removed from Netty, just in case someone's still using them. - private static final Set SHOULD_NOT_PROXY_HOP_BY_HOP_HEADERS = ImmutableSet.of( - HttpHeaderNames.CONNECTION.toString(), - HttpHeaderNames.KEEP_ALIVE.toString(), - HttpHeaderNames.PROXY_AUTHENTICATE.toString(), - HttpHeaderNames.PROXY_AUTHORIZATION.toString(), - HttpHeaderNames.TE.toString(), - HttpHeaderNames.TRAILER.toString(), - /* Note: Not removing Transfer-Encoding since LittleProxy does not normally re-chunk content. - HttpHeaderNames.TRANSFER_ENCODING.toString(), */ - HttpHeaderNames.UPGRADE.toString() - ); - - private static final Logger LOG = LoggerFactory.getLogger(ProxyUtils.class); - - private static final TimeZone GMT = TimeZone.getTimeZone("GMT"); - - /** - * Splits comma-separated header values (such as Connection) into their individual tokens. - */ - private static final Splitter COMMA_SEPARATED_HEADER_VALUE_SPLITTER = Splitter.on(',').trimResults().omitEmptyStrings(); - - /** - * Date format pattern used to parse HTTP date headers in RFC 1123 format. - */ - private static final String PATTERN_RFC1123 = "EEE, dd MMM yyyy HH:mm:ss zzz"; - - // Schemes are case-insensitive: - // http://tools.ietf.org/html/rfc3986#section-3.1 - private static Pattern HTTP_PREFIX = Pattern.compile("^https?://.*", - Pattern.CASE_INSENSITIVE); - - /** - * Strips the host from a URI string. This will turn "http://host.com/path" - * into "/path". - * - * @param uri - * The URI to transform. - * @return A string with the URI stripped. - */ - public static String stripHost(final String uri) { - if (!HTTP_PREFIX.matcher(uri).matches()) { - // It's likely a URI path, not the full URI (i.e. the host is - // already stripped). - return uri; - } - final String noHttpUri = StringUtils.substringAfter(uri, "://"); - final int slashIndex = noHttpUri.indexOf("/"); - if (slashIndex == -1) { - return "/"; - } - return noHttpUri.substring(slashIndex); - } - - /** - * Formats the given date according to the RFC 1123 pattern. - * - * @param date - * The date to format. - * @return An RFC 1123 formatted date string. - * - * @see #PATTERN_RFC1123 - */ - public static String formatDate(final Date date) { - return formatDate(date, PATTERN_RFC1123); - } - - /** - * Formats the given date according to the specified pattern. The pattern - * must conform to that used by the {@link SimpleDateFormat simple date - * format} class. - * - * @param date - * The date to format. - * @param pattern - * The pattern to use for formatting the date. - * @return A formatted date string. - * - * @throws IllegalArgumentException - * If the given date pattern is invalid. - * - * @see SimpleDateFormat - */ - public static String formatDate(final Date date, final String pattern) { - if (date == null) - throw new IllegalArgumentException("date is null"); - if (pattern == null) - throw new IllegalArgumentException("pattern is null"); - - final SimpleDateFormat formatter = new SimpleDateFormat(pattern, - Locale.US); - formatter.setTimeZone(GMT); - return formatter.format(date); - } - - /** - * If an HttpObject implements the market interface LastHttpContent, it - * represents the last chunk of a transfer. - * - * @see io.netty.handler.codec.http.LastHttpContent - */ - public static boolean isLastChunk(final HttpObject httpObject) { - return httpObject instanceof LastHttpContent; - } - - /** - * If an HttpObject is not the last chunk, then that means there are other - * chunks that will follow. - * - * @see io.netty.handler.codec.http.FullHttpMessage - */ - public static boolean isChunked(final HttpObject httpObject) { - return !isLastChunk(httpObject); - } - - /** - * Parses the host and port an HTTP request is being sent to. - * - * @param httpRequest - * The request. - * @return The host and port string. - */ - public static String parseHostAndPort(final HttpRequest httpRequest) { - return parseHostAndPort(httpRequest.uri()); - } - - /** - * Parses the host and port an HTTP request is being sent to. - * - * @param uri - * The URI. - * @return The host and port string. - */ - public static String parseHostAndPort(final String uri) { - final String tempUri; - if (!HTTP_PREFIX.matcher(uri).matches()) { - // Browsers particularly seem to send requests in this form when - // they use CONNECT. - tempUri = uri; - } else { - // We can't just take a substring from a hard-coded index because it - // could be either http or https. - tempUri = StringUtils.substringAfter(uri, "://"); - } - final String hostAndPort; - if (tempUri.contains("/")) { - hostAndPort = tempUri.substring(0, tempUri.indexOf("/")); - } else { - hostAndPort = tempUri; - } - return hostAndPort; - } - - /** - * Make a copy of the response including all mutable fields. - * - * @param original - * The original response to copy from. - * @return The copy with all mutable fields from the original. - */ - public static HttpResponse copyMutableResponseFields( - final HttpResponse original) { - - HttpResponse copy; - if (original instanceof DefaultFullHttpResponse) { - ByteBuf content = ((DefaultFullHttpResponse) original).content(); - copy = new DefaultFullHttpResponse(original.protocolVersion(), - original.status(), content); - } else { - copy = new DefaultHttpResponse(original.protocolVersion(), - original.status()); - } - final Collection headerNames = original.headers().names(); - for (final String name : headerNames) { - final List values = original.headers().getAll(name); - copy.headers().set(name, values); - } - return copy; - } - - /** - * Adds the Via header to specify that the message has passed through the proxy. The specified alias will be - * appended to the Via header line. The alias may be the hostname of the machine proxying the request, or a - * pseudonym. From RFC 7230, section 5.7.1: - *

-         The received-by portion of the field value is normally the host and
-         optional port number of a recipient server or client that
-         subsequently forwarded the message.  However, if the real host is
-         considered to be sensitive information, a sender MAY replace it with
-         a pseudonym.
-     * 
- * - * - * @param httpMessage HTTP message to add the Via header to - * @param alias the alias to provide in the Via header for this proxy - */ - public static void addVia(HttpMessage httpMessage, String alias) { - String newViaHeader = String.valueOf(httpMessage.protocolVersion().majorVersion()) + - '.' + - httpMessage.protocolVersion().minorVersion() + - ' ' + - alias; - - final List vias; - if (httpMessage.headers().contains(HttpHeaderNames.VIA)) { - List existingViaHeaders = httpMessage.headers().getAll(HttpHeaderNames.VIA); - vias = new ArrayList<>(existingViaHeaders); - vias.add(newViaHeader); - } else { - vias = Collections.singletonList(newViaHeader); - } - - httpMessage.headers().set(HttpHeaderNames.VIA, vias); + /** + * Hop-by-hop headers that should be removed when proxying, as defined by the (HTTP 1.1 spec, section + * 13.5.1). Transfer-Encoding is NOT included in this list, since LittleProxy does not + * typically modify the transfer encoding. See also {@link #shouldRemoveHopByHopHeader(String)}. + * + *

Header names are stored as lowercase to make case-insensitive comparisons easier. + */ + @SuppressWarnings( + "deprecation") // Don't remove header names from this set until they're removed from Netty, + // just in case someone's still using them. + private static final Set SHOULD_NOT_PROXY_HOP_BY_HOP_HEADERS = + Set.of( + HttpHeaderNames.CONNECTION.toString(), + HttpHeaderNames.KEEP_ALIVE.toString(), + HttpHeaderNames.PROXY_AUTHENTICATE.toString(), + HttpHeaderNames.PROXY_AUTHORIZATION.toString(), + HttpHeaderNames.TE.toString(), + HttpHeaderNames.TRAILER.toString(), + /* Note: Not removing Transfer-Encoding since LittleProxy does not normally re-chunk content. + HttpHeaderNames.TRANSFER_ENCODING.toString(), */ + HttpHeaderNames.UPGRADE.toString()); + + private static final Logger LOG = LoggerFactory.getLogger(ProxyUtils.class); + + // Schemes are case-insensitive: + // https://tools.ietf.org/html/rfc3986#section-3.1 + private static final Pattern HTTP_PREFIX = + Pattern.compile("^(http|ws)s?://.*", Pattern.CASE_INSENSITIVE); + + private ProxyUtils() {} + + /** + * Strips the host from a URI string. This will turn "https://host.com/path" into "/path". + * + * @param uri The URI to transform. + * @return A string with the URI stripped. + */ + public static String stripHost(final String uri) { + if (!HTTP_PREFIX.matcher(uri).matches()) { + // It's likely a URI path, not the full URI (i.e. the host is + // already stripped). + return uri; } - - /** - * Returns true if the specified string is either "true" or - * "on" ignoring case. - * - * @param val - * The string in question. - * @return true if the specified string is either "true" or - * "on" ignoring case, otherwise false. - */ - public static boolean isTrue(final String val) { - return checkTrueOrFalse(val, "true", "on"); + final String noHttpUri = StringUtils.substringAfter(uri, "://"); + final int slashIndex = noHttpUri.indexOf("/"); + if (slashIndex == -1) { + return "/"; } - - /** - * Returns true if the specified string is either "false" or - * "off" ignoring case. - * - * @param val - * The string in question. - * @return true if the specified string is either "false" or - * "off" ignoring case, otherwise false. - */ - public static boolean isFalse(final String val) { - return checkTrueOrFalse(val, "false", "off"); + return noHttpUri.substring(slashIndex); + } + + /** + * If an HttpObject implements the market interface LastHttpContent, it represents the last chunk + * of a transfer. + * + * @see io.netty.handler.codec.http.LastHttpContent + */ + public static boolean isLastChunk(final HttpObject httpObject) { + return httpObject instanceof LastHttpContent; + } + + /** + * If an HttpObject is not the last chunk, then that means there are other chunks that will + * follow. + * + * @see io.netty.handler.codec.http.FullHttpMessage + */ + public static boolean isChunked(final HttpObject httpObject) { + return !isLastChunk(httpObject); + } + + /** + * Parses the host and port an HTTP request is being sent to. + * + * @param httpRequest The request. + * @return The host and port string. + */ + @NonNull + @CheckReturnValue + public static String parseHostAndPort(@NonNull final HttpRequest httpRequest) { + return parseHostAndPort(httpRequest.uri()); + } + + /** + * Parses the host and port an HTTP request is being sent to. + * + * @param uri The URI. + * @return The host and port string. + */ + @NonNull + @CheckReturnValue + public static String parseHostAndPort(@NonNull final String uri) { + final String tempUri; + if (!HTTP_PREFIX.matcher(uri).matches()) { + // Browsers particularly seem to send requests in this form when + // they use CONNECT. + tempUri = uri; + } else { + // We can't just take a substring from a hard-coded index because it + // could be either http or https. + tempUri = StringUtils.substringAfter(uri, "://"); } - - public static boolean extractBooleanDefaultFalse(final Properties props, - final String key) { - final String throttle = props.getProperty(key); - if (StringUtils.isNotBlank(throttle)) { - return throttle.trim().equalsIgnoreCase("true"); - } - return false; + final String hostAndPort; + if (tempUri.contains("/")) { + hostAndPort = tempUri.substring(0, tempUri.indexOf("/")); + } else { + hostAndPort = tempUri; } + return hostAndPort; + } - public static boolean extractBooleanDefaultTrue(final Properties props, - final String key) { - final String throttle = props.getProperty(key); - if (StringUtils.isNotBlank(throttle)) { - return throttle.trim().equalsIgnoreCase("true"); - } - return true; - } - - public static int extractInt(final Properties props, final String key) { - return extractInt(props, key, -1); - } - - public static int extractInt(final Properties props, final String key, int defaultValue) { - final String readThrottleString = props.getProperty(key); - if (StringUtils.isNotBlank(readThrottleString) && NumberUtils.isCreatable(readThrottleString)) { - return Integer.parseInt(readThrottleString); - } - return defaultValue; + public static InetSocketAddress resolveSocketAddress(String address) { + if (Strings.isBlank(address)) { + return null; } + String[] parts = address.split(":"); + String host = parts[0]; + int port = Integer.parseInt(parts[1]); - public static boolean isCONNECT(HttpObject httpObject) { - return httpObject instanceof HttpRequest && HttpMethod.CONNECT.equals(((HttpRequest) httpObject).method()); + // remove hooks for IPv6 + if (host.startsWith("[") && host.endsWith("]")) { + host = host.substring(1, host.length() - 1); } - /** - * Returns true if the specified HttpRequest is a HEAD request. - * - * @param httpRequest http request - * @return true if request is a HEAD, otherwise false - */ - public static boolean isHEAD(HttpRequest httpRequest) { - return HttpMethod.HEAD.equals(httpRequest.method()); + return new InetSocketAddress(host, port); + } + + /** + * Make a copy of the response including all mutable fields. + * + * @param original The original response to copy from. + * @return The copy with all mutable fields from the original. + */ + public static HttpResponse copyMutableResponseFields(final HttpResponse original) { + + HttpResponse copy; + if (original instanceof DefaultFullHttpResponse) { + ByteBuf content = ((DefaultFullHttpResponse) original).content(); + copy = new DefaultFullHttpResponse(original.protocolVersion(), original.status(), content); + } else { + copy = new DefaultHttpResponse(original.protocolVersion(), original.status()); } - - private static boolean checkTrueOrFalse(final String val, - final String str1, final String str2) { - final String str = val.trim(); - return StringUtils.isNotBlank(str) - && (str.equalsIgnoreCase(str1) || str.equalsIgnoreCase(str2)); + final Collection headerNames = original.headers().names(); + for (final String name : headerNames) { + final List values = original.headers().getAll(name); + copy.headers().set(name, values); } - - /** - * Returns true if the HTTP message cannot contain an entity body, according to the HTTP spec. This code is taken directly - * from {@link io.netty.handler.codec.http.HttpObjectDecoder#isContentAlwaysEmpty(HttpMessage)}. - * - * @param msg HTTP message - * @return true if the HTTP message is always empty, false if the message may have entity content. - */ - public static boolean isContentAlwaysEmpty(HttpMessage msg) { - if (msg instanceof HttpResponse) { - HttpResponse res = (HttpResponse) msg; - int code = res.status().code(); - - // Correctly handle return codes of 1xx. - // - // See: - // - http://www.w3.org/Protocols/rfc2616/rfc2616-sec4.html Section 4.4 - // - https://github.com/netty/netty/issues/222 - if (code >= 100 && code < 200) { - // According to RFC 7231, section 6.1, 1xx responses have no content (https://tools.ietf.org/html/rfc7231#section-6.2): - // 1xx responses are terminated by the first empty line after - // the status-line (the empty line signaling the end of the header - // section). - - // Hixie 76 websocket handshake responses contain a 16-byte body, so their content is not empty; but Hixie 76 - // was a draft specification that was superceded by RFC 6455. Since it is rarely used and doesn't conform to - // RFC 7231, we do not support or make special allowance for Hixie 76 responses. - return true; - } - - switch (code) { - case 204: case 205: case 304: - return true; - } - } - return false; + return copy; + } + + /** + * Adds the Via header to specify that the message has passed through the proxy. The specified + * alias will be appended to the Via header line. The alias may be the hostname of the machine + * proxying the request, or a pseudonym. From RFC 7230, section 5.7.1: + * + *

+   * The received-by portion of the field value is normally the host and
+   * optional port number of a recipient server or client that
+   * subsequently forwarded the message.  However, if the real host is
+   * considered to be sensitive information, a sender MAY replace it with
+   * a pseudonym.
+   * 
+ * + * @param httpMessage HTTP message to add the Via header to + * @param alias the alias to provide in the Via header for this proxy + */ + public static void addVia(HttpMessage httpMessage, String alias) { + String newViaHeader = + String.valueOf(httpMessage.protocolVersion().majorVersion()) + + '.' + + httpMessage.protocolVersion().minorVersion() + + ' ' + + alias; + + final List vias; + if (httpMessage.headers().contains(HttpHeaderNames.VIA)) { + List existingViaHeaders = httpMessage.headers().getAll(HttpHeaderNames.VIA); + vias = new ArrayList<>(existingViaHeaders); + vias.add(newViaHeader); + } else { + vias = Collections.singletonList(newViaHeader); } - /** - * Returns true if the HTTP response from the server is expected to indicate its own message length/end-of-message. Returns false - * if the server is expected to indicate the end of the HTTP entity by closing the connection. - *

- * This method is based on the allowed message length indicators in the HTTP specification, section 4.4: - *

-         4.4 Message Length
-         The transfer-length of a message is the length of the message-body as it appears in the message; that is, after any transfer-codings have been applied. When a message-body is included with a message, the transfer-length of that body is determined by one of the following (in order of precedence):
-
-         1.Any response message which "MUST NOT" include a message-body (such as the 1xx, 204, and 304 responses and any response to a HEAD request) is always terminated by the first empty line after the header fields, regardless of the entity-header fields present in the message.
-         2.If a Transfer-Encoding header field (section 14.41) is present and has any value other than "identity", then the transfer-length is defined by use of the "chunked" transfer-coding (section 3.6), unless the message is terminated by closing the connection.
-         3.If a Content-Length header field (section 14.13) is present, its decimal value in OCTETs represents both the entity-length and the transfer-length. The Content-Length header field MUST NOT be sent if these two lengths are different (i.e., if a Transfer-Encoding
-         header field is present). If a message is received with both a Transfer-Encoding header field and a Content-Length header field, the latter MUST be ignored.
-         [LP note: multipart/byteranges support has been removed from the HTTP 1.1 spec by RFC 7230, section A.2. Since it is seldom used, LittleProxy does not check for it.]
-         5.By the server closing the connection. (Closing the connection cannot be used to indicate the end of a request body, since that would leave no possibility for the server to send back a response.)
-     * 
- * - * The rules for Transfer-Encoding are clarified in RFC 7230, section 3.3.1 and 3.3.3 (3): - *
-         If any transfer coding other than
-         chunked is applied to a response payload body, the sender MUST either
-         apply chunked as the final transfer coding or terminate the message
-         by closing the connection.
-     * 
- * - * - * @param response the HTTP response object - * @return true if the message will indicate its own message length, or false if the server is expected to indicate the message length by closing the connection - */ - public static boolean isResponseSelfTerminating(HttpResponse response) { - if (isContentAlwaysEmpty(response)) { - return true; - } - - // if there is a Transfer-Encoding value, determine whether the final encoding is "chunked", which makes the message self-terminating - List allTransferEncodingHeaders = getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, response); - if (!allTransferEncodingHeaders.isEmpty()) { - String finalEncoding = allTransferEncodingHeaders.get(allTransferEncodingHeaders.size() - 1); - - // per #3 above: "If a message is received with both a Transfer-Encoding header field and a Content-Length header field, the latter MUST be ignored." - // since the Transfer-Encoding field is present, the message is self-terminating if and only if the final Transfer-Encoding value is "chunked" - return HttpHeaderValues.CHUNKED.toString().equals(finalEncoding); - } - - String contentLengthHeader = response.headers().get(HttpHeaderNames.CONTENT_LENGTH); - return contentLengthHeader != null && !contentLengthHeader.isEmpty(); - - // not checking for multipart/byteranges, since it is seldom used and its use as a message length indicator was removed in RFC 7230 - - // none of the other message length indicators are present, so the only way the server can indicate the end - // of this message is to close the connection + httpMessage.headers().set(HttpHeaderNames.VIA, vias); + } + + /** + * Returns true if the specified string is either "true" or "on" ignoring case. + * + * @param val The string in question. + * @return true if the specified string is either "true" or "on" ignoring case, + * otherwise false. + */ + public static boolean isTrue(final String val) { + return checkTrueOrFalse(val, "true", "on"); + } + + /** + * Returns true if the specified string is either "false" or "off" ignoring case. + * + * @param val The string in question. + * @return true if the specified string is either "false" or "off" ignoring case, + * otherwise false. + */ + public static boolean isFalse(final String val) { + return checkTrueOrFalse(val, "false", "off"); + } + + public static boolean extractBooleanDefaultFalse(final Properties props, final String key) { + final String throttle = props.getProperty(key); + if (StringUtils.isNotBlank(throttle)) { + return "true".equalsIgnoreCase(throttle.trim()); } + return false; + } - /** - * Retrieves all comma-separated values for headers with the specified name on the HttpMessage. Any whitespace (spaces - * or tabs) surrounding the values will be removed. Empty values (e.g. two consecutive commas, or a value followed - * by a comma and no other value) will be removed; they will not appear as empty elements in the returned list. - * If the message contains repeated headers, their values will be added to the returned list in the order in which - * the headers appear. For example, if a message has headers like: - *
-     *     Transfer-Encoding: gzip,deflate
-     *     Transfer-Encoding: chunked
-     * 
- * This method will return a list of three values: "gzip", "deflate", "chunked". - *

- * Placing values on multiple header lines is allowed under certain circumstances - * in RFC 2616 section 4.2, and in RFC 7230 section 3.2.2 quoted here: - *

-     A sender MUST NOT generate multiple header fields with the same field
-     name in a message unless either the entire field value for that
-     header field is defined as a comma-separated list [i.e., #(values)]
-     or the header field is a well-known exception (as noted below).
-
-     A recipient MAY combine multiple header fields with the same field
-     name into one "field-name: field-value" pair, without changing the
-     semantics of the message, by appending each subsequent field value to
-     the combined field value in order, separated by a comma.  The order
-     in which header fields with the same field name are received is
-     therefore significant to the interpretation of the combined field
-     value; a proxy MUST NOT change the order of these field values when
-     forwarding a message.
-     * 
- * @param headerName the name of the header for which values will be retrieved - * @param httpMessage the HTTP message whose header values will be retrieved - * @return a list of single header values, or an empty list if the header was not present in the message or contained no values - */ - public static List getAllCommaSeparatedHeaderValues(AsciiString headerName, HttpMessage httpMessage) { - List allHeaders = httpMessage.headers().getAll(headerName); - if (allHeaders.isEmpty()) { - return Collections.emptyList(); - } - - ImmutableList.Builder headerValues = ImmutableList.builder(); - for (String header : allHeaders) { - List commaSeparatedValues = splitCommaSeparatedHeaderValues(header); - headerValues.addAll(commaSeparatedValues); - } + public static int extractInt(final Properties props, final String key) { + return extractInt(props, key, -1); + } - return headerValues.build(); + public static int extractInt(final Properties props, final String key, int defaultValue) { + final String readThrottleString = props.getProperty(key); + if (StringUtils.isNotBlank(readThrottleString) && NumberUtils.isCreatable(readThrottleString)) { + return Integer.parseInt(readThrottleString); } + return defaultValue; + } - /** - * Duplicates the status line and headers of an HttpResponse object. Does not duplicate any content associated with that response. - * - * @param originalResponse HttpResponse to be duplicated - * @return a new HttpResponse with the same status line and headers - */ - public static HttpResponse duplicateHttpResponse(HttpResponse originalResponse) { - DefaultHttpResponse newResponse = new DefaultHttpResponse(originalResponse.protocolVersion(), originalResponse.status()); - newResponse.headers().add(originalResponse.headers()); - - return newResponse; - } + public static long extractLong(final Properties props, final String key) { + return extractLong(props, key, -1); + } - /** - * Attempts to resolve the local machine's hostname. - * - * @return the local machine's hostname, or null if a hostname cannot be determined - */ - public static String getHostName() { - try { - return InetAddress.getLocalHost().getHostName(); - } catch (IOException | RuntimeException e) { - LOG.debug("Ignored exception", e); - } // An exception here must not stop the proxy. Android could throw a - // runtime exception, since it not allows network access in the main - // process. - - LOG.info("Could not lookup localhost"); - return null; + public static long extractLong(final Properties props, final String key, long defaultValue) { + final String readThrottleString = props.getProperty(key); + if (StringUtils.isNotBlank(readThrottleString) && NumberUtils.isCreatable(readThrottleString)) { + return Long.parseLong(readThrottleString); } - - /** - * Determines if the specified header should be removed from the proxied response because it is a hop-by-hop header, as defined by the - * HTTP 1.1 spec in section 13.5.1. The comparison is case-insensitive, so "Connection" will be treated the same as "connection" or "CONNECTION". - * From http://www.w3.org/Protocols/rfc2616/rfc2616-sec13.html#sec13.5.1 : - *
-       The following HTTP/1.1 headers are hop-by-hop headers:
-        - Connection
-        - Keep-Alive
-        - Proxy-Authenticate
-        - Proxy-Authorization
-        - TE
-        - Trailers [LittleProxy note: actual header name is Trailer]
-        - Transfer-Encoding [LittleProxy note: this header is not normally removed when proxying, since the proxy does not re-chunk
-                            responses. The exception is when an HttpObjectAggregator is enabled, which aggregates chunked content and removes
-                            the 'Transfer-Encoding: chunked' header itself.]
-        - Upgrade
-
-       All other headers defined by HTTP/1.1 are end-to-end headers.
-     * 
- * - * @param headerName the header name - * @return true if this header is a hop-by-hop header and should be removed when proxying, otherwise false - */ - public static boolean shouldRemoveHopByHopHeader(String headerName) { - return SHOULD_NOT_PROXY_HOP_BY_HOP_HEADERS.contains(headerName); + return defaultValue; + } + + public static boolean isCONNECT(HttpObject httpObject) { + return httpObject instanceof HttpRequest + && HttpMethod.CONNECT.equals(((HttpRequest) httpObject).method()); + } + + /** + * Returns true if the specified HttpRequest is a HEAD request. + * + * @param httpRequest http request + * @return true if request is a HEAD, otherwise false + */ + public static boolean isHEAD(HttpRequest httpRequest) { + return httpRequest != null && HttpMethod.HEAD.equals(httpRequest.method()); + } + + private static boolean checkTrueOrFalse(final String val, final String str1, final String str2) { + final String str = val.trim(); + return StringUtils.isNotBlank(str) + && (str.equalsIgnoreCase(str1) || str.equalsIgnoreCase(str2)); + } + + /** + * Returns true if the HTTP message cannot contain an entity body, according to the HTTP spec. + * This code is taken directly from {@link + * io.netty.handler.codec.http.HttpObjectDecoder#isContentAlwaysEmpty(HttpMessage)}. + * + * @param msg HTTP message + * @return true if the HTTP message is always empty, false if the message may have entity + * content. + */ + public static boolean isContentAlwaysEmpty(HttpMessage msg) { + if (msg instanceof HttpResponse) { + HttpResponse res = (HttpResponse) msg; + int code = res.status().code(); + + // Correctly handle return codes of 1xx. + // + // See: + // - https://www.w3.org/Protocols/rfc2616/rfc2616-sec4.html Section 4.4 + // - https://github.com/netty/netty/issues/222 + if (code >= 100 && code < 200) { + // According to RFC 7231, section 6.1, 1xx responses have no content + // (https://tools.ietf.org/html/rfc7231#section-6.2): + // 1xx responses are terminated by the first empty line after + // the status-line (the empty line signaling the end of the header + // section). + + // Hixie 76 websocket handshake responses contain a 16-byte body, so their content is not + // empty; but Hixie 76 + // was a draft specification that was superceded by RFC 6455. Since it is rarely used and + // doesn't conform to + // RFC 7231, we do not support or make special allowance for Hixie 76 responses. + return true; + } + + switch (code) { + case 204: + case 205: + case 304: + return true; + } } - - /** - * Splits comma-separated header values into tokens. For example, if the value of the Connection header is "Transfer-Encoding, close", - * this method will return "Transfer-Encoding" and "close". This method strips trims any optional whitespace from - * the tokens. Unlike {@link #getAllCommaSeparatedHeaderValues(AsciiString, HttpMessage)}, this method only operates on - * a single header value, rather than all instances of the header in a message. - * - * @param headerValue the un-tokenized header value (must not be null) - * @return all tokens within the header value, or an empty list if there are no values - */ - public static List splitCommaSeparatedHeaderValues(String headerValue) { - return ImmutableList.copyOf(COMMA_SEPARATED_HEADER_VALUE_SPLITTER.split(headerValue)); + return false; + } + + /** + * Returns true if the HTTP response from the server is expected to indicate its own message + * length/end-of-message. Returns false if the server is expected to indicate the end of the HTTP + * entity by closing the connection. + * + *

This method is based on the allowed message length indicators in the HTTP specification, + * section 4.4: + * + *

+   * 4.4 Message Length
+   * The transfer-length of a message is the length of the message-body as it appears in the message; that is, after any transfer-codings have been applied. When a message-body is included with a message, the transfer-length of that body is determined by one of the following (in order of precedence):
+   *
+   * 1.Any response message which "MUST NOT" include a message-body (such as the 1xx, 204, and 304 responses and any response to a HEAD request) is always terminated by the first empty line after the header fields, regardless of the entity-header fields present in the message.
+   * 2.If a Transfer-Encoding header field (section 14.41) is present and has any value other than "identity", then the transfer-length is defined by use of the "chunked" transfer-coding (section 3.6), unless the message is terminated by closing the connection.
+   * 3.If a Content-Length header field (section 14.13) is present, its decimal value in OCTETs represents both the entity-length and the transfer-length. The Content-Length header field MUST NOT be sent if these two lengths are different (i.e., if a Transfer-Encoding
+   * header field is present). If a message is received with both a Transfer-Encoding header field and a Content-Length header field, the latter MUST be ignored.
+   * [LP note: multipart/byteranges support has been removed from the HTTP 1.1 spec by RFC 7230, section A.2. Since it is seldom used, LittleProxy does not check for it.]
+   * 5.By the server closing the connection. (Closing the connection cannot be used to indicate the end of a request body, since that would leave no possibility for the server to send back a response.)
+   * 
+ * + * The rules for Transfer-Encoding are clarified in RFC 7230, section 3.3.1 and 3.3.3 (3): + * + *
+   * If any transfer coding other than
+   * chunked is applied to a response payload body, the sender MUST either
+   * apply chunked as the final transfer coding or terminate the message
+   * by closing the connection.
+   * 
+ * + * @param response the HTTP response object + * @return true if the message will indicate its own message length, or false if the server is + * expected to indicate the message length by closing the connection + */ + public static boolean isResponseSelfTerminating(HttpResponse response) { + if (isContentAlwaysEmpty(response)) { + return true; } - /** - * Determines if UDT is available on the classpath. - * - * @return true if UDT is available - */ - public static boolean isUdtAvailable() { - try { - return NioUdtProvider.BYTE_PROVIDER != null; - } catch (NoClassDefFoundError e) { - return false; - } + // if there is a Transfer-Encoding value, determine whether the final encoding is "chunked", + // which makes the message self-terminating + List allTransferEncodingHeaders = + getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, response); + if (!allTransferEncodingHeaders.isEmpty()) { + String finalEncoding = allTransferEncodingHeaders.get(allTransferEncodingHeaders.size() - 1); + + // per #3 above: "If a message is received with both a Transfer-Encoding header field and a + // Content-Length header field, the latter MUST be ignored." + // since the Transfer-Encoding field is present, the message is self-terminating if and only + // if the final Transfer-Encoding value is "chunked" + return HttpHeaderValues.CHUNKED.toString().equals(finalEncoding); } - /** - * Creates a new {@link FullHttpResponse} with the specified String as the body contents (encoded using UTF-8). - * - * @param httpVersion HTTP version of the response - * @param status HTTP status code - * @param body body to include in the FullHttpResponse; will be UTF-8 encoded - * @return new http response object - */ - public static FullHttpResponse createFullHttpResponse(HttpVersion httpVersion, - HttpResponseStatus status, - String body) { - byte[] bytes = body.getBytes(StandardCharsets.UTF_8); - ByteBuf content = Unpooled.copiedBuffer(bytes); - - return createFullHttpResponse(httpVersion, status, "text/html; charset=utf-8", content, bytes.length); + String contentLengthHeader = response.headers().get(HttpHeaderNames.CONTENT_LENGTH); + return contentLengthHeader != null && !contentLengthHeader.isEmpty(); + + // not checking for multipart/byteranges, since it is seldom used and its use as a message + // length indicator was removed in RFC 7230 + + // none of the other message length indicators are present, so the only way the server can + // indicate the end + // of this message is to close the connection + } + + /** + * Retrieves all comma-separated values for headers with the specified name on the HttpMessage. + * Any whitespace (spaces or tabs) surrounding the values will be removed. Empty values (e.g. two + * consecutive commas, or a value followed by a comma and no other value) will be removed; they + * will not appear as empty elements in the returned list. If the message contains repeated + * headers, their values will be added to the returned list in the order in which the headers + * appear. For example, if a message has headers like: + * + *
+   *     Transfer-Encoding: gzip,deflate
+   *     Transfer-Encoding: chunked
+   * 
+ * + * This method will return a list of three values: "gzip", "deflate", "chunked". + * + *

Placing values on multiple header lines is allowed under certain circumstances in RFC 2616 + * section 4.2, and in RFC 7230 section 3.2.2 quoted here: + * + *

+   * A sender MUST NOT generate multiple header fields with the same field
+   * name in a message unless either the entire field value for that
+   * header field is defined as a comma-separated list [i.e., #(values)]
+   * or the header field is a well-known exception (as noted below).
+   *
+   * A recipient MAY combine multiple header fields with the same field
+   * name into one "field-name: field-value" pair, without changing the
+   * semantics of the message, by appending each subsequent field value to
+   * the combined field value in order, separated by a comma.  The order
+   * in which header fields with the same field name are received is
+   * therefore significant to the interpretation of the combined field
+   * value; a proxy MUST NOT change the order of these field values when
+   * forwarding a message.
+   * 
+ * + * @param headerName the name of the header for which values will be retrieved + * @param httpMessage the HTTP message whose header values will be retrieved + * @return a list of single header values, or an empty list if the header was not present in the + * message or contained no values + */ + public static List getAllCommaSeparatedHeaderValues( + AsciiString headerName, HttpMessage httpMessage) { + List allHeaders = httpMessage.headers().getAll(headerName); + if (allHeaders.isEmpty()) { + return Collections.emptyList(); } - /** - * Creates a new {@link FullHttpResponse} with no body content - * - * @param httpVersion HTTP version of the response - * @param status HTTP status code - * @return new http response object - */ - public static FullHttpResponse createFullHttpResponse(HttpVersion httpVersion, - HttpResponseStatus status) { - return createFullHttpResponse(httpVersion, status, null, null, 0); + return allHeaders.stream() + .map(ProxyUtils::splitCommaSeparatedHeaderValues) + .flatMap(List::stream) + .collect(toList()); + } + + /** + * Duplicates the status line and headers of an HttpResponse object. Does not duplicate any + * content associated with that response. + * + * @param originalResponse HttpResponse to be duplicated + * @return a new HttpResponse with the same status line and headers + */ + public static HttpResponse duplicateHttpResponse(HttpResponse originalResponse) { + DefaultHttpResponse newResponse = + new DefaultHttpResponse(originalResponse.protocolVersion(), originalResponse.status()); + newResponse.headers().add(originalResponse.headers()); + + return newResponse; + } + + /** + * Attempts to resolve the local machine's hostname. + * + * @return the local machine's hostname, or null if a hostname cannot be determined + */ + @Nullable + public static String getHostName() { + return Hostname.getHostName(); + } + + /** + * Determines if the specified header should be removed from the proxied response because it is a + * hop-by-hop header, as defined by the HTTP 1.1 spec in section 13.5.1. The comparison is + * case-insensitive, so "Connection" will be treated the same as "connection" or "CONNECTION". + * From section + * 13.5.1 : + * + *
+   * The following HTTP/1.1 headers are hop-by-hop headers:
+   * - Connection
+   * - Keep-Alive
+   * - Proxy-Authenticate
+   * - Proxy-Authorization
+   * - TE
+   * - Trailers [LittleProxy note: actual header name is Trailer]
+   * - Transfer-Encoding [LittleProxy note: this header is not normally removed when proxying, since the proxy does not re-chunk
+   * responses. The exception is when an HttpObjectAggregator is enabled, which aggregates chunked content and removes
+   * the 'Transfer-Encoding: chunked' header itself.]
+   * - Upgrade
+   *
+   * All other headers defined by HTTP/1.1 are end-to-end headers.
+   * 
+ * + * @param headerName the header name + * @return true if this header is a hop-by-hop header and should be removed when proxying, + * otherwise false + */ + public static boolean shouldRemoveHopByHopHeader(String headerName) { + return SHOULD_NOT_PROXY_HOP_BY_HOP_HEADERS.contains(headerName); + } + + /** + * Removes all headers that should not be forwarded. See RFC 2616 13.5.1 End-to-end and Hop-by-hop + * Headers. + * + * @param headers The headers to modify + */ + public static void stripHopByHopHeaders(HttpHeaders headers) { + // Not explicitly documented, but remove is case-insensitive as HTTP header handling function + // should be + for (String headerName : SHOULD_NOT_PROXY_HOP_BY_HOP_HEADERS) { + headers.remove(headerName); } - - /** - * Creates a new {@link FullHttpResponse} with the specified body. - * - * @param httpVersion HTTP version of the response - * @param status HTTP status code - * @param contentType the Content-Type of the body - * @param body body to include in the FullHttpResponse; if null - * @param contentLength number of bytes to send in the Content-Length header; should equal the number of bytes in the ByteBuf - * @return new http response object - */ - public static FullHttpResponse createFullHttpResponse(HttpVersion httpVersion, - HttpResponseStatus status, - String contentType, - ByteBuf body, - int contentLength) { - DefaultFullHttpResponse response; - - if (body != null) { - response = new DefaultFullHttpResponse(httpVersion, status, body); - response.headers().set(HttpHeaderNames.CONTENT_LENGTH, contentLength); - response.headers().set(HttpHeaderNames.CONTENT_TYPE, contentType); - } else { - response = new DefaultFullHttpResponse(httpVersion, status); - } - - return response; + } + + /** + * Splits comma-separated header values into tokens. + * + *

For example, if the value of the Connection header is "Transfer-Encoding, close", this + * method will return "Transfer-Encoding" and "close". + * + *

This method strips trims any optional whitespace from the tokens. Unlike {@link + * #getAllCommaSeparatedHeaderValues(AsciiString, HttpMessage)}, this method only operates on a + * single header value, rather than all instances of the header in a message. + * + * @param headerValue the un-tokenized header value (must not be null) + * @return all tokens within the header value, or an empty list if there are no values + */ + public static List splitCommaSeparatedHeaderValues(String headerValue) { + return Stream.of(headerValue.split(",")) + .map(String::trim) + .filter(s -> !s.isEmpty()) + .collect(toList()); + } + + /** + * Creates a new {@link FullHttpResponse} with the specified String as the body contents (encoded + * using UTF-8). + * + * @param httpVersion HTTP version of the response + * @param status HTTP status code + * @param body body to include in the FullHttpResponse; will be UTF-8 encoded + * @return new http response object + */ + public static FullHttpResponse createFullHttpResponse( + HttpVersion httpVersion, HttpResponseStatus status, String body) { + byte[] bytes = body.getBytes(StandardCharsets.UTF_8); + ByteBuf content = Unpooled.copiedBuffer(bytes); + + return createFullHttpResponse( + httpVersion, status, "text/html; charset=utf-8", content, bytes.length); + } + + /** + * Creates a new {@link FullHttpResponse} with no content + * + * @param httpVersion HTTP version of the response + * @param status HTTP status code + * @return new http response object + */ + public static FullHttpResponse createFullHttpResponse( + HttpVersion httpVersion, HttpResponseStatus status) { + return createFullHttpResponse(httpVersion, status, null, null, 0); + } + + /** + * Creates a new {@link FullHttpResponse} with the specified body. + * + * @param httpVersion HTTP version of the response + * @param status HTTP status code + * @param contentType the Content-Type of the body + * @param body body to include in the FullHttpResponse; if null + * @param contentLength number of bytes to send in the Content-Length header; should equal the + * number of bytes in the ByteBuf + * @return new http response object + */ + public static FullHttpResponse createFullHttpResponse( + HttpVersion httpVersion, + HttpResponseStatus status, + String contentType, + ByteBuf body, + int contentLength) { + DefaultFullHttpResponse response; + + if (body != null) { + response = new DefaultFullHttpResponse(httpVersion, status, body); + response.headers().set(HttpHeaderNames.CONTENT_LENGTH, contentLength); + response.headers().set(HttpHeaderNames.CONTENT_TYPE, contentType); + } else { + response = new DefaultFullHttpResponse(httpVersion, status); } - /** - * Given an HttpHeaders instance, removes 'sdch' from the 'Accept-Encoding' - * header list (if it exists) and returns the modified instance. - * - * Removes all occurrences of 'sdch' from the 'Accept-Encoding' header. - * @param headers The headers to modify. - */ - public static void removeSdchEncoding(HttpHeaders headers) { - List encodings = headers.getAll(HttpHeaderNames.ACCEPT_ENCODING); - headers.remove(HttpHeaderNames.ACCEPT_ENCODING); - - for (String encoding : encodings) { - if (encoding != null) { - // The former regex should remove occurrences of 'sdch' while the - // latter regex should take care of the dangling comma case when - // 'sdch' was the first element in the list and there are other - // encodings. - encoding = encoding.replaceAll(",? *(sdch|SDCH)", "").replaceFirst("^ *, *", ""); - - if (StringUtils.isNotBlank(encoding)) { - headers.add(HttpHeaderNames.ACCEPT_ENCODING, encoding); - } - } + return response; + } + + /** + * Given an HttpHeaders instance, removes 'sdch' from the 'Accept-Encoding' header list (if it + * exists) and returns the modified instance. + * + *

Removes all occurrences of 'sdch' from the 'Accept-Encoding' header. + * + * @param headers The headers to modify. + */ + public static void removeSdchEncoding(HttpHeaders headers) { + List encodings = headers.getAll(HttpHeaderNames.ACCEPT_ENCODING); + headers.remove(HttpHeaderNames.ACCEPT_ENCODING); + + for (String encoding : encodings) { + if (encoding != null) { + // The former regex should remove occurrences of 'sdch' while the + // latter regex should take care of the dangling comma case when + // 'sdch' was the first element in the list and there are other + // encodings. + encoding = encoding.replaceAll(",? *(sdch|SDCH)", "").replaceFirst("^ *, *", ""); + + if (StringUtils.isNotBlank(encoding)) { + headers.add(HttpHeaderNames.ACCEPT_ENCODING, encoding); } + } } + } + + /** + * Tests whether the given response indicates that the connection is switching to the WebSocket + * protocol. + * + * @param response the response to check. + * @return true if switching to the WebSocket protocol; false otherwise; + */ + public static boolean isSwitchingToWebSocketProtocol(HttpResponse response) { + return (response.status() == HttpResponseStatus.SWITCHING_PROTOCOLS) + && response.headers().contains(HttpHeaderNames.CONNECTION, HttpHeaderNames.UPGRADE, true) + && response.headers().contains(HttpHeaderNames.UPGRADE, "websocket", true); + } + + /** + * Tests whether the given request indicates that the connection is switching to the WebSocket + * protocol. + * + * @param request the request to check. + * @return true if switching to the WebSocket protocol; false otherwise; + */ + public static boolean isSwitchingToWebSocketProtocol(HttpRequest request) { + return request.headers().contains(HttpHeaderNames.CONNECTION, HttpHeaderNames.UPGRADE, true) + && request.headers().contains(HttpHeaderNames.UPGRADE, "websocket", true); + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/ServerConnectionPool.java b/src/main/java/org/littleshoot/proxy/impl/ServerConnectionPool.java new file mode 100644 index 00000000..c91abb09 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/ServerConnectionPool.java @@ -0,0 +1,147 @@ +package org.littleshoot.proxy.impl; + +import io.netty.channel.Channel; +import io.netty.handler.codec.http.HttpRequest; +import java.net.InetSocketAddress; +import java.time.Duration; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.HttpFilters; + +/** + * Interface for pooling ProxyToServerConnection instances. + * + *

This interface allows swapping different pooling implementations: + * + *

    + *
  • {@link ConcurrentMapServerConnectionPool} - Simple ConcurrentHashMap-based pool + *
+ */ +public interface ServerConnectionPool { + + /** + * Gets a connection for the given host and port, or creates one if it doesn't exist. + * + * @param serverHostAndPort the server host and port key + * @param chainedProxyAddress the address of the resolved chained proxy, used to segregate + * connections by upstream route (null for direct connections) + * @param clientConnection the client connection that needs the server connection + * @param initialFilters the initial HTTP filters + * @param initialHttpRequest the initial HTTP request + * @return the ProxyToServerConnection, or null if creation failed or pool is exhausted + */ + @Nullable ProxyToServerConnection getOrCreateConnection( + String serverHostAndPort, + @Nullable InetSocketAddress chainedProxyAddress, + ClientToProxyConnection clientConnection, + HttpFilters initialFilters, + HttpRequest initialHttpRequest); + + /** + * Releases a connection back to the pool after a request is complete. + * + * @param connection the connection to release + */ + void releaseConnection(ProxyToServerConnection connection); + + /** + * Registers a pending request for HTTP pipelining support. + * + * @param channel the server channel + * @param clientConnection the client connection that made the request + * @param request the HTTP request + * @param filters the filters active for this request + */ + void registerPendingRequest( + Channel channel, + ClientToProxyConnection clientConnection, + HttpRequest request, + HttpFilters filters); + + /** + * Gets and removes the oldest pending request for the given channel (FIFO order for pipelining). + * + * @param channel the server channel + * @return the oldest pending request, or null if none found + */ + @Nullable PendingRequest removePendingRequest(Channel channel); + + /** + * Gets the oldest pending request without removing it. + * + * @param channel the server channel + * @return the oldest pending request, or null if none found + */ + @Nullable PendingRequest peekPendingRequest(Channel channel); + + /** + * Drains and removes all pending requests for the given channel. + * + * @param channel the server channel + */ + void drainPendingRequests(Channel channel); + + /** + * Removes a connection from the pool when it's disconnected. + * + * @param connection the connection being removed + */ + void removeConnection(ProxyToServerConnection connection); + + /** Closes all connections in the pool. */ + void closeAll(); + + /** Returns the maximum number of connections allowed per host. */ + int getMaxConnectionsPerHost(); + + /** Returns the maximum total number of connections allowed in the pool. */ + int getMaxConnections(); + + /** + * Sets the idle timeout for connections. Connections idle for longer than this duration will be + * evicted. + * + * @param idleTimeout the idle timeout duration, or null to disable + */ + void setIdleTimeout(@Nullable Duration idleTimeout); + + /** Returns the configured idle timeout, or null if not set. */ + @Nullable Duration getIdleTimeout(); + + /** + * Enables or disables connection validation. When enabled, connections are validated before being + * borrowed from the pool to ensure they are still functional. + * + * @param validationEnabled true to enable validation, false to disable + */ + void setConnectionValidationEnabled(boolean validationEnabled); + + /** Returns true if connection validation is enabled. */ + boolean isConnectionValidationEnabled(); + + /** + * Returns current pool metrics. + * + * @return PoolMetrics with current statistics + */ + PoolMetrics getMetrics(); + + default String computePoolKey( + String serverHostAndPort, @Nullable InetSocketAddress chainedProxyAddress) { + if (chainedProxyAddress == null) { + return serverHostAndPort + ":direct"; + } + if (chainedProxyAddress.getAddress() == null) { + // Unresolved address - use hostname string instead + return serverHostAndPort + + ":" + + chainedProxyAddress.getHostString() + + ":" + + chainedProxyAddress.getPort(); + } + return serverHostAndPort + + ":" + + chainedProxyAddress.getAddress().getHostAddress() + + ":" + + chainedProxyAddress.getPort(); + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/ServerConnectionPoolConfig.java b/src/main/java/org/littleshoot/proxy/impl/ServerConnectionPoolConfig.java new file mode 100644 index 00000000..1eac78aa --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/ServerConnectionPoolConfig.java @@ -0,0 +1,91 @@ +package org.littleshoot.proxy.impl; + +import static java.util.Objects.requireNonNull; + +import java.time.Duration; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.ServerConnectionPoolType; + +/** Configuration for the server connection pool. */ +public class ServerConnectionPoolConfig { + private boolean enabled = false; + private ServerConnectionPoolType poolType = ServerConnectionPoolType.CONCURRENT_MAP; + private int maxConnectionsPerHost = + ConcurrentMapServerConnectionPool.DEFAULT_MAX_CONNECTIONS_PER_HOST; + private int maxConnections = ConcurrentMapServerConnectionPool.DEFAULT_MAX_TOTAL_CONNECTIONS; + @Nullable private Duration idleTimeout; + private boolean poolSharedMitmConnections = false; + private boolean poolPerRequestInMitm = false; + + public boolean isEnabled() { + return enabled; + } + + public ServerConnectionPoolConfig setEnabled(boolean enabled) { + this.enabled = enabled; + return this; + } + + public ServerConnectionPoolType getPoolType() { + return poolType; + } + + public ServerConnectionPoolConfig setPoolType(ServerConnectionPoolType poolType) { + this.poolType = requireNonNull(poolType, "poolType must not be null"); + return this; + } + + public int getMaxConnectionsPerHost() { + return maxConnectionsPerHost; + } + + public ServerConnectionPoolConfig setMaxConnectionsPerHost(int maxConnectionsPerHost) { + if (maxConnectionsPerHost <= 0) { + throw new IllegalArgumentException( + "maxConnectionsPerHost must be positive: " + maxConnectionsPerHost); + } + this.maxConnectionsPerHost = maxConnectionsPerHost; + return this; + } + + public int getMaxConnections() { + return maxConnections; + } + + public ServerConnectionPoolConfig setMaxConnections(int maxConnections) { + if (maxConnections <= 0) { + throw new IllegalArgumentException("maxConnections must be positive: " + maxConnections); + } + this.maxConnections = maxConnections; + return this; + } + + @Nullable + public Duration getIdleTimeout() { + return idleTimeout; + } + + public ServerConnectionPoolConfig setIdleTimeout(@Nullable Duration idleTimeout) { + this.idleTimeout = idleTimeout; + return this; + } + + public boolean isPoolSharedMitmConnections() { + return poolSharedMitmConnections; + } + + public ServerConnectionPoolConfig setPoolSharedMitmConnections( + boolean poolSharedMitmConnections) { + this.poolSharedMitmConnections = poolSharedMitmConnections; + return this; + } + + public boolean isPoolPerRequestInMitm() { + return poolPerRequestInMitm; + } + + public ServerConnectionPoolConfig setPoolPerRequestInMitm(boolean poolPerRequestInMitm) { + this.poolPerRequestInMitm = poolPerRequestInMitm; + return this; + } +} diff --git a/src/main/java/org/littleshoot/proxy/impl/ServerGroup.java b/src/main/java/org/littleshoot/proxy/impl/ServerGroup.java index a9e6ea6d..0a723b4e 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ServerGroup.java +++ b/src/main/java/org/littleshoot/proxy/impl/ServerGroup.java @@ -1,291 +1,316 @@ package org.littleshoot.proxy.impl; import io.netty.channel.EventLoopGroup; -import io.netty.channel.udt.nio.NioUdtProvider; -import org.littleshoot.proxy.HttpProxyServer; -import org.littleshoot.proxy.TransportProtocol; -import org.littleshoot.proxy.UnknownTransportProtocolException; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - import java.nio.channels.spi.SelectorProvider; import java.util.ArrayList; import java.util.EnumMap; +import java.util.HashSet; import java.util.List; +import java.util.Set; import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; +import org.littleshoot.proxy.HttpProxyServer; +import org.littleshoot.proxy.TransportProtocol; +import org.littleshoot.proxy.UnknownTransportProtocolException; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; /** - * Manages thread pools for one or more proxy server instances. When servers are created, they must register with the - * ServerGroup using {@link #registerProxyServer(HttpProxyServer)}, and when they shut down, must unregister with the - * ServerGroup using {@link #unregisterProxyServer(HttpProxyServer, boolean)}. + * Manages thread pools for one or more proxy server instances. When servers are created, they must + * register with the ServerGroup using {@link #registerProxyServer(HttpProxyServer)}, and when they + * shut down, must unregister with the ServerGroup using {@link + * #unregisterProxyServer(HttpProxyServer, boolean)}. */ public class ServerGroup { - private static final Logger log = LoggerFactory.getLogger(ServerGroup.class); - - /** - * The default number of threads to accept incoming requests from clients. (Requests are serviced by worker threads, - * not acceptor threads.) - */ - public static final int DEFAULT_INCOMING_ACCEPTOR_THREADS = 2; - - /** - * The default number of threads to service incoming requests from clients. - */ - public static final int DEFAULT_INCOMING_WORKER_THREADS = 8; - - /** - * The default number of threads to service outgoing requests to servers. - */ - public static final int DEFAULT_OUTGOING_WORKER_THREADS = 8; - - /** - * Global counter for the {@link #serverGroupId}. - */ - private static final AtomicInteger serverGroupCount = new AtomicInteger(0); - - /** - * A name for this ServerGroup to use in naming threads. - */ - private final String name; - - /** - * The ID of this server group. Forms part of the name of each thread created for this server group. Useful for - * differentiating threads when multiple proxy instances are running. - */ - private final int serverGroupId; - - private final int incomingAcceptorThreads; - private final int incomingWorkerThreads; - private final int outgoingWorkerThreads; - - /** - * List of all servers registered to use this ServerGroup. Any access to this list should be synchronized using the - * {@link #SERVER_REGISTRATION_LOCK}. - */ - public final List registeredServers = new ArrayList<>(1); - - /** - * A mapping of {@link TransportProtocol}s to their initialized {@link ProxyThreadPools}. Each transport uses a - * different thread pool, since the initialization parameters are different. - */ - private final EnumMap protocolThreadPools = new EnumMap<>(TransportProtocol.class); - - /** - * A mapping of selector providers to transport protocols. Avoids special-casing each transport protocol during - * transport protocol initialization. - */ - private static final EnumMap TRANSPORT_PROTOCOL_SELECTOR_PROVIDERS = new EnumMap<>(TransportProtocol.class); - static { - TRANSPORT_PROTOCOL_SELECTOR_PROVIDERS.put(TransportProtocol.TCP, SelectorProvider.provider()); - - // allow the proxy to operate without UDT support. this allows clients that do not use UDT to exclude the barchart - // dependency completely. - if (ProxyUtils.isUdtAvailable()) { - TRANSPORT_PROTOCOL_SELECTOR_PROVIDERS.put(TransportProtocol.UDT, NioUdtProvider.BYTE_PROVIDER); - } else { - log.debug("UDT provider not found on classpath. UDT transport will not be available."); - } - } - - /** - * True when this ServerGroup is stopped. - */ - private final AtomicBoolean stopped = new AtomicBoolean(false); - - /** - * Creates a new ServerGroup instance for a proxy. Threads created for this ServerGroup will have the specified - * ServerGroup name in the Thread name. This constructor does not actually initialize any thread pools; instead, - * thread pools for specific transport protocols are lazily initialized as needed. - * - * @param name ServerGroup name to include in thread names - * @param incomingAcceptorThreads number of acceptor threads per protocol - * @param incomingWorkerThreads number of client-to-proxy worker threads per protocol - * @param outgoingWorkerThreads number of proxy-to-server worker threads per protocol - */ - public ServerGroup(String name, int incomingAcceptorThreads, int incomingWorkerThreads, int outgoingWorkerThreads) { - this.name = name; - this.serverGroupId = serverGroupCount.getAndIncrement(); - this.incomingAcceptorThreads = incomingAcceptorThreads; - this.incomingWorkerThreads = incomingWorkerThreads; - this.outgoingWorkerThreads = outgoingWorkerThreads; - } - - /** - * Lock for initializing any transport protocols. - */ - private final Object THREAD_POOL_INIT_LOCK = new Object(); - - /** - * Retrieves the {@link ProxyThreadPools} for the specified transport protocol. Lazily initializes the thread pools - * for the transport protocol if they have not yet been initialized. If the protocol has already been initialized, - * this method returns immediately, without synchronization. If initialization is necessary, the initialization - * process creates the acceptor and worker threads necessary to service requests to/from the proxy. - *

- * This method is thread-safe; no external locking is necessary. - * - * @param protocol transport protocol to retrieve thread pools for - * @return thread pools for the specified transport protocol - */ - private ProxyThreadPools getThreadPoolsForProtocol(TransportProtocol protocol) { - // if the thread pools have not been initialized for this protocol, initialize them + private static final Logger log = LoggerFactory.getLogger(ServerGroup.class); + + /** + * The default number of threads to accept incoming requests from clients. (Requests are serviced + * by worker threads, not acceptor threads.) + */ + public static final int DEFAULT_INCOMING_ACCEPTOR_THREADS = 2; + + /** The default number of threads to service incoming requests from clients. */ + public static final int DEFAULT_INCOMING_WORKER_THREADS = 8; + + /** The default number of threads to service outgoing requests to servers. */ + public static final int DEFAULT_OUTGOING_WORKER_THREADS = 8; + + /** Global counter for the {@link #serverGroupId}. */ + private static final AtomicInteger serverGroupCount = new AtomicInteger(0); + + /** A name for this ServerGroup to use in naming threads. */ + private final String name; + + /** + * The ID of this server group. Forms part of the name of each thread created for this server + * group. Useful for differentiating threads when multiple proxy instances are running. + */ + private final int serverGroupId; + + private final int incomingAcceptorThreads; + private final int incomingWorkerThreads; + private final int outgoingWorkerThreads; + private final boolean autoStop; + + /** + * List of all servers registered to use this ServerGroup. Any access to this list should be + * synchronized using the {@link #SERVER_REGISTRATION_LOCK}. + */ + public final Set registeredServers = new HashSet<>(1); + + /** + * A mapping of {@link TransportProtocol}s to their initialized {@link ProxyThreadPools}. Each + * transport uses a different thread pool, since the initialization parameters are different. + */ + private final EnumMap protocolThreadPools = + new EnumMap<>(TransportProtocol.class); + + /** + * A mapping of selector providers to transport protocols. Avoids special-casing each transport + * protocol during transport protocol initialization. + */ + private static final EnumMap + TRANSPORT_PROTOCOL_SELECTOR_PROVIDERS = new EnumMap<>(TransportProtocol.class); + + static { + TRANSPORT_PROTOCOL_SELECTOR_PROVIDERS.put(TransportProtocol.TCP, SelectorProvider.provider()); + } + + /** True when this ServerGroup is stopped. */ + private final AtomicBoolean stopped = new AtomicBoolean(false); + + /** + * Creates a new ServerGroup instance for a proxy. Threads created for this ServerGroup will have + * the specified ServerGroup name in the Thread name. This constructor does not actually + * initialize any thread pools; instead, thread pools for specific transport protocols are lazily + * initialized as needed. + * + * @param name ServerGroup name to include in thread names + * @param incomingAcceptorThreads number of acceptor threads per protocol + * @param incomingWorkerThreads number of client-to-proxy worker threads per protocol + * @param outgoingWorkerThreads number of proxy-to-server worker threads per protocol + */ + public ServerGroup( + String name, + int incomingAcceptorThreads, + int incomingWorkerThreads, + int outgoingWorkerThreads) { + this(name, incomingAcceptorThreads, incomingWorkerThreads, outgoingWorkerThreads, true); + } + + /** + * Creates a new ServerGroup instance for a proxy. Threads created for this ServerGroup will have + * the specified ServerGroup name in the Thread name. This constructor does not actually + * initialize any thread pools; instead, thread pools for specific transport protocols are lazily + * initialized as needed. + * + * @param name ServerGroup name to include in thread names + * @param incomingAcceptorThreads number of acceptor threads per protocol + * @param incomingWorkerThreads number of client-to-proxy worker threads per protocol + * @param outgoingWorkerThreads number of proxy-to-server worker threads per protocol + * @param autoStop if this group should stop after removal of the last proxy server + */ + public ServerGroup( + String name, + int incomingAcceptorThreads, + int incomingWorkerThreads, + int outgoingWorkerThreads, + boolean autoStop) { + this.name = name; + this.serverGroupId = serverGroupCount.getAndIncrement(); + this.incomingAcceptorThreads = incomingAcceptorThreads; + this.incomingWorkerThreads = incomingWorkerThreads; + this.outgoingWorkerThreads = outgoingWorkerThreads; + this.autoStop = autoStop; + } + + /** Lock for initializing any transport protocols. */ + private final Object THREAD_POOL_INIT_LOCK = new Object(); + + /** + * Retrieves the {@link ProxyThreadPools} for the specified transport protocol. Lazily initializes + * the thread pools for the transport protocol if they have not yet been initialized. If the + * protocol has already been initialized, this method returns immediately, without + * synchronization. If initialization is necessary, the initialization process creates the + * acceptor and worker threads necessary to service requests to/from the proxy. + * + *

This method is thread-safe; no external locking is necessary. + * + * @param protocol transport protocol to retrieve thread pools for + * @return thread pools for the specified transport protocol + */ + private ProxyThreadPools getThreadPoolsForProtocol(TransportProtocol protocol) { + // if the thread pools have not been initialized for this protocol, initialize them + if (protocolThreadPools.get(protocol) == null) { + synchronized (THREAD_POOL_INIT_LOCK) { if (protocolThreadPools.get(protocol) == null) { - synchronized (THREAD_POOL_INIT_LOCK) { - if (protocolThreadPools.get(protocol) == null) { - log.debug("Initializing thread pools for {} with {} acceptor threads, {} incoming worker threads, and {} outgoing worker threads", - protocol, incomingAcceptorThreads, incomingWorkerThreads, outgoingWorkerThreads); - - SelectorProvider selectorProvider = TRANSPORT_PROTOCOL_SELECTOR_PROVIDERS.get(protocol); - if (selectorProvider == null) { - throw new UnknownTransportProtocolException(protocol); - } - - ProxyThreadPools threadPools = new ProxyThreadPools(selectorProvider, - incomingAcceptorThreads, - incomingWorkerThreads, - outgoingWorkerThreads, - name, - serverGroupId); - protocolThreadPools.put(protocol, threadPools); - } - } + log.debug( + "Initializing thread pools for {} with {} acceptor threads, {} incoming worker threads, and {} outgoing worker threads", + protocol, + incomingAcceptorThreads, + incomingWorkerThreads, + outgoingWorkerThreads); + + SelectorProvider selectorProvider = TRANSPORT_PROTOCOL_SELECTOR_PROVIDERS.get(protocol); + if (selectorProvider == null) { + throw new UnknownTransportProtocolException(protocol); + } + + ProxyThreadPools threadPools = + new ProxyThreadPools( + selectorProvider, + incomingAcceptorThreads, + incomingWorkerThreads, + outgoingWorkerThreads, + name, + serverGroupId); + protocolThreadPools.put(protocol, threadPools); } - - return protocolThreadPools.get(protocol); + } } - /** - * Lock controlling access to the {@link #registerProxyServer(HttpProxyServer)} and {@link #unregisterProxyServer(HttpProxyServer, boolean)} - * methods. - */ - private final Object SERVER_REGISTRATION_LOCK = new Object(); - - /** - * Registers the specified proxy server as a consumer of this server group. The server group will not be shut down - * until the proxy unregisters itself. - * - * @param proxyServer proxy server instance to register - */ - public void registerProxyServer(HttpProxyServer proxyServer) { - synchronized (SERVER_REGISTRATION_LOCK) { - registeredServers.add(proxyServer); - } + return protocolThreadPools.get(protocol); + } + + /** + * Lock controlling access to the {@link #registerProxyServer(HttpProxyServer)} and {@link + * #unregisterProxyServer(HttpProxyServer, boolean)} methods. + */ + private final Object SERVER_REGISTRATION_LOCK = new Object(); + + /** + * Registers the specified proxy server as a consumer of this server group. The server group will + * not be shut down until the proxy unregisters itself. + * + * @param proxyServer proxy server instance to register + */ + public void registerProxyServer(HttpProxyServer proxyServer) { + synchronized (SERVER_REGISTRATION_LOCK) { + registeredServers.add(proxyServer); } - - /** - * Unregisters the specified proxy server from this server group. If this was the last registered proxy server, the - * server group will be shut down. - * - * @param proxyServer proxy server instance to unregister - * @param graceful when true, the server group shutdown (if necessary) will be graceful - */ - public void unregisterProxyServer(HttpProxyServer proxyServer, boolean graceful) { - synchronized (SERVER_REGISTRATION_LOCK) { - boolean wasRegistered = registeredServers.remove(proxyServer); - if (!wasRegistered) { - log.warn("Attempted to unregister proxy server from ServerGroup that it was not registered with. Was the proxy unregistered twice?"); - } - - if (registeredServers.isEmpty()) { - log.debug("Proxy server unregistered from ServerGroup. No proxy servers remain registered, so shutting down ServerGroup."); - - shutdown(graceful); - } else { - log.debug("Proxy server unregistered from ServerGroup. Not shutting down ServerGroup ({} proxy servers remain registered).", registeredServers.size()); - } - } + } + + /** + * Unregisters the specified proxy server from this server group. If this was the last registered + * proxy server, the server group will be shut down. + * + * @param proxyServer proxy server instance to unregister + * @param graceful when true, the server group shutdown (if necessary) will be graceful + */ + public void unregisterProxyServer(HttpProxyServer proxyServer, boolean graceful) { + synchronized (SERVER_REGISTRATION_LOCK) { + boolean wasRegistered = registeredServers.remove(proxyServer); + if (!wasRegistered) { + log.warn( + "Attempted to unregister proxy server from ServerGroup that it was not registered with. Was the proxy unregistered twice?"); + } + + if (registeredServers.isEmpty() && autoStop) { + log.debug( + "Proxy server unregistered from ServerGroup. No proxy servers remain registered, so shutting down ServerGroup."); + + shutdown(graceful); + } else { + log.debug( + "Proxy server unregistered from ServerGroup. Not shutting down ServerGroup ({} proxy servers remain registered).", + registeredServers.size()); + } + } + } + + /** + * Shuts down all event loops owned by this server group. + * + * @param graceful when true, event loops will "gracefully" terminate, waiting for submitted tasks + * to finish + */ + public void shutdown(boolean graceful) { + if (!stopped.compareAndSet(false, true)) { + log.info("Shutdown requested, but ServerGroup is already stopped. Doing nothing."); + + return; } - /** - * Shuts down all event loops owned by this server group. - * - * @param graceful when true, event loops will "gracefully" terminate, waiting for submitted tasks to finish - */ - private void shutdown(boolean graceful) { - if (!stopped.compareAndSet(false, true)) { - log.info("Shutdown requested, but ServerGroup is already stopped. Doing nothing."); - - return; - } - - log.info("Shutting down server group event loops " + (graceful ? "(graceful)" : "(non-graceful)")); - - // loop through all event loops managed by this server group. this includes acceptor and worker event loops - // for both TCP and UDP transport protocols. - List allEventLoopGroups = new ArrayList<>(); - - for (ProxyThreadPools threadPools : protocolThreadPools.values()) { - allEventLoopGroups.addAll(threadPools.getAllEventLoops()); - } - - for (EventLoopGroup group : allEventLoopGroups) { - if (graceful) { - group.shutdownGracefully(); - } else { - group.shutdownGracefully(0, 0, TimeUnit.SECONDS); - } - } - - if (graceful) { - for (EventLoopGroup group : allEventLoopGroups) { - try { - group.awaitTermination(60, TimeUnit.SECONDS); - } catch (InterruptedException e) { - Thread.currentThread().interrupt(); - - log.warn("Interrupted while shutting down event loop"); - } - } - } + log.info( + "Shutting down server group event loops {}", graceful ? "(graceful)" : "(non-graceful)"); - log.debug("Done shutting down server group"); - } + // loop through all event loops managed by this server group. this includes acceptor and worker + // event loops + // for both TCP and UDP transport protocols. + List allEventLoopGroups = new ArrayList<>(); - /** - * Retrieves the client-to-proxy acceptor thread pool for the specified protocol. Initializes the pool if it has not - * yet been initialized. - *

- * This method is thread-safe; no external locking is necessary. - * - * @param protocol transport protocol to retrieve the thread pool for - * @return the client-to-proxy acceptor thread pool - */ - public EventLoopGroup getClientToProxyAcceptorPoolForTransport(TransportProtocol protocol) { - return getThreadPoolsForProtocol(protocol).getClientToProxyAcceptorPool(); + for (ProxyThreadPools threadPools : protocolThreadPools.values()) { + allEventLoopGroups.addAll(threadPools.getAllEventLoops()); } - /** - * Retrieves the client-to-proxy acceptor worker pool for the specified protocol. Initializes the pool if it has not - * yet been initialized. - *

- * This method is thread-safe; no external locking is necessary. - * - * @param protocol transport protocol to retrieve the thread pool for - * @return the client-to-proxy worker thread pool - */ - public EventLoopGroup getClientToProxyWorkerPoolForTransport(TransportProtocol protocol) { - return getThreadPoolsForProtocol(protocol).getClientToProxyWorkerPool(); + for (EventLoopGroup group : allEventLoopGroups) { + if (graceful) { + group.shutdownGracefully(); + } else { + group.shutdownGracefully(0, 0, TimeUnit.SECONDS); + } } - /** - * Retrieves the proxy-to-server worker thread pool for the specified protocol. Initializes the pool if it has not - * yet been initialized. - *

- * This method is thread-safe; no external locking is necessary. - * - * @param protocol transport protocol to retrieve the thread pool for - * @return the proxy-to-server worker thread pool - */ - public EventLoopGroup getProxyToServerWorkerPoolForTransport(TransportProtocol protocol) { - return getThreadPoolsForProtocol(protocol).getProxyToServerWorkerPool(); - } + if (graceful) { + for (EventLoopGroup group : allEventLoopGroups) { + try { + group.awaitTermination(60, TimeUnit.SECONDS); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); - /** - * @return true if this ServerGroup has already been stopped - */ - public boolean isStopped() { - return stopped.get(); + log.warn("Interrupted while shutting down event loop"); + } + } } + log.debug("Done shutting down server group"); + } + + /** + * Retrieves the client-to-proxy acceptor thread pool for the specified protocol. Initializes the + * pool if it has not yet been initialized. + * + *

This method is thread-safe; no external locking is necessary. + * + * @param protocol transport protocol to retrieve the thread pool for + * @return the client-to-proxy acceptor thread pool + */ + public EventLoopGroup getClientToProxyAcceptorPoolForTransport(TransportProtocol protocol) { + return getThreadPoolsForProtocol(protocol).getClientToProxyAcceptorPool(); + } + + /** + * Retrieves the client-to-proxy acceptor worker pool for the specified protocol. Initializes the + * pool if it has not yet been initialized. + * + *

This method is thread-safe; no external locking is necessary. + * + * @param protocol transport protocol to retrieve the thread pool for + * @return the client-to-proxy worker thread pool + */ + public EventLoopGroup getClientToProxyWorkerPoolForTransport(TransportProtocol protocol) { + return getThreadPoolsForProtocol(protocol).getClientToProxyWorkerPool(); + } + + /** + * Retrieves the proxy-to-server worker thread pool for the specified protocol. Initializes the + * pool if it has not yet been initialized. + * + *

This method is thread-safe; no external locking is necessary. + * + * @param protocol transport protocol to retrieve the thread pool for + * @return the proxy-to-server worker thread pool + */ + public EventLoopGroup getProxyToServerWorkerPoolForTransport(TransportProtocol protocol) { + return getThreadPoolsForProtocol(protocol).getProxyToServerWorkerPool(); + } + + /** + * @return true if this ServerGroup has already been stopped + */ + public boolean isStopped() { + return stopped.get(); + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/SupplierEx.java b/src/main/java/org/littleshoot/proxy/impl/SupplierEx.java new file mode 100644 index 00000000..66cf3851 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/SupplierEx.java @@ -0,0 +1,6 @@ +package org.littleshoot.proxy.impl; + +@FunctionalInterface +interface SupplierEx { + T get() throws Exception; +} diff --git a/src/main/java/org/littleshoot/proxy/impl/ThreadPoolConfiguration.java b/src/main/java/org/littleshoot/proxy/impl/ThreadPoolConfiguration.java index b8e7b058..91545fa2 100644 --- a/src/main/java/org/littleshoot/proxy/impl/ThreadPoolConfiguration.java +++ b/src/main/java/org/littleshoot/proxy/impl/ThreadPoolConfiguration.java @@ -1,62 +1,62 @@ package org.littleshoot.proxy.impl; /** - * Configuration object for the proxy's thread pools. Controls the number of acceptor and worker threads in the Netty - * {@link io.netty.channel.EventLoopGroup} used by the proxy. + * Configuration object for the proxy's thread pools. Controls the number of acceptor and worker + * threads in the Netty {@link io.netty.channel.EventLoopGroup} used by the proxy. */ public class ThreadPoolConfiguration { - private int acceptorThreads = ServerGroup.DEFAULT_INCOMING_ACCEPTOR_THREADS; - private int clientToProxyWorkerThreads = ServerGroup.DEFAULT_INCOMING_WORKER_THREADS; - private int proxyToServerWorkerThreads = ServerGroup.DEFAULT_OUTGOING_WORKER_THREADS; - - public int getClientToProxyWorkerThreads() { - return clientToProxyWorkerThreads; - } - - /** - * Set the number of client-to-proxy worker threads to create. Worker threads perform the actual processing of - * client requests. The default value is {@link ServerGroup#DEFAULT_INCOMING_WORKER_THREADS}. - * - * @param clientToProxyWorkerThreads number of client-to-proxy worker threads to create - * @return this thread pool configuration instance, for chaining - */ - public ThreadPoolConfiguration withClientToProxyWorkerThreads(int clientToProxyWorkerThreads) { - this.clientToProxyWorkerThreads = clientToProxyWorkerThreads; - return this; - } - - public int getAcceptorThreads() { - return acceptorThreads; - } - - /** - * Set the number of acceptor threads to create. Acceptor threads accept HTTP connections from the client and queue - * them for processing by client-to-proxy worker threads. The default value is - * {@link ServerGroup#DEFAULT_INCOMING_ACCEPTOR_THREADS}. - * - * @param acceptorThreads number of acceptor threads to create - * @return this thread pool configuration instance, for chaining - */ - public ThreadPoolConfiguration withAcceptorThreads(int acceptorThreads) { - this.acceptorThreads = acceptorThreads; - return this; - } - - public int getProxyToServerWorkerThreads() { - return proxyToServerWorkerThreads; - } - - /** - * Set the number of proxy-to-server worker threads to create. Proxy-to-server worker threads make requests to - * upstream servers and process responses from the server. The default value is - * {@link ServerGroup#DEFAULT_OUTGOING_WORKER_THREADS}. - * - * @param proxyToServerWorkerThreads number of proxy-to-server worker threads to create - * @return this thread pool configuration instance, for chaining - */ - public ThreadPoolConfiguration withProxyToServerWorkerThreads(int proxyToServerWorkerThreads) { - this.proxyToServerWorkerThreads = proxyToServerWorkerThreads; - return this; - } - + private int acceptorThreads = ServerGroup.DEFAULT_INCOMING_ACCEPTOR_THREADS; + private int clientToProxyWorkerThreads = ServerGroup.DEFAULT_INCOMING_WORKER_THREADS; + private int proxyToServerWorkerThreads = ServerGroup.DEFAULT_OUTGOING_WORKER_THREADS; + + public int getClientToProxyWorkerThreads() { + return clientToProxyWorkerThreads; + } + + /** + * Set the number of client-to-proxy worker threads to create. Worker threads perform the actual + * processing of client requests. The default value is {@link + * ServerGroup#DEFAULT_INCOMING_WORKER_THREADS}. + * + * @param clientToProxyWorkerThreads number of client-to-proxy worker threads to create + * @return this thread pool configuration instance, for chaining + */ + public ThreadPoolConfiguration withClientToProxyWorkerThreads(int clientToProxyWorkerThreads) { + this.clientToProxyWorkerThreads = clientToProxyWorkerThreads; + return this; + } + + public int getAcceptorThreads() { + return acceptorThreads; + } + + /** + * Set the number of acceptor threads to create. Acceptor threads accept HTTP connections from the + * client and queue them for processing by client-to-proxy worker threads. The default value is + * {@link ServerGroup#DEFAULT_INCOMING_ACCEPTOR_THREADS}. + * + * @param acceptorThreads number of acceptor threads to create + * @return this thread pool configuration instance, for chaining + */ + public ThreadPoolConfiguration withAcceptorThreads(int acceptorThreads) { + this.acceptorThreads = acceptorThreads; + return this; + } + + public int getProxyToServerWorkerThreads() { + return proxyToServerWorkerThreads; + } + + /** + * Set the number of proxy-to-server worker threads to create. Proxy-to-server worker threads make + * requests to upstream servers and process responses from the server. The default value is {@link + * ServerGroup#DEFAULT_OUTGOING_WORKER_THREADS}. + * + * @param proxyToServerWorkerThreads number of proxy-to-server worker threads to create + * @return this thread pool configuration instance, for chaining + */ + public ThreadPoolConfiguration withProxyToServerWorkerThreads(int proxyToServerWorkerThreads) { + this.proxyToServerWorkerThreads = proxyToServerWorkerThreads; + return this; + } } diff --git a/src/main/java/org/littleshoot/proxy/impl/WebSocketFramePipeHandler.java b/src/main/java/org/littleshoot/proxy/impl/WebSocketFramePipeHandler.java new file mode 100644 index 00000000..32a19d81 --- /dev/null +++ b/src/main/java/org/littleshoot/proxy/impl/WebSocketFramePipeHandler.java @@ -0,0 +1,71 @@ +package org.littleshoot.proxy.impl; + +import static java.util.Objects.requireNonNull; + +import io.netty.buffer.ByteBuf; +import io.netty.channel.Channel; +import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.ChannelInboundHandlerAdapter; +import java.util.function.Supplier; +import org.littleshoot.proxy.HttpFilters; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * A {@link ChannelInboundHandlerAdapter} that forwards raw WebSocket frame bytes to the peer + * connection and optionally notifies an {@link HttpFilters} observer via {@link + * HttpFilters#webSocketFrameReceived(Supplier, boolean)}. + * + *

Installed on both the client-to-proxy and proxy-to-server channels after a WebSocket upgrade, + * replacing the HTTP codec pipeline. + */ +public class WebSocketFramePipeHandler extends ChannelInboundHandlerAdapter { + private static final Logger log = LoggerFactory.getLogger(WebSocketFramePipeHandler.class); + private final ProxyConnection sink; + private final HttpFilters filters; + private final boolean fromClient; + + public WebSocketFramePipeHandler( + final ProxyConnection sink, final HttpFilters filters, final boolean fromClient) { + this.sink = requireNonNull(sink, "sink cannot be null"); + this.filters = filters; + this.fromClient = fromClient; + } + + @Override + public void channelRead(final ChannelHandlerContext ctx, final Object msg) { + if (filters != null && msg instanceof ByteBuf) { + try { + filters.webSocketFrameReceived(new WebSocketFrameBytes((ByteBuf) msg), fromClient); + } catch (Exception e) { + log.error("Failed to notify listeners that websocket frame received", e); + } + } + Channel channel = sink.channel; + if (channel != null) { + channel.writeAndFlush(msg); + } + } + + @Override + public void channelInactive(final ChannelHandlerContext ctx) { + if (!sink.getCurrentState().isDisconnectingOrDisconnected()) { + sink.disconnect(); + } + } + + private static class WebSocketFrameBytes implements Supplier { + private final ByteBuf message; + + private WebSocketFrameBytes(ByteBuf message) { + this.message = message; + } + + @Override + public byte[] get() { + byte[] frameBytes = new byte[message.readableBytes()]; + message.getBytes(message.readerIndex(), frameBytes); + return frameBytes; + } + } +} diff --git a/src/main/resources/littleproxy_async_log4j2.xml b/src/main/resources/littleproxy_async_log4j2.xml new file mode 100644 index 00000000..75bcc8df --- /dev/null +++ b/src/main/resources/littleproxy_async_log4j2.xml @@ -0,0 +1,51 @@ + + + + + + + + + + + + + + %d{ISO8601} %-5p [%t] %c{2} (%F:%L).%M() - %m%n + + + + + + + + + + %m%n + + + + + + + + + + + + + + + + + + + + + + + \ No newline at end of file diff --git a/src/main/resources/littleproxy_default_log4j2.xml b/src/main/resources/littleproxy_default_log4j2.xml new file mode 100644 index 00000000..3bbeedab --- /dev/null +++ b/src/main/resources/littleproxy_default_log4j2.xml @@ -0,0 +1,50 @@ + + + + + + + + + + + + + + %d{ISO8601} %-5p [%t] %c{2} (%F:%L).%M() - %m%n + + + + + + + + + + + + %d{ISO8601} %-5p [%t] %c{2} (%F:%L).%M() - %m%n + + + + + + + + + + + + + + + + + + + + + diff --git a/src/test/java/org/littleshoot/proxy/AbstractProxyTest.java b/src/test/java/org/littleshoot/proxy/AbstractProxyTest.java index 26a53d9e..65f9b4d8 100644 --- a/src/test/java/org/littleshoot/proxy/AbstractProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/AbstractProxyTest.java @@ -1,6 +1,12 @@ package org.littleshoot.proxy; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.buildHttpClient; + import io.netty.handler.codec.http.HttpRequest; +import java.io.IOException; +import java.util.concurrent.atomic.AtomicInteger; +import javax.net.ssl.SSLSession; import org.apache.http.HttpEntity; import org.apache.http.HttpHost; import org.apache.http.HttpResponse; @@ -11,348 +17,300 @@ import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.util.EntityUtils; import org.eclipse.jetty.server.Server; -import org.junit.After; -import org.junit.Before; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; - -import javax.net.ssl.SSLSession; -import java.net.InetSocketAddress; -import java.util.concurrent.atomic.AtomicInteger; - -import static org.hamcrest.Matchers.greaterThan; -import static org.junit.Assert.assertEquals; -import static org.junit.Assert.assertThat; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; /** - * Base for tests that test the proxy. This base class encapsulates all of the - * testing infrastructure. + * Base for tests that test the proxy. This base class encapsulates all the testing infrastructure. */ public abstract class AbstractProxyTest { - protected static final String DEFAULT_RESOURCE = "/"; - - protected int webServerPort = -1; - protected int httpsWebServerPort = -1; - - protected HttpHost webHost; - protected HttpHost httpsWebHost; - - /** - * The server used by the tests. - */ - protected HttpProxyServer proxyServer; - - /** - * Holds the most recent response after executing a test method. - */ - protected String lastResponse; - - /** - * The web server that provides the back-end. - */ - private Server webServer; - - protected AtomicInteger bytesReceivedFromClient; - protected AtomicInteger requestsReceivedFromClient; - protected AtomicInteger bytesSentToServer; - protected AtomicInteger requestsSentToServer; - protected AtomicInteger bytesReceivedFromServer; - protected AtomicInteger responsesReceivedFromServer; - protected AtomicInteger bytesSentToClient; - protected AtomicInteger responsesSentToClient; - protected AtomicInteger clientConnects; - protected AtomicInteger clientSSLHandshakeSuccesses; - protected AtomicInteger clientDisconnects; - - @Before - public void initializeCounters() { - bytesReceivedFromClient = new AtomicInteger(0); - requestsReceivedFromClient = new AtomicInteger(0); - bytesSentToServer = new AtomicInteger(0); - requestsSentToServer = new AtomicInteger(0); - bytesReceivedFromServer = new AtomicInteger(0); - responsesReceivedFromServer = new AtomicInteger(0); - bytesSentToClient = new AtomicInteger(0); - responsesSentToClient = new AtomicInteger(0); - clientConnects = new AtomicInteger(0); - clientSSLHandshakeSuccesses = new AtomicInteger(0); - clientDisconnects = new AtomicInteger(0); - } - - @Before - public void runSetUp() throws Exception { - webServer = TestUtils.startWebServer(true); - - // find out what ports the HTTP and HTTPS connectors were bound to - httpsWebServerPort = TestUtils.findLocalHttpsPort(webServer); - if (httpsWebServerPort < 0) { - throw new RuntimeException("HTTPS connector should already be open and listening, but port was " + webServerPort); - } - - webServerPort = TestUtils.findLocalHttpPort(webServer); - if (webServerPort < 0) { - throw new RuntimeException("HTTP connector should already be open and listening, but port was " + webServerPort); - } - - webHost = new HttpHost("127.0.0.1", webServerPort); - httpsWebHost = new HttpHost("127.0.0.1", httpsWebServerPort, "https"); - - setUp(); - } - - protected abstract void setUp() throws Exception; - - @After - public void runTearDown() throws Exception { - try { - tearDown(); - } finally { - try { - if (this.proxyServer != null) { - this.proxyServer.abort(); - } - } finally { - if (this.webServer != null) { - webServer.stop(); - } - } - } - } - - protected void tearDown() throws Exception { + protected static final String DEFAULT_RESOURCE = "/"; + private static final String DEFAULT_JKS_KEYSTORE_PATH = "target/littleproxy_keystore.jks"; + + protected int webServerPort = -1; + protected int httpsWebServerPort = -1; + + protected HttpHost webHost; + protected HttpHost httpsWebHost; + + /** The server used by the tests. */ + protected HttpProxyServer proxyServer; + + /** Holds the most recent response after executing a test method. */ + protected String lastResponse; + + /** The web server that provides the back-end. */ + private Server webServer; + + private final AtomicInteger bytesReceivedFromClient = new AtomicInteger(0); + private final AtomicInteger requestsReceivedFromClient = new AtomicInteger(0); + private final AtomicInteger bytesSentToServer = new AtomicInteger(0); + private final AtomicInteger requestsSentToServer = new AtomicInteger(0); + private final AtomicInteger bytesReceivedFromServer = new AtomicInteger(0); + private final AtomicInteger responsesReceivedFromServer = new AtomicInteger(0); + private final AtomicInteger bytesSentToClient = new AtomicInteger(0); + private final AtomicInteger responsesSentToClient = new AtomicInteger(0); + private final AtomicInteger clientConnects = new AtomicInteger(0); + private final AtomicInteger clientSSLHandshakeSuccesses = new AtomicInteger(0); + private final AtomicInteger clientDisconnects = new AtomicInteger(0); + + protected Logger logger() { + return LoggerFactory.getLogger(getClass()); + } + + @BeforeEach + final void runSetUp() throws Exception { + webServer = TestUtils.startWebServer(true, DEFAULT_JKS_KEYSTORE_PATH); + + // find out what ports the HTTP and HTTPS connectors were bound to + httpsWebServerPort = TestUtils.findLocalHttpsPort(webServer); + if (httpsWebServerPort < 0) { + throw new RuntimeException( + "HTTPS connector should already be open and listening, but port was " + webServerPort); } - /** - * Override this to specify a username to use when authenticating with - * proxy. - */ - protected String getUsername() { - return null; + webServerPort = TestUtils.findLocalHttpPort(webServer); + if (webServerPort < 0) { + throw new RuntimeException( + "HTTP connector should already be open and listening, but port was " + webServerPort); } - /** - * Override this to specify a password to use when authenticating with - * proxy. - */ - protected String getPassword() { - return null; - } - - protected void assertReceivedBadGateway(ResponseInfo response) { - assertEquals("Received: " + response, 502, response.getStatusCode()); - } - - protected ResponseInfo httpPostWithApacheClient( - HttpHost host, String resourceUrl, boolean isProxied) - throws Exception { - final boolean supportSsl = true; - String username = getUsername(); - String password = getPassword(); - try (CloseableHttpClient httpClient = TestUtils.buildHttpClient( - isProxied, supportSsl, proxyServer.getListenAddress().getPort(), username, password)) { - final HttpPost request = new HttpPost(resourceUrl); - request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); - - final StringEntity entity = new StringEntity("adsf", "UTF-8"); - entity.setChunked(true); - request.setEntity(entity); - - final HttpResponse response = httpClient.execute(host, request); - final HttpEntity resEntity = response.getEntity(); - return new ResponseInfo(response.getStatusLine().getStatusCode(), - EntityUtils.toString(resEntity)); - } - } - - protected ResponseInfo httpGetWithApacheClient(HttpHost host, - String resourceUrl, boolean isProxied, boolean callHeadFirst) - throws Exception { - final boolean supportSsl = true; - String username = getUsername(); - String password = getPassword(); - try (CloseableHttpClient httpClient = TestUtils.buildHttpClient( - isProxied, supportSsl, proxyServer.getListenAddress().getPort(), username, password)){ - Integer contentLength = null; - if (callHeadFirst) { - HttpHead request = new HttpHead(resourceUrl); - request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); - HttpResponse response = httpClient.execute(host, request); - contentLength = new Integer(response.getFirstHeader( - "Content-Length").getValue()); - } - - HttpGet request = new HttpGet(resourceUrl); - request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); - - HttpResponse response = httpClient.execute(host, request); - HttpEntity resEntity = response.getEntity(); - - if (contentLength != null) { - assertEquals( - "Content-Length from GET should match that from HEAD", - contentLength, - new Integer(response.getFirstHeader("Content-Length") - .getValue())); - } - return new ResponseInfo(response.getStatusLine().getStatusCode(), - EntityUtils.toString(resEntity)); + webHost = new HttpHost("127.0.0.1", webServerPort); + httpsWebHost = new HttpHost("127.0.0.1", httpsWebServerPort, "https"); + logger().info("Started webserver http:{}, https:{}", webServerPort, httpsWebServerPort); + + setUp(); + logger().info("Started proxy server {}", proxyServer.getListenAddress()); + } + + protected abstract void setUp() throws Exception; + + @AfterEach + final void runTearDown() throws Exception { + try { + tearDown(); + } finally { + try { + if (proxyServer != null) { + logger().info("Stop proxy server {}", proxyServer.getListenAddress()); + proxyServer.abort(); } - } - - protected String compareProxiedAndUnproxiedPOST(HttpHost host, - String resourceUrl) throws Exception { - ResponseInfo proxiedResponse = httpPostWithApacheClient(host, - resourceUrl, true); - if (expectBadGatewayForEverything()) { - assertReceivedBadGateway(proxiedResponse); - } else { - ResponseInfo unproxiedResponse = httpPostWithApacheClient(host, - resourceUrl, false); - assertEquals(unproxiedResponse, proxiedResponse); - checkStatistics(host); + } finally { + if (webServer != null) { + logger().info("Stop webserver http:{}, https:{}", webServerPort, httpsWebServerPort); + webServer.stop(); } - return proxiedResponse.getBody(); + } } - - protected String compareProxiedAndUnproxiedGET(HttpHost host, - String resourceUrl) throws Exception { - ResponseInfo proxiedResponse = httpGetWithApacheClient(host, - resourceUrl, true, false); - if (expectBadGatewayForEverything()) { - assertReceivedBadGateway(proxiedResponse); - } else { - ResponseInfo unproxiedResponse = httpGetWithApacheClient(host, - resourceUrl, false, false); - assertEquals(unproxiedResponse, proxiedResponse); - checkStatistics(host); - } - return proxiedResponse.getBody(); + } + + protected void tearDown() throws Exception {} + + /** Override this to specify a username to use when authenticating with proxy. */ + protected String getUsername() { + return null; + } + + /** Override this to specify a password to use when authenticating with proxy. */ + protected String getPassword() { + return null; + } + + protected void assertReceivedBadGateway(ResponseInfo response) { + assertThat(response.getStatusCode()).as("Received: %s", response).isEqualTo(502); + } + + protected ResponseInfo httpPostWithApacheClient( + HttpHost host, String resourceUrl, boolean isProxied) { + final boolean supportSsl = true; + String username = getUsername(); + String password = getPassword(); + try (CloseableHttpClient httpClient = + buildHttpClient( + isProxied, supportSsl, proxyServer.getListenAddress().getPort(), username, password)) { + final HttpPost request = new HttpPost(resourceUrl); + request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); + + final StringEntity entity = new StringEntity("adsf", "UTF-8"); + entity.setChunked(true); + request.setEntity(entity); + + final HttpResponse response = httpClient.execute(host, request); + final HttpEntity resEntity = response.getEntity(); + return new ResponseInfo( + response.getStatusLine().getStatusCode(), EntityUtils.toString(resEntity)); + } catch (IOException e) { + throw new RuntimeException(e); } - - private void checkStatistics(HttpHost host) { - boolean isHTTPS = host.getSchemeName().equalsIgnoreCase("HTTPS"); - int numberOfExpectedClientInteractions = 1; - int numberOfExpectedServerInteractions = 1; - if (isAuthenticating()) { - numberOfExpectedClientInteractions += 1; - } - if (isHTTPS && isMITM()) { - numberOfExpectedClientInteractions += 1; - numberOfExpectedServerInteractions += 1; - } - if (isHTTPS && !isChained()) { - numberOfExpectedServerInteractions -= 1; - } - assertThat(bytesReceivedFromClient.get(), greaterThan(0)); - assertEquals(numberOfExpectedClientInteractions, - requestsReceivedFromClient.get()); - assertThat(bytesSentToServer.get(), greaterThan(0)); - assertEquals(numberOfExpectedServerInteractions, - requestsSentToServer.get()); - assertThat(bytesReceivedFromServer.get(), greaterThan(0)); - assertEquals(numberOfExpectedServerInteractions, - responsesReceivedFromServer.get()); - assertThat(bytesSentToClient.get(), greaterThan(0)); - assertEquals(numberOfExpectedClientInteractions, - responsesSentToClient.get()); + } + + protected ResponseInfo httpGetWithApacheClient( + HttpHost host, String resourceUrl, boolean isProxied, boolean callHeadFirst) { + final boolean supportSsl = true; + String username = getUsername(); + String password = getPassword(); + try (CloseableHttpClient httpClient = + buildHttpClient( + isProxied, supportSsl, proxyServer.getListenAddress().getPort(), username, password)) { + Integer contentLength = null; + if (callHeadFirst) { + HttpHead request = new HttpHead(resourceUrl); + request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); + HttpResponse response = httpClient.execute(host, request); + contentLength = Integer.valueOf(response.getFirstHeader("Content-Length").getValue()); + } + + HttpGet request = new HttpGet(resourceUrl); + request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); + + HttpResponse response = httpClient.execute(host, request); + HttpEntity resEntity = response.getEntity(); + + if (contentLength != null) { + assertThat(Integer.valueOf(response.getFirstHeader("Content-Length").getValue())) + .as("Content-Length from GET should match that from HEAD") + .isEqualTo(contentLength); + } + return new ResponseInfo( + response.getStatusLine().getStatusCode(), EntityUtils.toString(resEntity)); + } catch (IOException e) { + throw new RuntimeException(e); } - - /** - * Override this to indicate that the proxy is chained. - */ - protected boolean isChained() { - return false; + } + + protected String compareProxiedAndUnproxiedPOST(HttpHost host, String resourceUrl) { + ResponseInfo proxiedResponse = httpPostWithApacheClient(host, resourceUrl, true); + if (expectBadGatewayForEverything()) { + assertReceivedBadGateway(proxiedResponse); + } else { + ResponseInfo unproxiedResponse = httpPostWithApacheClient(host, resourceUrl, false); + assertThat(proxiedResponse).isEqualTo(unproxiedResponse); + checkStatistics(host); } - - /** - * Override this to indicate that the test uses authentication. - */ - protected boolean isAuthenticating() { - return false; + return proxiedResponse.getBody(); + } + + protected String compareProxiedAndUnproxiedGET(HttpHost host, String resourceUrl) { + ResponseInfo proxiedResponse = httpGetWithApacheClient(host, resourceUrl, true, false); + if (expectBadGatewayForEverything()) { + assertReceivedBadGateway(proxiedResponse); + } else { + ResponseInfo unproxiedResponse = httpGetWithApacheClient(host, resourceUrl, false, false); + assertThat(proxiedResponse).isEqualTo(unproxiedResponse); + checkStatistics(host); } - - protected boolean isMITM() { - return false; + return proxiedResponse.getBody(); + } + + private void checkStatistics(HttpHost host) { + boolean isHTTPS = "HTTPS".equalsIgnoreCase(host.getSchemeName()); + int numberOfExpectedClientInteractions = 1; + int numberOfExpectedServerInteractions = 1; + if (isAuthenticating()) { + numberOfExpectedClientInteractions += 1; } - - protected boolean expectBadGatewayForEverything() { - return false; + if (isHTTPS && isMITM()) { + numberOfExpectedClientInteractions += 1; + numberOfExpectedServerInteractions += 1; } - - protected HttpProxyServerBootstrap bootstrapProxy() { - return DefaultHttpProxyServer.bootstrap().plusActivityTracker( - new ActivityTracker() { - @Override - public void bytesReceivedFromClient( - FlowContext flowContext, - int numberOfBytes) { - bytesReceivedFromClient.addAndGet(numberOfBytes); - } - - @Override - public void requestReceivedFromClient( - FlowContext flowContext, - HttpRequest httpRequest) { - requestsReceivedFromClient.incrementAndGet(); - } - - @Override - public void bytesSentToServer(FullFlowContext flowContext, - int numberOfBytes) { - bytesSentToServer.addAndGet(numberOfBytes); - } - - @Override - public void requestSentToServer( - FullFlowContext flowContext, - HttpRequest httpRequest) { - requestsSentToServer.incrementAndGet(); - } - - @Override - public void bytesReceivedFromServer( - FullFlowContext flowContext, - int numberOfBytes) { - bytesReceivedFromServer.addAndGet(numberOfBytes); - } - - @Override - public void responseReceivedFromServer( - FullFlowContext flowContext, - io.netty.handler.codec.http.HttpResponse httpResponse) { - responsesReceivedFromServer.incrementAndGet(); - } - - @Override - public void bytesSentToClient(FlowContext flowContext, - int numberOfBytes) { - bytesSentToClient.addAndGet(numberOfBytes); - } - - @Override - public void responseSentToClient( - FlowContext flowContext, - io.netty.handler.codec.http.HttpResponse httpResponse) { - responsesSentToClient.incrementAndGet(); - } - - @Override - public void clientConnected(InetSocketAddress clientAddress) { - clientConnects.incrementAndGet(); - } - - @Override - public void clientSSLHandshakeSucceeded( - InetSocketAddress clientAddress, - SSLSession sslSession) { - clientSSLHandshakeSuccesses.incrementAndGet(); - } - - @Override - public void clientDisconnected( - InetSocketAddress clientAddress, - SSLSession sslSession) { - clientDisconnects.incrementAndGet(); - } - }); + if (isHTTPS && !isChained()) { + numberOfExpectedServerInteractions -= 1; } + assertThat(bytesReceivedFromClient.get()).isGreaterThan(0); + assertThat(requestsReceivedFromClient.get()).isEqualTo(numberOfExpectedClientInteractions); + assertThat(bytesSentToServer.get()).isGreaterThan(0); + assertThat(requestsSentToServer.get()).isEqualTo(numberOfExpectedServerInteractions); + assertThat(bytesReceivedFromServer.get()).isGreaterThan(0); + assertThat(responsesReceivedFromServer.get()).isEqualTo(numberOfExpectedServerInteractions); + assertThat(bytesSentToClient.get()).isGreaterThan(0); + assertThat(responsesSentToClient.get()).isEqualTo(numberOfExpectedClientInteractions); + } + + /** Override this to indicate that the proxy is chained. */ + protected boolean isChained() { + return false; + } + + /** Override this to indicate that the test does use authentication. */ + protected boolean isAuthenticating() { + return false; + } + + protected boolean isMITM() { + return false; + } + + protected boolean expectBadGatewayForEverything() { + return false; + } + + protected HttpProxyServerBootstrap bootstrapProxy() { + return DefaultHttpProxyServer.bootstrap() + .plusActivityTracker( + new ActivityTrackerAdapter() { + @Override + public void bytesReceivedFromClient(FlowContext flowContext, int numberOfBytes) { + bytesReceivedFromClient.addAndGet(numberOfBytes); + } + + @Override + public void requestReceivedFromClient( + FlowContext flowContext, HttpRequest httpRequest) { + requestsReceivedFromClient.incrementAndGet(); + } + + @Override + public void bytesSentToServer(FullFlowContext flowContext, int numberOfBytes) { + bytesSentToServer.addAndGet(numberOfBytes); + } + + @Override + public void requestSentToServer( + FullFlowContext flowContext, HttpRequest httpRequest) { + requestsSentToServer.incrementAndGet(); + } + + @Override + public void bytesReceivedFromServer(FullFlowContext flowContext, int numberOfBytes) { + bytesReceivedFromServer.addAndGet(numberOfBytes); + } + + @Override + public void responseReceivedFromServer( + FullFlowContext flowContext, + io.netty.handler.codec.http.HttpResponse httpResponse) { + responsesReceivedFromServer.incrementAndGet(); + } + + @Override + public void bytesSentToClient(FlowContext flowContext, int numberOfBytes) { + bytesSentToClient.addAndGet(numberOfBytes); + } + + @Override + public void responseSentToClient( + FlowContext flowContext, io.netty.handler.codec.http.HttpResponse httpResponse) { + responsesSentToClient.incrementAndGet(); + } + + @Override + public void clientConnected(FlowContext flowContext) { + clientConnects.incrementAndGet(); + } + + @Override + public void clientSSLHandshakeSucceeded( + FlowContext flowContext, SSLSession sslSession) { + clientSSLHandshakeSuccesses.incrementAndGet(); + } + + @Override + public void clientDisconnected(FlowContext flowContext, SSLSession sslSession) { + clientDisconnects.incrementAndGet(); + } + }); + } } diff --git a/src/test/java/org/littleshoot/proxy/AuthenticatingProxyWithChainingTest.java b/src/test/java/org/littleshoot/proxy/AuthenticatingProxyWithChainingTest.java index 1d7dfb15..6612e4a1 100644 --- a/src/test/java/org/littleshoot/proxy/AuthenticatingProxyWithChainingTest.java +++ b/src/test/java/org/littleshoot/proxy/AuthenticatingProxyWithChainingTest.java @@ -1,63 +1,63 @@ package org.littleshoot.proxy; -import io.netty.handler.codec.http.HttpRequest; -import org.junit.Assert; -import org.littleshoot.proxy.impl.ClientDetails; +import static org.assertj.core.api.Assertions.assertThat; +import io.netty.handler.codec.http.HttpRequest; import java.util.Queue; +import org.littleshoot.proxy.impl.ClientDetails; -/** - * Tests a single proxy that requires username/password authentication. - */ +/** Tests a single proxy that requires username/password authentication. */ public class AuthenticatingProxyWithChainingTest extends BaseProxyTest - implements ProxyAuthenticator, ChainedProxyManager { - - private ClientDetails savedClientDetails; - - @Override - protected void setUp() { - this.proxyServer = bootstrapProxy() - .withPort(0) - .withProxyAuthenticator(this) - .withChainProxyManager(this) - .start(); - } - - @Override - protected String getUsername() { - return "user1"; - } - - @Override - protected String getPassword() { - return "user2"; - } - - @Override - public boolean authenticate(String userName, String password) { - return getUsername().equals(userName) && getPassword().equals(password); - } - - @Override - protected boolean isAuthenticating() { - return true; - } - - @Override - public String getRealm() { - return null; - } - - @Override - public void lookupChainedProxies(HttpRequest httpRequest, Queue chainedProxies, ClientDetails clientDetails) { - savedClientDetails = clientDetails; - chainedProxies.add(ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION); - } - - @Override - protected void tearDown() throws Exception { - super.tearDown(); - Assert.assertEquals(getUsername(), savedClientDetails.getUserName()); - Assert.assertTrue(savedClientDetails.getClientAddress().getAddress().isLoopbackAddress()); - } + implements ProxyAuthenticator, ChainedProxyManager { + + private ClientDetails savedClientDetails; + + @Override + protected void setUp() { + proxyServer = + bootstrapProxy() + .withPort(0) + .withProxyAuthenticator(this) + .withChainProxyManager(this) + .start(); + } + + @Override + protected String getUsername() { + return "user1"; + } + + @Override + protected String getPassword() { + return "user2"; + } + + @Override + public boolean authenticate(String userName, String password) { + return getUsername().equals(userName) && getPassword().equals(password); + } + + @Override + protected boolean isAuthenticating() { + return true; + } + + @Override + public String getRealm() { + return null; + } + + @Override + public void lookupChainedProxies( + HttpRequest httpRequest, Queue chainedProxies, ClientDetails clientDetails) { + savedClientDetails = clientDetails; + chainedProxies.add(ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION); + } + + @Override + protected void tearDown() throws Exception { + super.tearDown(); + assertThat(savedClientDetails.getUserName()).isEqualTo(getUsername()); + assertThat(savedClientDetails.getClientAddress().getAddress().isLoopbackAddress()).isTrue(); + } } diff --git a/src/test/java/org/littleshoot/proxy/AuthenticationCalledOnceTest.java b/src/test/java/org/littleshoot/proxy/AuthenticationCalledOnceTest.java new file mode 100644 index 00000000..b378e699 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/AuthenticationCalledOnceTest.java @@ -0,0 +1,189 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.buildHttpClient; + +import java.io.IOException; +import java.util.concurrent.atomic.AtomicInteger; +import org.apache.http.HttpHost; +import org.apache.http.HttpResponse; +import org.apache.http.client.methods.HttpGet; +import org.apache.http.impl.client.CloseableHttpClient; +import org.eclipse.jetty.server.Server; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * Test to reproduce issue #56: Proxy server authentication called twice. + * + *

This test verifies that the ProxyAuthenticator.authenticate() method is called exactly once + * per authentication attempt, not multiple times. + * + *

The bug: When a client makes a request through an authenticating proxy: 1. First request + * without credentials -> 407 Proxy Authentication Required 2. Client retries with credentials -> + * authenticate() is called TWICE instead of once + * + * @see Issue #56 + */ +public class AuthenticationCalledOnceTest { + + private static final Logger logger = LoggerFactory.getLogger(AuthenticationCalledOnceTest.class); + + private Server webServer; + private HttpProxyServer proxyServer; + private int webServerPort; + private int httpsWebServerPort; + private HttpHost webHost; + private HttpHost httpsWebHost; + + // Counter to track how many times authenticate() is called + private final AtomicInteger authenticateCallCount = new AtomicInteger(0); + + private static final String DEFAULT_JKS_KEYSTORE_PATH = "target/littleproxy_keystore.jks"; + private static final String DEFAULT_RESOURCE = "/"; + private static final String USERNAME = "testuser"; + private static final String PASSWORD = "testpass"; + + @BeforeEach + public void setUp() throws Exception { + webServer = org.littleshoot.proxy.TestUtils.startWebServer(true, DEFAULT_JKS_KEYSTORE_PATH); + + webServerPort = org.littleshoot.proxy.TestUtils.findLocalHttpPort(webServer); + httpsWebServerPort = org.littleshoot.proxy.TestUtils.findLocalHttpsPort(webServer); + + webHost = new HttpHost("127.0.0.1", webServerPort); + httpsWebHost = new HttpHost("127.0.0.1", httpsWebServerPort, "https"); + + logger.info("Started webserver http:{}, https:{}", webServerPort, httpsWebServerPort); + + // Create proxy with authentication + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withProxyAuthenticator( + new ProxyAuthenticator() { + @Override + public boolean authenticate(String username, String password) { + int count = authenticateCallCount.incrementAndGet(); + logger.info("authenticate() called - count: {}", count); + return USERNAME.equals(username) && PASSWORD.equals(password); + } + + @Override + public String getRealm() { + return "TestRealm"; + } + }) + .start(); + + logger.info("Started proxy server on port {}", proxyServer.getListenAddress().getPort()); + } + + @AfterEach + public void tearDown() throws Exception { + try { + if (proxyServer != null) { + logger.info("Stopping proxy server"); + proxyServer.abort(); + } + } finally { + if (webServer != null) { + logger.info("Stopping webserver"); + webServer.stop(); + } + } + } + + /** + * Test case for issue #56: ProxyAuthenticator.authenticate() is called twice. + * + *

When a client makes a request through an authenticating proxy: 1. First request without + * credentials -> 407 Proxy Authentication Required (authenticate NOT called) 2. Client retries + * with credentials -> authenticate() should be called ONCE + * + *

Bug: Currently authenticate() is being called TWICE during step 2 + */ + @Test + public void testAuthenticateCalledOnceDuringAuthentication() throws IOException { + // Reset counter before test + authenticateCallCount.set(0); + + // Make a request through the proxy WITH authentication credentials + // The HTTP client will automatically handle the 407 challenge + try (CloseableHttpClient httpClient = + buildHttpClient( + true, // isProxied + false, // supportSsl (we're using HTTP not HTTPS) + proxyServer.getListenAddress().getPort(), + USERNAME, + PASSWORD)) { + + final HttpGet request = new HttpGet("http://127.0.0.1:" + webServerPort + DEFAULT_RESOURCE); + final HttpResponse response = httpClient.execute(webHost, request); + + logger.info("Response status: {}", response.getStatusLine().getStatusCode()); + + // Verify the request was successful + assertThat(response.getStatusLine().getStatusCode()).isEqualTo(200); + + // CRITICAL ASSERTION: authenticate() should be called exactly ONCE + // Currently, due to the bug, it's being called twice + int callCount = authenticateCallCount.get(); + logger.info("Total authenticate() calls: {}", callCount); + + // This assertion will FAIL with the current bug (count will be 2) + // After fixing issue #56, this should pass (count should be 1) + assertThat(callCount) + .as( + "authenticate() should be called exactly once during authentication, not %d times", + callCount) + .isEqualTo(1); + } + } + + /** + * Additional test: verify that subsequent requests on an authenticated connection do NOT trigger + * additional authenticate() calls. + * + *

Note: This test may show the authenticate being called multiple times if the HTTP client + * creates new connections. This is separate from issue #56. + */ + @Test + public void testSubsequentRequestsDoNotReauthenticate() throws IOException { + // Reset counter + authenticateCallCount.set(0); + + // Use a single client connection to make multiple requests + try (CloseableHttpClient httpClient = + buildHttpClient( + true, false, proxyServer.getListenAddress().getPort(), USERNAME, PASSWORD)) { + + // First request - triggers authentication + HttpGet request1 = new HttpGet("http://127.0.0.1:" + webServerPort + DEFAULT_RESOURCE); + HttpResponse response1 = httpClient.execute(webHost, request1); + assertThat(response1.getStatusLine().getStatusCode()).isEqualTo(200); + + int countAfterFirstRequest = authenticateCallCount.get(); + logger.info("After first request, authenticate() called {} times", countAfterFirstRequest); + + // Second request - should NOT trigger authentication again on the same connection + HttpGet request2 = new HttpGet("http://127.0.0.1:" + webServerPort + DEFAULT_RESOURCE); + HttpResponse response2 = httpClient.execute(webHost, request2); + assertThat(response2.getStatusLine().getStatusCode()).isEqualTo(200); + + int countAfterSecondRequest = authenticateCallCount.get(); + logger.info("After second request, authenticate() called {} times", countAfterSecondRequest); + + // This test might fail if HTTP client creates new connections + // The key assertion for issue #56 is the first test case + // This test is to verify connection reuse behavior + assertThat(countAfterSecondRequest) + .as("authenticate() should not be called again for subsequent requests") + .isLessThanOrEqualTo(countAfterFirstRequest + 1); // Allow at most 1 additional call + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/BadClientAuthenticationTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/BadClientAuthenticationTCPChainedProxyTest.java index da084a94..8a23ff87 100644 --- a/src/test/java/org/littleshoot/proxy/BadClientAuthenticationTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/BadClientAuthenticationTCPChainedProxyTest.java @@ -1,52 +1,37 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -/** - * Tests that clients are authenticated and that if they're missing certs, we - * get an error. - */ -public class BadClientAuthenticationTCPChainedProxyTest extends - BaseChainedProxyTest { - private final SslEngineSource serverSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - private final SslEngineSource clientSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks", false, false); - - @Override - protected boolean expectBadGatewayForEverything() { +/** Tests that clients are authenticated and that if they're missing certs, we get an error. */ +public final class BadClientAuthenticationTCPChainedProxyTest extends BaseChainedProxyTest { + private final SslEngineSource serverSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks", false, false); + private final SslEngineSource clientSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks", false, false); + + @Override + protected boolean expectBadGatewayForEverything() { + return true; + } + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(serverSslEngineSource); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { return true; - } - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(serverSslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return clientSslEngineSource.newSslEngine(); - } - }; - } + } + + @Override + public SSLEngine newSslEngine() { + return clientSslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/BadServerAuthenticationTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/BadServerAuthenticationTCPChainedProxyTest.java index 68eb934a..191216c6 100644 --- a/src/test/java/org/littleshoot/proxy/BadServerAuthenticationTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/BadServerAuthenticationTCPChainedProxyTest.java @@ -1,52 +1,37 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -/** - * Tests that servers are authenticated and that if they're missing certs, we - * get an error. - */ -public class BadServerAuthenticationTCPChainedProxyTest extends - BaseChainedProxyTest { - protected final SslEngineSource serverSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - protected final SslEngineSource clientSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_2.jks"); - - @Override - protected boolean expectBadGatewayForEverything() { +/** Tests that servers are authenticated and that if they're missing certs, we get an error. */ +public class BadServerAuthenticationTCPChainedProxyTest extends BaseChainedProxyTest { + protected final SslEngineSource serverSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks"); + protected final SslEngineSource clientSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_2.jks"); + + @Override + protected boolean expectBadGatewayForEverything() { + return true; + } + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(serverSslEngineSource); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { return true; - } - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(serverSslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return clientSslEngineSource.newSslEngine(); - } - }; - } + } + + @Override + public SSLEngine newSslEngine() { + return clientSslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/BaseChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/BaseChainedProxyTest.java index 0c901390..c4850c8b 100644 --- a/src/test/java/org/littleshoot/proxy/BaseChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/BaseChainedProxyTest.java @@ -1,136 +1,134 @@ package org.littleshoot.proxy; -import io.netty.handler.codec.http.HttpRequest; -import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import static org.assertj.core.api.Assertions.assertThat; +import io.netty.handler.codec.http.HttpRequest; +import java.io.IOException; import java.net.InetAddress; import java.net.InetSocketAddress; import java.net.UnknownHostException; import java.util.concurrent.ConcurrentSkipListSet; import java.util.concurrent.atomic.AtomicLong; - -import static org.hamcrest.Matchers.in; -import static org.hamcrest.Matchers.is; -import static org.junit.Assert.assertEquals; -import static org.junit.Assert.assertThat; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; /** - * Base class for tests that test a proxy chained to an upstream proxy. In - * addition to the usual assertions, this also asserts that every request sent - * by the downstream proxy was received by the upstream proxy. + * Base class for tests that test a proxy chained to an upstream proxy. In addition to the usual + * assertions, this also asserts that every request sent by the downstream proxy was received by the + * upstream proxy. */ -public abstract class BaseChainedProxyTest extends BaseProxyTest { - protected final AtomicLong REQUESTS_SENT_BY_DOWNSTREAM = new AtomicLong( - 0L); - protected final AtomicLong REQUESTS_RECEIVED_BY_UPSTREAM = new AtomicLong( - 0L); - protected final ConcurrentSkipListSet TRANSPORTS_USED = new ConcurrentSkipListSet<>(); - - protected final ActivityTracker DOWNSTREAM_TRACKER = new ActivityTrackerAdapter() { +abstract class BaseChainedProxyTest extends BaseProxyTest { + protected final AtomicLong REQUESTS_SENT_BY_DOWNSTREAM = new AtomicLong(0L); + protected final AtomicLong REQUESTS_RECEIVED_BY_UPSTREAM = new AtomicLong(0L); + protected final ConcurrentSkipListSet TRANSPORTS_USED = + new ConcurrentSkipListSet<>(); + + protected final ActivityTracker DOWNSTREAM_TRACKER = + new ActivityTrackerAdapter() { @Override - public void requestSentToServer(FullFlowContext flowContext, - io.netty.handler.codec.http.HttpRequest httpRequest) { - REQUESTS_SENT_BY_DOWNSTREAM.incrementAndGet(); - TRANSPORTS_USED.add(flowContext.getChainedProxy() - .getTransportProtocol()); + public void requestSentToServer( + FullFlowContext flowContext, io.netty.handler.codec.http.HttpRequest httpRequest) { + REQUESTS_SENT_BY_DOWNSTREAM.incrementAndGet(); + TRANSPORTS_USED.add(flowContext.getChainedProxy().getTransportProtocol()); } - }; + }; - protected final ActivityTracker UPSTREAM_TRACKER = new ActivityTrackerAdapter() { + protected final ActivityTracker UPSTREAM_TRACKER = + new ActivityTrackerAdapter() { @Override - public void requestReceivedFromClient(FlowContext flowContext, - HttpRequest httpRequest) { - REQUESTS_RECEIVED_BY_UPSTREAM.incrementAndGet(); + public void requestReceivedFromClient(FlowContext flowContext, HttpRequest httpRequest) { + REQUESTS_RECEIVED_BY_UPSTREAM.incrementAndGet(); } - }; - - protected HttpProxyServer upstreamProxy; - - @Override - protected void setUp() { - REQUESTS_SENT_BY_DOWNSTREAM.set(0); - REQUESTS_RECEIVED_BY_UPSTREAM.set(0); - TRANSPORTS_USED.clear(); - this.upstreamProxy = upstreamProxy().start(); - this.proxyServer = bootstrapProxy() - .withName("Downstream") - .withPort(0) - .withChainProxyManager(chainedProxyManager()) - .plusActivityTracker(DOWNSTREAM_TRACKER).start(); + }; + + protected HttpProxyServer upstreamProxy; + + @Override + protected void setUp() throws IOException { + REQUESTS_SENT_BY_DOWNSTREAM.set(0); + REQUESTS_RECEIVED_BY_UPSTREAM.set(0); + TRANSPORTS_USED.clear(); + upstreamProxy = upstreamProxy().start(); + proxyServer = + bootstrapProxy() + .withName("Downstream") + .withPort(0) + .withChainProxyManager(chainedProxyManager()) + .plusActivityTracker(DOWNSTREAM_TRACKER) + .start(); + } + + protected HttpProxyServerBootstrap upstreamProxy() { + return DefaultHttpProxyServer.bootstrap() + .withName("Upstream") + .withPort(0) + .plusActivityTracker(UPSTREAM_TRACKER); + } + + protected ChainedProxyManager chainedProxyManager() { + return (httpRequest, chainedProxies, clientDetails) -> chainedProxies.add(newChainedProxy()); + } + + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy(); + } + + @Override + protected void tearDown() { + if (upstreamProxy != null) { + upstreamProxy.abort(); } + } - protected HttpProxyServerBootstrap upstreamProxy() { - return DefaultHttpProxyServer.bootstrap() - .withName("Upstream") - .withPort(0) - .plusActivityTracker(UPSTREAM_TRACKER); - } - - protected ChainedProxyManager chainedProxyManager() { - return (httpRequest, chainedProxies, clientDetails) -> chainedProxies.add(newChainedProxy()); + @Override + public void testSimplePostRequest() { + super.testSimplePostRequest(); + if (isChained() && !expectBadGatewayForEverything()) { + assertThatUpstreamProxyReceivedSentRequests(); } + } - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy(); + @Override + public void testSimpleGetRequest() { + super.testSimpleGetRequest(); + if (isChained() && !expectBadGatewayForEverything()) { + assertThatUpstreamProxyReceivedSentRequests(); } + } - @Override - protected void tearDown() { - this.upstreamProxy.abort(); + @Override + public void testProxyWithBadAddress() { + super.testProxyWithBadAddress(); + if (isChained() && !expectBadGatewayForEverything()) { + assertThatUpstreamProxyReceivedSentRequests(); } - - @Override - public void testSimplePostRequest() throws Exception { - super.testSimplePostRequest(); - if (isChained() && !expectBadGatewayForEverything()) { - assertThatUpstreamProxyReceivedSentRequests(); - } - } - + } + + @Override + protected boolean isChained() { + return true; + } + + private void assertThatUpstreamProxyReceivedSentRequests() { + assertThat(REQUESTS_SENT_BY_DOWNSTREAM.get()) + .as("Upstream proxy should have seen every request sent by downstream proxy") + .isEqualTo(REQUESTS_RECEIVED_BY_UPSTREAM.get()); + assertThat(TRANSPORTS_USED) + .as("1 and only 1 transport protocol should have been used to upstream proxy") + .hasSize(1); + assertThat(TRANSPORTS_USED) + .as("Correct transport should have been used") + .contains(newChainedProxy().getTransportProtocol()); + } + + protected class BaseChainedProxy extends ChainedProxyAdapter { @Override - public void testSimpleGetRequest() throws Exception { - super.testSimpleGetRequest(); - if (isChained() && !expectBadGatewayForEverything()) { - assertThatUpstreamProxyReceivedSentRequests(); - } - } - - @Override - public void testProxyWithBadAddress() throws Exception { - super.testProxyWithBadAddress(); - if (isChained() && !expectBadGatewayForEverything()) { - assertThatUpstreamProxyReceivedSentRequests(); - } - } - - @Override - protected boolean isChained() { - return true; - } - - private void assertThatUpstreamProxyReceivedSentRequests() { - assertEquals( - "Upstream proxy should have seen every request sent by downstream proxy", - REQUESTS_SENT_BY_DOWNSTREAM.get(), - REQUESTS_RECEIVED_BY_UPSTREAM.get()); - assertEquals( - "1 and only 1 transport protocol should have been used to upstream proxy", - 1, TRANSPORTS_USED.size()); - assertThat("Correct transport should have been used", - newChainedProxy().getTransportProtocol(), is(in(TRANSPORTS_USED))); - } - - protected class BaseChainedProxy extends ChainedProxyAdapter { - @Override - public InetSocketAddress getChainedProxyAddress() { - try { - return new InetSocketAddress(InetAddress - .getByName("127.0.0.1"), - upstreamProxy.getListenAddress().getPort()); - } catch (UnknownHostException uhe) { - throw new RuntimeException( - "Unable to resolve 127.0.0.1?!"); - } - } + public InetSocketAddress getChainedProxyAddress() { + try { + return new InetSocketAddress( + InetAddress.getByName("127.0.0.1"), upstreamProxy.getListenAddress().getPort()); + } catch (UnknownHostException uhe) { + throw new RuntimeException("Unable to resolve 127.0.0.1?!"); + } } + } } diff --git a/src/test/java/org/littleshoot/proxy/BaseChainedSocksProxyTest.java b/src/test/java/org/littleshoot/proxy/BaseChainedSocksProxyTest.java index b8b32fb0..0dc167fd 100644 --- a/src/test/java/org/littleshoot/proxy/BaseChainedSocksProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/BaseChainedSocksProxyTest.java @@ -1,5 +1,7 @@ package org.littleshoot.proxy; +import static org.assertj.core.api.Assertions.fail; + import io.netty.bootstrap.ServerBootstrap; import io.netty.channel.ChannelFuture; import io.netty.channel.EventLoopGroup; @@ -8,70 +10,72 @@ import io.netty.example.socksproxy.SocksServerInitializer; import io.netty.handler.logging.LogLevel; import io.netty.handler.logging.LoggingHandler; - import java.net.InetSocketAddress; -import static org.junit.Assert.fail; +abstract class BaseChainedSocksProxyTest extends BaseProxyTest { + private EventLoopGroup socksBossGroup; + private EventLoopGroup socksWorkerGroup; + private int socksPort; -abstract public class BaseChainedSocksProxyTest extends BaseProxyTest { - private EventLoopGroup socksBossGroup; - private EventLoopGroup socksWorkerGroup; - private int socksPort; + protected abstract ChainedProxyType getSocksProxyType(); - abstract protected ChainedProxyType getSocksProxyType(); + @Override + protected void setUp() throws Exception { + initializeSocksServer(); + proxyServer = + bootstrapProxy() + .withName("Downstream") + .withPort(0) + .withChainProxyManager(chainedProxyManager()) + .start(); + } - @Override - protected void setUp() throws Exception { - initializeSocksServer(); - this.proxyServer = bootstrapProxy() - .withName("Downstream") - .withPort(0) - .withChainProxyManager(chainedProxyManager()) - .start(); + @Override + protected void tearDown() { + if (socksBossGroup != null) { + socksBossGroup.shutdownGracefully(); } - - @Override - protected void tearDown() { - if (socksBossGroup != null) { - socksBossGroup.shutdownGracefully(); - } - if (socksWorkerGroup != null) { - socksWorkerGroup.shutdownGracefully(); - } + if (socksWorkerGroup != null) { + socksWorkerGroup.shutdownGracefully(); } + } - private void initializeSocksServer() throws Exception { - socksBossGroup = new NioEventLoopGroup(1); - socksWorkerGroup = new NioEventLoopGroup(); + protected void initializeSocksServer() throws Exception { + socksBossGroup = new NioEventLoopGroup(1); + socksWorkerGroup = new NioEventLoopGroup(); - ServerBootstrap bootstrap = new ServerBootstrap(); - bootstrap.group(socksBossGroup, socksWorkerGroup) - .channel(NioServerSocketChannel.class) - .handler(new LoggingHandler(LogLevel.DEBUG)) - .childHandler(new SocksServerInitializer()); + ServerBootstrap bootstrap = new ServerBootstrap(); + bootstrap + .group(socksBossGroup, socksWorkerGroup) + .channel(NioServerSocketChannel.class) + .handler(new LoggingHandler(LogLevel.DEBUG)) + .childHandler(new SocksServerInitializer()); - ChannelFuture channelFuture = bootstrap.bind(0).sync(); - socksPort = ((InetSocketAddress)channelFuture.channel().localAddress()).getPort(); - } + ChannelFuture channelFuture = bootstrap.bind(0).sync(); + socksPort = ((InetSocketAddress) channelFuture.channel().localAddress()).getPort(); + } - private ChainedProxyManager chainedProxyManager() { - return (httpRequest, chainedProxies, details) -> chainedProxies.add(new ChainedProxyAdapter() { - @Override - public InetSocketAddress getChainedProxyAddress() { + protected ChainedProxyManager chainedProxyManager() { + return (httpRequest, chainedProxies, details) -> + chainedProxies.add( + new ChainedProxyAdapter() { + @Override + public InetSocketAddress getChainedProxyAddress() { return new InetSocketAddress("127.0.0.1", socksPort); - } - @Override - public ChainedProxyType getChainedProxyType() { + } + + @Override + public ChainedProxyType getChainedProxyType() { final ChainedProxyType socksProxyType = getSocksProxyType(); switch (socksProxyType) { - case SOCKS4: - case SOCKS5: - return socksProxyType; - default: - fail(socksProxyType + " is not a type of SOCKS proxy"); - throw new UnknownChainedProxyTypeException(socksProxyType); + case SOCKS4: + case SOCKS5: + return socksProxyType; + default: + fail(socksProxyType + " is not a type of SOCKS proxy"); + throw new UnknownChainedProxyTypeException(socksProxyType); } - } - }); - } + } + }); + } } diff --git a/src/test/java/org/littleshoot/proxy/BaseProxyTest.java b/src/test/java/org/littleshoot/proxy/BaseProxyTest.java index 846545c1..ebe84284 100644 --- a/src/test/java/org/littleshoot/proxy/BaseProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/BaseProxyTest.java @@ -1,58 +1,52 @@ package org.littleshoot.proxy; import org.apache.http.HttpHost; -import org.junit.Test; +import org.junit.jupiter.api.Test; /** - * Base for tests that test the proxy. This base class encapsulates all of the - * tests and test conditions. Sub-classes should provide different - * {@link #setUp()} and {@link #tearDown()} methods for testing different - * configurations of the proxy (e.g. single versus chained, tunneling, etc.). + * Base for tests that test the proxy. This base class encapsulates all the tests and test + * conditions. Subclasses should provide different {@link #setUp()} and {@link #tearDown()} methods + * for testing different configurations of the proxy (e.g. single versus chained, tunneling, etc.). */ -public abstract class BaseProxyTest extends AbstractProxyTest { - @Test - public void testSimpleGetRequest() throws Exception { - lastResponse = - compareProxiedAndUnproxiedGET(webHost, DEFAULT_RESOURCE); - } - - @Test - public void testSimpleGetRequestOverHTTPS() throws Exception { - lastResponse = - compareProxiedAndUnproxiedGET(httpsWebHost, DEFAULT_RESOURCE); - } - - @Test - public void testSimplePostRequest() throws Exception { - lastResponse = - compareProxiedAndUnproxiedPOST(webHost, DEFAULT_RESOURCE); - } - - @Test - public void testSimplePostRequestOverHTTPS() throws Exception { - lastResponse = - compareProxiedAndUnproxiedPOST(httpsWebHost, DEFAULT_RESOURCE); - } - - /** - * This test tests a HEAD followed by a GET for the same resource, making - * sure that the requests complete and that the Content-Length matches. - */ - @Test - public void testHeadRequestFollowedByGet() throws Exception { - httpGetWithApacheClient(webHost, DEFAULT_RESOURCE, true, true); - } - - @Test - public void testProxyWithBadAddress() - throws Exception { - // This test used to try connecting to "test.localhost" and that worked for for local builds, but resulted in - // the wrong error (405 instead of 502) on the build server due to nginx. So, switched it to localhost:17, - // which should work as long as there's not a web server running on the QOTD port. - ResponseInfo response = - httpPostWithApacheClient(new HttpHost("localhost", 17), - DEFAULT_RESOURCE, true); - assertReceivedBadGateway(response); - } - +abstract class BaseProxyTest extends AbstractProxyTest { + @Test + public void testSimpleGetRequest() { + lastResponse = compareProxiedAndUnproxiedGET(webHost, DEFAULT_RESOURCE); + } + + @Test + public void testSimpleGetRequestOverHTTPS() { + lastResponse = compareProxiedAndUnproxiedGET(httpsWebHost, DEFAULT_RESOURCE); + } + + @Test + public void testSimplePostRequest() { + lastResponse = compareProxiedAndUnproxiedPOST(webHost, DEFAULT_RESOURCE); + } + + @Test + public void testSimplePostRequestOverHTTPS() { + lastResponse = compareProxiedAndUnproxiedPOST(httpsWebHost, DEFAULT_RESOURCE); + } + + /** + * This test tests a HEAD followed by a GET for the same resource, making sure that the requests + * complete and that the Content-Length matches. + */ + @Test + public void testHeadRequestFollowedByGet() { + httpGetWithApacheClient(webHost, DEFAULT_RESOURCE, true, true); + } + + @Test + public void testProxyWithBadAddress() { + // This test used to try connecting to "test.localhost" and that worked for local builds, but + // resulted in + // the wrong error (405 instead of 502) on the build server due to nginx. So, switched it to + // localhost:17, + // which should work as long as there's not a web server running on the QOTD port. + ResponseInfo response = + httpPostWithApacheClient(new HttpHost("localhost", 17), DEFAULT_RESOURCE, true); + assertReceivedBadGateway(response); + } } diff --git a/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackTest.java b/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackTest.java index 0edb6db1..3b8b5148 100644 --- a/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackTest.java +++ b/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackTest.java @@ -1,6 +1,6 @@ package org.littleshoot.proxy; -import org.junit.Assert; +import static org.assertj.core.api.Assertions.assertThat; import java.net.InetAddress; import java.net.InetSocketAddress; @@ -8,50 +8,49 @@ import java.util.concurrent.atomic.AtomicBoolean; /** - * Tests a proxy chained to a missing downstream proxy. When the downstream - * proxy is unavailable, the downstream proxy should just fall back to a direct - * connection. + * Tests a proxy chained to a missing downstream proxy. When the downstream proxy is unavailable, + * the downstream proxy should just fall back to a direct connection. */ -public class ChainedProxyWithFallbackTest extends BaseProxyTest { - private AtomicBoolean unableToConnect = new AtomicBoolean(false); - - @Override - protected void setUp() { - unableToConnect.set(false); - this.proxyServer = bootstrapProxy() - .withName("Downstream") - .withPort(0) - .withChainProxyManager((httpRequest, chainedProxies, clientDetails) -> { - chainedProxies.add(new ChainedProxyAdapter() { +public final class ChainedProxyWithFallbackTest extends BaseProxyTest { + private final AtomicBoolean unableToConnect = new AtomicBoolean(false); + + @Override + protected void setUp() { + unableToConnect.set(false); + proxyServer = + bootstrapProxy() + .withName("Downstream") + .withPort(0) + .withChainProxyManager( + (httpRequest, chainedProxies, clientDetails) -> { + chainedProxies.add( + new ChainedProxyAdapter() { @Override public InetSocketAddress getChainedProxyAddress() { - try { - // using unconnectable port 0 - return new InetSocketAddress(InetAddress.getByName("127.0.0.1"), 0); - } catch (UnknownHostException uhe) { - throw new RuntimeException( - "Unable to resolve 127.0.0.1?!"); - } + try { + // using unconnectable port 0 + return new InetSocketAddress(InetAddress.getByName("127.0.0.1"), 0); + } catch (UnknownHostException uhe) { + throw new RuntimeException("Unable to resolve 127.0.0.1?!"); + } } @Override public void connectionFailed(Throwable cause) { - unableToConnect.set(true); + unableToConnect.set(true); } + }); - }); - - chainedProxies - .add(ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION); + chainedProxies.add(ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION); }) - .start(); - } - - @Override - protected void tearDown() throws Exception { - super.tearDown(); - Assert.assertTrue( - "We should have been told that we were unable to connect", - unableToConnect.get()); - } + .start(); + } + + @Override + protected void tearDown() throws Exception { + super.tearDown(); + assertThat(unableToConnect.get()) + .as("We should have been told that we were unable to connect") + .isTrue(); + } } diff --git a/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToDirectDueToSSLTest.java b/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToDirectDueToSSLTest.java index 62a8377a..69ab9566 100644 --- a/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToDirectDueToSSLTest.java +++ b/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToDirectDueToSSLTest.java @@ -1,30 +1,28 @@ package org.littleshoot.proxy; /** - * Tests a proxy chained to a downstream proxy with an untrusted SSL cert. When - * the downstream proxy is unavailable, the downstream proxy should just fall - * back to a direct connection. + * Tests a proxy chained to a downstream proxy with an untrusted SSL cert. When the downstream proxy + * is unavailable, the downstream proxy should just fall back to a direct connection. */ -public class ChainedProxyWithFallbackToDirectDueToSSLTest extends - BadServerAuthenticationTCPChainedProxyTest { - @Override - protected boolean isChained() { - // Set this to false since we don't actually expect anything to go - // through the chained proxy - return false; - } +public final class ChainedProxyWithFallbackToDirectDueToSSLTest + extends BadServerAuthenticationTCPChainedProxyTest { + @Override + protected boolean isChained() { + // Set this to false since we don't actually expect anything to go + // through the chained proxy + return false; + } - @Override - protected boolean expectBadGatewayForEverything() { - return false; - } + @Override + protected boolean expectBadGatewayForEverything() { + return false; + } - protected ChainedProxyManager chainedProxyManager() { - return (httpRequest, chainedProxies, clientDetails) -> { - // This first one has a bad cert - chainedProxies.add(newChainedProxy()); - chainedProxies - .add(ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION); - }; - } + protected ChainedProxyManager chainedProxyManager() { + return (httpRequest, chainedProxies, clientDetails) -> { + // This first one has a bad cert + chainedProxies.add(newChainedProxy()); + chainedProxies.add(ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION); + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToOtherChainedProxyDueToSSLTest.java b/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToOtherChainedProxyDueToSSLTest.java index 835fb7ae..c3b99071 100644 --- a/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToOtherChainedProxyDueToSSLTest.java +++ b/src/test/java/org/littleshoot/proxy/ChainedProxyWithFallbackToOtherChainedProxyDueToSSLTest.java @@ -3,38 +3,33 @@ import javax.net.ssl.SSLEngine; /** - * Tests a proxy chained to a downstream proxy with an untrusted SSL cert. When - * the downstream proxy is unavailable, the downstream proxy should just fall - * back to a the next chained proxy. + * Tests a proxy chained to a downstream proxy with an untrusted SSL cert. When the downstream proxy + * is unavailable, the downstream proxy should just fall back to the next chained proxy. */ -public class ChainedProxyWithFallbackToOtherChainedProxyDueToSSLTest extends - BadServerAuthenticationTCPChainedProxyTest { - @Override - protected boolean expectBadGatewayForEverything() { - return false; - } +public final class ChainedProxyWithFallbackToOtherChainedProxyDueToSSLTest + extends BadServerAuthenticationTCPChainedProxyTest { + @Override + protected boolean expectBadGatewayForEverything() { + return false; + } - protected ChainedProxyManager chainedProxyManager() { - return (httpRequest, chainedProxies, clientDetails) -> { - // This first one has a bad cert - chainedProxies.add(newChainedProxy()); - // This 2nd one should work - chainedProxies.add(new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } + protected ChainedProxyManager chainedProxyManager() { + return (httpRequest, chainedProxies, clientDetails) -> { + // This first one has a bad cert + chainedProxies.add(newChainedProxy()); + // This 2nd one should work + chainedProxies.add( + new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { + return true; + } - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return serverSslEngineSource.newSslEngine(); - } - }); - }; - } + @Override + public SSLEngine newSslEngine() { + return serverSslEngineSource.newSslEngine(); + } + }); + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/ClientAuthenticationNotRequiredTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/ClientAuthenticationNotRequiredTCPChainedProxyTest.java index 0a80d649..05d0e36c 100644 --- a/src/test/java/org/littleshoot/proxy/ClientAuthenticationNotRequiredTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/ClientAuthenticationNotRequiredTCPChainedProxyTest.java @@ -1,48 +1,37 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; - -import static org.littleshoot.proxy.TransportProtocol.TCP; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; /** - * Tests that when client authentication is not required, it doesn't matter what - * certs the client sends. + * Tests that when client authentication is not required, it doesn't matter what certs the client + * sends. */ -public class ClientAuthenticationNotRequiredTCPChainedProxyTest extends - BaseChainedProxyTest { - private final SslEngineSource serverSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - private final SslEngineSource clientSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks", false, false); - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(serverSslEngineSource) - .withAuthenticateSslClients(false); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return clientSslEngineSource.newSslEngine(); - } - }; - } +public final class ClientAuthenticationNotRequiredTCPChainedProxyTest extends BaseChainedProxyTest { + private final SslEngineSource serverSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks"); + private final SslEngineSource clientSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks", false, false); + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy() + .withSslEngineSource(serverSslEngineSource) + .withAuthenticateSslClients(false); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { + return true; + } + + @Override + public SSLEngine newSslEngine() { + return clientSslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/ClonedProxyTest.java b/src/test/java/org/littleshoot/proxy/ClonedProxyTest.java index 85571946..df1b9560 100644 --- a/src/test/java/org/littleshoot/proxy/ClonedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/ClonedProxyTest.java @@ -1,120 +1,99 @@ package org.littleshoot.proxy; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.test.HttpClientUtil.performLocalHttpGet; + +import com.github.tomakehurst.wiremock.WireMockServer; import org.apache.http.HttpResponse; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import org.littleshoot.proxy.test.HttpClientUtil; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; - -import static org.junit.Assert.assertEquals; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; - -public class ClonedProxyTest { - private ClientAndServer mockServer; - private int mockServerPort; - - private HttpProxyServer originalProxy; - private HttpProxyServer clonedProxy; - @Before - public void setUp() { - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); - } - - @After - public void tearDown() { - try { - if (mockServer != null) { - mockServer.stop(); - } - } finally { - try { - if (originalProxy != null) { - originalProxy.abort(); - } - } finally { - if (clonedProxy != null) { - clonedProxy.abort(); - } - } +public final class ClonedProxyTest { + private WireMockServer mockServer; + private int mockServerPort; + + private HttpProxyServer originalProxy; + private HttpProxyServer clonedProxy; + + @BeforeEach + void setUp() { + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + } + + @AfterEach + void tearDown() { + try { + if (mockServer != null) { + mockServer.stop(); + } + } finally { + try { + if (originalProxy != null) { + originalProxy.abort(); } + } finally { + if (clonedProxy != null) { + clonedProxy.abort(); + } + } } - - @Test - public void testClonedProxyHandlesRequests() { - originalProxy = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withName("original") - .start(); - clonedProxy = originalProxy.clone() - .withName("clone") - .start(); - - mockServer.when(request() - .withMethod("GET") - .withPath("/testClonedProxyHandlesRequests"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("success") - ); - - HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + mockServerPort + "/testClonedProxyHandlesRequests", clonedProxy); - assertEquals("Expected to receive a 200 when making a request using the cloned proxy server", 200, response.getStatusLine().getStatusCode()); - } - - @Test - public void testStopClonedProxyDoesNotStopOriginalServer() { - originalProxy = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withName("original") - .start(); - clonedProxy = originalProxy.clone() - .withName("clone") - .start(); - - clonedProxy.abort(); - - mockServer.when(request() - .withMethod("GET") - .withPath("/testClonedProxyHandlesRequests"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("success") - ); - - HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + mockServerPort + "/testClonedProxyHandlesRequests", originalProxy); - assertEquals("Expected to receive a 200 when making a request using the cloned proxy server", 200, response.getStatusLine().getStatusCode()); - } - - @Test - public void testStopOriginalServerDoesNotStopClonedServer() { - originalProxy = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withName("original") - .start(); - clonedProxy = originalProxy.clone() - .withName("clone") - .start(); - - originalProxy.abort(); - - mockServer.when(request() - .withMethod("GET") - .withPath("/testClonedProxyHandlesRequests"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("success") - ); - - HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + mockServerPort + "/testClonedProxyHandlesRequests", clonedProxy); - assertEquals("Expected to receive a 200 when making a request using the cloned proxy server", 200, response.getStatusLine().getStatusCode()); - } + } + + @Test + public void testClonedProxyHandlesRequests() { + originalProxy = DefaultHttpProxyServer.bootstrap().withPort(0).withName("original").start(); + clonedProxy = originalProxy.clone().withName("clone").start(); + + mockServer.stubFor( + get(urlEqualTo("/testClonedProxyHandlesRequests")) + .willReturn(aResponse().withStatus(200).withBody("success"))); + + HttpResponse response = + performLocalHttpGet(mockServerPort, "/testClonedProxyHandlesRequests", clonedProxy); + assertThat(response.getStatusLine().getStatusCode()) + .as("Expected to receive a 200 when making a request using the cloned proxy server") + .isEqualTo(200); + } + + @Test + public void testStopClonedProxyDoesNotStopOriginalServer() { + originalProxy = DefaultHttpProxyServer.bootstrap().withPort(0).withName("original").start(); + clonedProxy = originalProxy.clone().withName("clone").start(); + + clonedProxy.abort(); + + mockServer.stubFor( + get(urlEqualTo("/testClonedProxyHandlesRequests")) + .willReturn(aResponse().withStatus(200).withBody("success"))); + + HttpResponse response = + performLocalHttpGet(mockServerPort, "/testClonedProxyHandlesRequests", originalProxy); + assertThat(response.getStatusLine().getStatusCode()) + .as("Expected to receive a 200 when making a request using the cloned proxy server") + .isEqualTo(200); + } + + @Test + public void testStopOriginalServerDoesNotStopClonedServer() { + originalProxy = DefaultHttpProxyServer.bootstrap().withPort(0).withName("original").start(); + clonedProxy = originalProxy.clone().withName("clone").start(); + + originalProxy.abort(); + + mockServer.stubFor( + get(urlEqualTo("/testClonedProxyHandlesRequests")) + .willReturn(aResponse().withStatus(200).withBody("success"))); + + HttpResponse response = + performLocalHttpGet(mockServerPort, "/testClonedProxyHandlesRequests", clonedProxy); + assertThat(response.getStatusLine().getStatusCode()) + .as("Expected to receive a 200 when making a request using the cloned proxy server") + .isEqualTo(200); + } } diff --git a/src/test/java/org/littleshoot/proxy/ConnectResponseFiltersTest.java b/src/test/java/org/littleshoot/proxy/ConnectResponseFiltersTest.java new file mode 100644 index 00000000..907963cc --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/ConnectResponseFiltersTest.java @@ -0,0 +1,239 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.buildHttpClient; + +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpObject; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import org.apache.http.client.methods.HttpGet; +import org.apache.http.impl.client.CloseableHttpClient; +import org.eclipse.jetty.server.Server; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.extras.TestMitmManager; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * Integration test to reproduce issue #77 + * + *

Issue: CONNECT Response Not Returned to HttpFilters + * + *

When using LittleProxy with MITM enabled and a custom HttpFilters, the CONNECT request goes + * through proxyToServerRequest() but the CONNECT response is never sent to serverToProxyResponse(). + * This is because when isConnecting() is true (during CONNECT tunnel establishment), the response + * goes to connectionFlow.read() instead of being passed to the filters. + */ +@NullMarked +class ConnectResponseFiltersTest { + private static final String DEFAULT_JKS_KEYSTORE_PATH = "target/littleproxy_keystore.jks"; + private static final Logger logger = LoggerFactory.getLogger(ConnectResponseFiltersTest.class); + private @Nullable Server webServer; + private @Nullable HttpProxyServer proxyServer; + private int httpsWebServerPort; + + @BeforeEach + void setUp() throws Exception { + // Start web server with both HTTP and HTTPS enabled + webServer = TestUtils.startWebServer(true, DEFAULT_JKS_KEYSTORE_PATH); + httpsWebServerPort = TestUtils.findLocalHttpsPort(webServer); + } + + @AfterEach + void tearDown() throws Exception { + try { + if (webServer != null) { + webServer.stop(); + } + } finally { + try { + if (proxyServer != null) { + proxyServer.abort(); + } + } catch (Exception e) { + // ignore + } + } + } + + /** + * Test that verifies the CONNECT response is returned to HttpFilters. + * + *

This test should FAIL if the bug exists (CONNECT response not going to filters) and PASS + * after the fix is applied. + */ + @Test + public void testConnectResponseIsReturnedToFilters() throws Exception { + // Track whether the CONNECT request was seen in proxyToServerRequest + final AtomicBoolean connectRequestSeen = new AtomicBoolean(false); + // Track whether the CONNECT response (200) was seen in serverToProxyResponse + final AtomicBoolean connectResponseSeen = new AtomicBoolean(false); + // Track how many 200 responses we've seen in serverToProxyResponse + final AtomicInteger connectResponseCount = new AtomicInteger(0); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @Override + public HttpFilters filterRequest(HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse proxyToServerRequest(HttpObject httpObject) { + if (httpObject instanceof HttpRequest request) { + if (HttpMethod.CONNECT.equals(request.method())) { + connectRequestSeen.set(true); + } + } + return null; + } + + @Override + public HttpObject serverToProxyResponse(HttpObject httpObject) { + if (httpObject instanceof HttpResponse response) { + // Check if this is a 200 response (CONNECT responses have 200 status) + if (response.status().code() == 200) { + connectResponseCount.incrementAndGet(); + // If we see a 200 response, it's likely the CONNECT tunnel response + // (since the actual GET request will also get a 200, but we'll have already + // seen the CONNECT request in proxyToServerRequest) + if (connectRequestSeen.get()) { + connectResponseSeen.set(true); + } + } + } + return httpObject; + } + }; + } + }; + + // Start proxy with MITM enabled (this is what triggers CONNECT tunneling) + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource(filtersSource) + .withManInTheMiddle(new TestMitmManager()) + .start(); + + // Give the proxy time to start + Thread.sleep(500); + + // Make an HTTPS request through the proxy using a client that trusts self-signed certs + String httpsUrl = "https://localhost:" + httpsWebServerPort + "/"; + CloseableHttpClient httpClient = + buildHttpClient(true, true, proxyServer.getListenAddress().getPort(), null, null); + HttpGet get = new HttpGet(httpsUrl); + org.apache.http.HttpResponse response = httpClient.execute(get); + httpClient.close(); + + // Wait for filters to be invoked + Thread.sleep(1000); + + // Verify the CONNECT request was seen in proxyToServerRequest + assertThat(connectRequestSeen.get()) + .as("CONNECT request should be seen in proxyToServerRequest filter") + .isTrue(); + + // This assertion should FAIL if the bug exists (issue #77) + // The CONNECT response should be returned to serverToProxyResponse + assertThat(connectResponseSeen.get()) + .as( + "CONNECT response (200) should be seen in serverToProxyResponse filter. " + + "This is the bug described in issue #77 - the CONNECT response is not being " + + "passed to HttpFilters when isConnecting() is true.") + .isTrue(); + + // Verify we got at least one 200 response (should be 2: CONNECT tunnel + GET response) + assertThat(connectResponseCount.get()) + .as("Should have received at least one 200 response in serverToProxyResponse") + .isGreaterThanOrEqualTo(1); + + // Verify the request actually succeeded + assertThat(response.getStatusLine().getStatusCode()) + .as("HTTPS request should succeed") + .isEqualTo(200); + } + + /** + * Additional test to verify that both the CONNECT tunnel establishment AND the subsequent HTTP + * request over the tunnel go through the filters correctly. + */ + @Test + public void testConnectResponseAndSubsequentRequestBothFiltered() throws Exception { + final StringBuilder filterCalls = new StringBuilder(); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @Override + public HttpFilters filterRequest(HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse proxyToServerRequest(HttpObject httpObject) { + if (httpObject instanceof HttpRequest request) { + filterCalls.append("proxyToServerRequest:").append(request.method()).append(","); + } + return null; + } + + @Override + public HttpObject serverToProxyResponse(HttpObject httpObject) { + if (httpObject instanceof HttpResponse response) { + filterCalls + .append("serverToProxyResponse:") + .append(response.status().code()) + .append("(originalMethod:") + .append(originalRequest.method()) + .append("),"); + } + return httpObject; + } + }; + } + }; + + // Start proxy with MITM enabled + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource(filtersSource) + .withManInTheMiddle(new TestMitmManager()) + .start(); + + // Give the proxy time to start + Thread.sleep(500); + + // Make an HTTPS request through the proxy using a client that trusts self-signed certs + String httpsUrl = "https://localhost:" + httpsWebServerPort + "/"; + CloseableHttpClient httpClient = + buildHttpClient(true, true, proxyServer.getListenAddress().getPort(), null, null); + HttpGet get = new HttpGet(httpsUrl); + httpClient.execute(get); + httpClient.close(); + + // Wait for filters to be invoked + Thread.sleep(1000); + + // Print what we captured for debugging + logger.info("Filter calls captured: {}", filterCalls); + + // Verify that both CONNECT (for tunnel) and GET (for actual request) are seen + String filterCallsAsString = filterCalls.toString(); + assertThat(filterCallsAsString) + .as("Both CONNECT and GET should appear in filter calls") + .contains("CONNECT"); + assertThat(filterCallsAsString) + .as( + "Both CONNECT response (200) and GET response (200) should appear in serverToProxyResponse") + .contains("serverToProxyResponse:200"); + } +} diff --git a/src/test/java/org/littleshoot/proxy/DefaultHostResolverTest.java b/src/test/java/org/littleshoot/proxy/DefaultHostResolverTest.java new file mode 100644 index 00000000..b9175da2 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/DefaultHostResolverTest.java @@ -0,0 +1,46 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.net.InetSocketAddress; +import java.net.UnknownHostException; +import org.junit.jupiter.api.Test; + +final class DefaultHostResolverTest { + private final DefaultHostResolver resolver = new DefaultHostResolver(); + + @Test + void resolveLocalhost() throws UnknownHostException { + InetSocketAddress address = resolver.resolve("localhost", 8080); + + assertThat(address).isNotNull(); + assertThat(address.getPort()).isEqualTo(8080); + assertThat(address.getAddress()).isNotNull(); + assertThat(address.getAddress().isLoopbackAddress()) + .as("localhost should resolve to 127.0.0.1") + .isTrue(); + } + + @Test + void resolveWithIpAddress() throws UnknownHostException { + InetSocketAddress address = resolver.resolve("127.0.0.1", 9090); + + assertThat(address).isNotNull(); + assertThat(address.getPort()).isEqualTo(9090); + assertThat(address.getAddress().getHostAddress()).isEqualTo("127.0.0.1"); + } + + @Test + void resolveUnknownHost() { + assertThatThrownBy(() -> resolver.resolve("this-host.invalid", 80)) + .isInstanceOf(UnknownHostException.class); + } + + @Test + void resolveReturnsCorrectPort() throws UnknownHostException { + InetSocketAddress address = resolver.resolve("localhost", 443); + + assertThat(address.getPort()).isEqualTo(443); + } +} diff --git a/src/test/java/org/littleshoot/proxy/DirectRequestTest.java b/src/test/java/org/littleshoot/proxy/DirectRequestTest.java index 5b93546b..d52812fb 100644 --- a/src/test/java/org/littleshoot/proxy/DirectRequestTest.java +++ b/src/test/java/org/littleshoot/proxy/DirectRequestTest.java @@ -1,143 +1,156 @@ package org.littleshoot.proxy; -import io.netty.handler.codec.http.*; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; -import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import org.littleshoot.proxy.test.HttpClientUtil; - -import javax.net.ssl.SSLException; +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.littleshoot.proxy.test.HttpClientUtil.performHttpGet; +import static org.littleshoot.proxy.test.HttpClientUtil.performLocalHttpGet; + +import io.netty.handler.codec.http.DefaultHttpResponse; +import io.netty.handler.codec.http.HttpHeaderNames; +import io.netty.handler.codec.http.HttpObject; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpVersion; import java.util.concurrent.atomic.AtomicBoolean; +import javax.net.ssl.SSLException; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import static org.hamcrest.Matchers.instanceOf; -import static org.junit.Assert.*; - -/** - * This class tests direct requests to the proxy server, which causes endless - * loops (#205). - */ -public class DirectRequestTest { - - private HttpProxyServer proxyServer; - - @Before - public void setUp() { - proxyServer = null; - } - - @After - public void tearDown() { - if (proxyServer != null) { - proxyServer.abort(); - } - } - - @Test(timeout = 5000) - public void testAnswerBadRequestInsteadOfEndlessLoop() { - - startProxyServer(); - - int proxyPort = proxyServer.getListenAddress().getPort(); - org.apache.http.HttpResponse response = HttpClientUtil.performHttpGet("http://127.0.0.1:" + proxyPort + "/directToProxy", proxyServer); - int statusCode = response.getStatusLine().getStatusCode(); - - assertEquals("Expected to receive an HTTP 400 from the server", 400, statusCode); - } - - @Test(timeout = 5000) - public void testAnswerFromFilterShouldBeServed() { - - startProxyServerWithFilterAnsweringStatusCode(403); +/** This class tests direct requests to the proxy server, which causes endless loops (#205). */ +@NullMarked +public final class DirectRequestTest { - int proxyPort = proxyServer.getListenAddress().getPort(); - org.apache.http.HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + proxyPort + "/directToProxy", proxyServer); - int statusCode = response.getStatusLine().getStatusCode(); + @Nullable private HttpProxyServer proxyServer; - assertEquals("Expected to receive an HTTP 403 from the server", 403, statusCode); + @AfterEach + void tearDown() { + if (proxyServer != null) { + proxyServer.abort(); } - - private void startProxyServerWithFilterAnsweringStatusCode(int statusCode) { - final HttpResponseStatus status = HttpResponseStatus.valueOf(statusCode); - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - return new DefaultHttpResponse(HttpVersion.HTTP_1_1, status); - } - }; - } + } + + @Test + @Timeout(5) + public void testAnswerBadRequestInsteadOfEndlessLoop() { + + HttpProxyServer proxyServer = startProxyServer(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + org.apache.http.HttpResponse response = + performHttpGet("http://127.0.0.1:" + proxyPort + "/directToProxy", proxyServer); + int statusCode = response.getStatusLine().getStatusCode(); + + assertThat(statusCode).as("Expected to receive an HTTP 400 from the server").isEqualTo(400); + } + + @Test + @Timeout(5) + public void testAnswerFromFilterShouldBeServed() { + + HttpProxyServer proxyServer = startProxyServerWithFilterAnsweringStatusCode(403); + + int proxyPort = proxyServer.getListenAddress().getPort(); + org.apache.http.HttpResponse response = + performLocalHttpGet(proxyPort, "/directToProxy", proxyServer); + int statusCode = response.getStatusLine().getStatusCode(); + + assertThat(statusCode).as("Expected to receive an HTTP 403 from the server").isEqualTo(403); + } + + private HttpProxyServer startProxyServerWithFilterAnsweringStatusCode(int statusCode) { + final HttpResponseStatus status = HttpResponseStatus.valueOf(statusCode); + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @Override + public HttpFilters filterRequest(HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Override + public HttpResponse clientToProxyRequest(HttpObject httpObject) { + return new DefaultHttpResponse(HttpVersion.HTTP_1_1, status); + } + }; + } }; - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .start(); - } - - @Test(timeout = 5000) - public void testHttpsShouldCancelConnection() { - startProxyServer(); - - int proxyPort = proxyServer.getListenAddress().getPort(); - - - try { - HttpClientUtil.performHttpGet("https://localhost:" + proxyPort + "/directToProxy", proxyServer); - } catch (RuntimeException e) { - Throwable cause = e.getCause(); - assertThat("Expected an SSL exception when attempting to perform an HTTPS GET directly to the proxy", cause, instanceOf(SSLException.class)); - } - } - - @Test(timeout = 5000) - public void testAllowRequestToOriginServerWithOverride() { - // verify that the filter is hit twice: first, on the request from the client, without a Via header; and second, when the proxy - // forwards the request to itself - final AtomicBoolean receivedRequestWithoutVia = new AtomicBoolean(); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withAllowRequestToOriginServer(true) - .withProxyAlias("testAllowRequestToOriginServerWithOverride") - .withFiltersSource(new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - HttpRequest request = (HttpRequest) httpObject; - String viaHeader = request.headers().get(HttpHeaderNames.VIA); - if (viaHeader != null && viaHeader.contains("testAllowRequestToOriginServerWithOverride")) { - return new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.NO_CONTENT); - } else { - receivedRequestWithoutVia.set(true); - } - } - return null; - } - }; - } + proxyServer = + DefaultHttpProxyServer.bootstrap().withPort(0).withFiltersSource(filtersSource).start(); + return proxyServer; + } + + @Test + @Timeout(5) + public void testHttpsShouldCancelConnection() { + HttpProxyServer proxyServer = startProxyServer(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + + assertThatThrownBy( + () -> performHttpGet("https://localhost:" + proxyPort + "/directToProxy", proxyServer)) + .isInstanceOf(RuntimeException.class) + .cause() + .as( + "Expected an SSL exception when attempting to perform an HTTPS GET directly to the proxy") + .isInstanceOf(SSLException.class); + } + + @Test + @Timeout(5) + public void testAllowRequestToOriginServerWithOverride() { + // verify that the filter is hit twice: first, on the request from the client, without a Via + // header; and second, when the proxy + // forwards the request to itself + final AtomicBoolean receivedRequestWithoutVia = new AtomicBoolean(); + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withAllowRequestToOriginServer(true) + .withProxyAlias("testAllowRequestToOriginServerWithOverride") + .withFiltersSource( + new HttpFiltersSourceAdapter() { + @Override + public HttpFilters filterRequest(HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse clientToProxyRequest(HttpObject httpObject) { + if (httpObject instanceof HttpRequest request) { + String viaHeader = request.headers().get(HttpHeaderNames.VIA); + if (viaHeader != null + && viaHeader.contains("testAllowRequestToOriginServerWithOverride")) { + return new DefaultHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.NO_CONTENT); + } else { + receivedRequestWithoutVia.set(true); + } + } + return null; + } + }; + } }) - .start(); + .start(); - int proxyPort = proxyServer.getListenAddress().getPort(); + int proxyPort = proxyServer.getListenAddress().getPort(); - org.apache.http.HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + proxyPort + "/originrequest", proxyServer); - int statusCode = response.getStatusLine().getStatusCode(); + org.apache.http.HttpResponse response = + performLocalHttpGet(proxyPort, "/originrequest", proxyServer); + int statusCode = response.getStatusLine().getStatusCode(); - assertEquals("Expected to receive a 204 response from the filter", 204, statusCode); + assertThat(statusCode).as("Expected to receive a 204 response from the filter").isEqualTo(204); - assertTrue("Expected to receive a request from the client without a Via header", receivedRequestWithoutVia.get()); - } - - private void startProxyServer() { - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - } + assertThat(receivedRequestWithoutVia.get()) + .as("Expected to receive a request from the client without a Via header") + .isTrue(); + } + private HttpProxyServer startProxyServer() { + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + return proxyServer; + } } diff --git a/src/test/java/org/littleshoot/proxy/EncryptedTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/EncryptedTCPChainedProxyTest.java index 7c164523..b8d74de5 100644 --- a/src/test/java/org/littleshoot/proxy/EncryptedTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/EncryptedTCPChainedProxyTest.java @@ -1,39 +1,32 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; +import org.junit.jupiter.api.parallel.Execution; +import org.junit.jupiter.api.parallel.ExecutionMode; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -public class EncryptedTCPChainedProxyTest extends BaseChainedProxyTest { - private final SslEngineSource sslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(sslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return sslEngineSource.newSslEngine(); - } - }; - } +@Execution(ExecutionMode.SAME_THREAD) +public final class EncryptedTCPChainedProxyTest extends BaseChainedProxyTest { + private final SslEngineSource sslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks"); + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(sslEngineSource); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { + return true; + } + + @Override + public SSLEngine newSslEngine() { + return sslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/EncryptedUDTChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/EncryptedUDTChainedProxyTest.java deleted file mode 100644 index f657629d..00000000 --- a/src/test/java/org/littleshoot/proxy/EncryptedUDTChainedProxyTest.java +++ /dev/null @@ -1,39 +0,0 @@ -package org.littleshoot.proxy; - -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - -import javax.net.ssl.SSLEngine; - -import static org.littleshoot.proxy.TransportProtocol.UDT; - -public class EncryptedUDTChainedProxyTest extends BaseChainedProxyTest { - private final SslEngineSource sslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(UDT) - .withSslEngineSource(sslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.UDT; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return sslEngineSource.newSslEngine(); - } - }; - } -} diff --git a/src/test/java/org/littleshoot/proxy/EndToEndStoppingTest.java b/src/test/java/org/littleshoot/proxy/EndToEndStoppingTest.java index 68aaf101..40208d46 100644 --- a/src/test/java/org/littleshoot/proxy/EndToEndStoppingTest.java +++ b/src/test/java/org/littleshoot/proxy/EndToEndStoppingTest.java @@ -1,196 +1,213 @@ package org.littleshoot.proxy; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static java.time.Duration.ofSeconds; +import static java.util.Locale.ROOT; +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assumptions.assumeThat; +import static org.littleshoot.proxy.TestUtils.createProxiedHttpClient; + +import com.github.tomakehurst.wiremock.WireMockServer; import io.netty.handler.codec.http.HttpObject; import io.netty.handler.codec.http.HttpRequest; +import java.nio.charset.StandardCharsets; import org.apache.commons.io.IOUtils; import org.apache.http.HttpEntity; import org.apache.http.HttpResponse; -import org.apache.http.client.HttpClient; import org.apache.http.client.methods.HttpGet; +import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.util.EntityUtils; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; +import org.jspecify.annotations.NonNull; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; +import org.littleshoot.proxy.test.EnableThreadDump; import org.openqa.selenium.Proxy; import org.openqa.selenium.WebDriver; +import org.openqa.selenium.chrome.ChromeDriver; +import org.openqa.selenium.chrome.ChromeOptions; import org.openqa.selenium.firefox.FirefoxDriver; import org.openqa.selenium.firefox.FirefoxOptions; import org.openqa.selenium.remote.CapabilityType; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import java.nio.charset.StandardCharsets; -import java.util.concurrent.TimeUnit; - -import static org.hamcrest.Matchers.greaterThan; -import static org.junit.Assert.assertEquals; -import static org.junit.Assert.assertThat; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; - /** - * End to end test making sure the proxy is able to service simple HTTP requests - * and stop at the end. Made into a unit test from isopov and nasis's - * contributions at: https://github.com/adamfisk/LittleProxy/issues/36 + * End-to-end test making sure the proxy is able to service simple HTTP requests and stop at the + * end. Made into a unit test from isopov and nasis's contributions at: ... */ -public class EndToEndStoppingTest { - private static final Logger log = LoggerFactory.getLogger(EndToEndStoppingTest.class); - - private ClientAndServer mockServer; - private int mockServerPort; - - @Before - public void setUp() { - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); - } - - @After - public void tearDown() { - if (mockServer != null) { - mockServer.stop(); - } - } - - /** - * This is a quick test from nasis that exhibits different behavior from - * unit tests because unit tests call System.exit(). The stop method should - * stop all non-daemon threads and should cause the JVM to exit without - * explicitly calling System.exit(), which running as an application - * properly tests. - */ - public static void main(final String[] args) { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - - Proxy proxy = new Proxy(); - proxy.setProxyType(Proxy.ProxyType.MANUAL); - String proxyStr = String.format("localhost:%d", proxyServer.getListenAddress().getPort()); - proxy.setHttpProxy(proxyStr); - proxy.setSslProxy(proxyStr); - - FirefoxOptions capability = new FirefoxOptions(); - capability.setCapability(CapabilityType.PROXY, proxy); - - String urlString = "http://www.yahoo.com/"; - WebDriver driver = new FirefoxDriver(capability); - driver.manage().timeouts().pageLoadTimeout(30, TimeUnit.SECONDS); - - driver.get(urlString); - - driver.close(); - System.out.println("Driver closed"); - - proxyServer.abort(); - System.out.println("Proxy stopped"); +@EnableThreadDump +public final class EndToEndStoppingTest { + private static final Logger log = LoggerFactory.getLogger(EndToEndStoppingTest.class); + + private WireMockServer mockServer; + private int mockServerPort; + + @BeforeEach + void setUp() { + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + } + + @AfterEach + void tearDown() { + if (mockServer != null) { + mockServer.stop(); } + } - @Test - public void testWithHttpClient() throws Exception { - mockServer.when(request() - .withMethod("GET") - .withPath("/success"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("Success!") - ); - final String url = "http://127.0.0.1:" + mockServerPort + "/success"; - final String[] sites = { url };// "https://www.google.com.ua"};//"https://exceptional.io"};//"http://www.google.com.ua"}; - for (final String site : sites) { - runSiteTestWithHttpClient(site); - } - } - - private void runSiteTestWithHttpClient(final String site) throws Exception { - // HttpResponse response = client.execute(get); - - // assertEquals(200, response.getStatusLine().getStatusCode()); - // EntityUtils.consume(response.getEntity()); - /* - * final HttpProxyServer ssl = new DefaultHttpProxyServer(PROXY_PORT, - * null, null, new SslHandshakeHandlerFactory(), new HttpRequestFilter() - * { - * - * @Override public void filter(HttpRequest httpRequest) { - * System.out.println("Request went through proxy"); } }); - */ - - final HttpProxyServer proxy = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public io.netty.handler.codec.http.HttpResponse proxyToServerRequest( - HttpObject httpObject) { - System.out.println("Request with through proxy"); - return null; - } - }; - } - }).start(); - - try { - final HttpClient client = TestUtils.createProxiedHttpClient(proxy.getListenAddress().getPort()); - - // final HttpPost get = new HttpPost(site); - final HttpGet get = new HttpGet(site); - - // client.getParams().setParameter(ConnRoutePNames.DEFAULT_PROXY, - // new HttpHost("75.101.134.244", PROXY_PORT)); - // new HttpHost("localhost", PROXY_PORT, "https")); - HttpResponse response = client.execute(get); - assertEquals(200, response.getStatusLine().getStatusCode()); - final HttpEntity entity = response.getEntity(); - final String body = IOUtils.toString(entity.getContent(), StandardCharsets.US_ASCII); - EntityUtils.consume(entity); - - log.info("Consuming entity -- got body: {}", body); - EntityUtils.consume(response.getEntity()); - - log.info("Stopping proxy"); - } finally { - if (proxy != null) { - proxy.abort(); - } - } - } + /** + * This is a quick test from nasis that exhibits different behavior from unit tests because unit + * tests call System.exit(). The stop method should stop all non-daemon threads and should cause + * the JVM to exit without explicitly calling System.exit(), which running as an application + * properly tests. + */ + public static void main(final String[] args) { + HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - // @Test - public void testWithWebDriver() { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); + Proxy proxy = createSeleniumProxy(proxyServer); - Proxy proxy = new Proxy(); - proxy.setProxyType(Proxy.ProxyType.MANUAL); - String proxyStr = String.format("localhost:%d", proxyServer.getListenAddress().getPort()); - proxy.setHttpProxy(proxyStr); - proxy.setSslProxy(proxyStr); + FirefoxOptions capability = new FirefoxOptions(); + capability.setCapability(CapabilityType.PROXY, proxy); - FirefoxOptions capability = new FirefoxOptions(); - capability.setCapability(CapabilityType.PROXY, proxy); + String urlString = "https://www.yahoo.com/"; + WebDriver driver = new FirefoxDriver(capability); + driver.manage().timeouts().pageLoadTimeout(ofSeconds(30)); - final String urlString = "http://www.yahoo.com/"; + driver.get(urlString); - // Note this will actually launch a browser!! - final WebDriver driver = new FirefoxDriver(capability); - driver.manage().timeouts().pageLoadTimeout(30, TimeUnit.SECONDS); + driver.close(); + System.out.println("Driver closed"); - driver.get(urlString); - final String source = driver.getPageSource(); + proxyServer.abort(); + System.out.println("Proxy stopped"); + } - // Just make sure it got something within reason. - assertThat(source.length(), greaterThan(100)); - driver.close(); + @Test + public void testWithHttpClient() throws Exception { + mockServer.stubFor( + get(urlEqualTo("/success")).willReturn(aResponse().withStatus(200).withBody("Success!"))); - proxyServer.abort(); + final String url = "http://127.0.0.1:" + mockServerPort + "/success"; + final String[] sites = {url}; + for (final String site : sites) { + runSiteTestWithHttpClient(site); } - + } + + private void runSiteTestWithHttpClient(final String site) throws Exception { + final HttpProxyServer proxy = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource( + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Override + public io.netty.handler.codec.http.HttpResponse proxyToServerRequest( + @NonNull HttpObject httpObject) { + System.out.println("Request with through proxy"); + return null; + } + }; + } + }) + .start(); + + try (CloseableHttpClient client = createProxiedHttpClient(proxy.getListenAddress().getPort())) { + // final HttpPost get = new HttpPost(site); + final HttpGet get = new HttpGet(site); + + // client.getParams().setParameter(ConnRoutePNames.DEFAULT_PROXY, + // new HttpHost("75.101.134.244", PROXY_PORT)); + // new HttpHost("localhost", PROXY_PORT, "https")); + HttpResponse response = client.execute(get); + assertThat(response.getStatusLine().getStatusCode()).isEqualTo(200); + final HttpEntity entity = response.getEntity(); + final String body = IOUtils.toString(entity.getContent(), StandardCharsets.US_ASCII); + EntityUtils.consume(entity); + + log.info("Consuming entity -- got body: {}", body); + EntityUtils.consume(response.getEntity()); + + log.info("Stopping proxy"); + } finally { + proxy.abort(); + } + } + + /** This test actually launches a browser! */ + @Test + @Timeout(60) + public void testWithWebDriver() { + String os = System.getProperty("os.name", "unknown").toLowerCase(ROOT); + log.info("OS: {} (is windows: {})", os, os.contains("win")); + assumeThat(os.contains("win")).isFalse(); + + assumeThat(isChromeDriverAvailable()) + .as("Chrome/chromedriver must be installed to run WebDriver tests") + .isTrue(); + + HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + + try { + Proxy proxy = createSeleniumProxy(proxyServer); + tryProxyWithBrowser(proxy); + } finally { + proxyServer.abort(); + } + } + + private static Proxy createSeleniumProxy(HttpProxyServer proxyServer) { + Proxy proxy = new Proxy(); + proxy.setProxyType(Proxy.ProxyType.MANUAL); + String proxyStr = String.format("localhost:%d", proxyServer.getListenAddress().getPort()); + proxy.setHttpProxy(proxyStr); + proxy.setSslProxy(proxyStr); + return proxy; + } + + private void tryProxyWithBrowser(Proxy proxy) { + ChromeOptions options = new ChromeOptions(); + options.setCapability(CapabilityType.PROXY, proxy); + options.addArguments("--headless=new"); + + WebDriver driver = new ChromeDriver(options); + try { + driver.manage().timeouts().pageLoadTimeout(ofSeconds(30)); + + driver.get("https://github.com/littleProxy/LittleProxy"); + String source = driver.getPageSource(); + + // Just make sure it got something within reason + assertThat(source).hasSizeGreaterThan(100); + } finally { + driver.quit(); + } + } + + private static boolean isChromeDriverAvailable() { + String[] commands = {"chromedriver", "google-chrome", "google-chrome-stable", "chrome"}; + for (String cmd : commands) { + try { + Process p = new ProcessBuilder("which", cmd).redirectErrorStream(true).start(); + int exitCode = p.waitFor(); + if (exitCode == 0) { + log.info("Found browser/driver: {}", cmd); + return true; + } + } catch (Exception e) { + // ignore + } + } + log.warn("No Chrome/chromedriver found on PATH - WebDriver tests will be skipped"); + return false; + } } diff --git a/src/test/java/org/littleshoot/proxy/HttpFilterTest.java b/src/test/java/org/littleshoot/proxy/HttpFilterTest.java index f4ca7128..a143dfd9 100644 --- a/src/test/java/org/littleshoot/proxy/HttpFilterTest.java +++ b/src/test/java/org/littleshoot/proxy/HttpFilterTest.java @@ -1,969 +1,1280 @@ package org.littleshoot.proxy; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.test.HttpClientUtil.performHttpGet; +import static org.littleshoot.proxy.test.HttpClientUtil.performLocalHttpGet; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import com.github.tomakehurst.wiremock.WireMockServer; import io.netty.channel.ChannelHandlerContext; import io.netty.handler.codec.http.*; -import org.eclipse.jetty.server.Server; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import org.littleshoot.proxy.test.HttpClientUtil; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; - -import javax.net.ssl.SSLEngine; import java.io.IOException; import java.net.InetSocketAddress; import java.net.Socket; import java.net.UnknownHostException; import java.util.LinkedList; import java.util.Queue; -import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; import java.util.concurrent.atomic.AtomicInteger; import java.util.concurrent.atomic.AtomicLongArray; import java.util.concurrent.atomic.AtomicReference; +import javax.net.ssl.SSLEngine; +import org.eclipse.jetty.server.Server; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.test.EnableThreadDump; +import org.littleshoot.proxy.test.HttpClientUtil; -import static org.hamcrest.Matchers.lessThan; -import static org.junit.Assert.*; -import static org.mockito.Mockito.mock; -import static org.mockito.Mockito.when; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; - -public class HttpFilterTest { - private Server webServer; - private HttpProxyServer proxyServer; - private int webServerPort; - - private ClientAndServer mockServer; - private int mockServerPort; - - @Before - public void setUp() throws Exception { - webServer = new Server(0); - webServer.start(); - webServerPort = TestUtils.findLocalHttpPort(webServer); - - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); - - } - - @After - public void tearDown() throws Exception { - try { - if (webServer != null) { - webServer.stop(); - } - } finally { - try { - if (proxyServer != null) { - proxyServer.abort(); - } - } finally { - if (mockServer != null) { - mockServer.stop(); - } - } +@EnableThreadDump +public final class HttpFilterTest { + private Server webServer; + private HttpProxyServer proxyServer; + private int webServerPort; + + private WireMockServer mockServer; + private int mockServerPort; + + @BeforeEach + void setUp() throws Exception { + webServer = new Server(0); + webServer.start(); + webServerPort = TestUtils.findLocalHttpPort(webServer); + + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + } + + @AfterEach + void tearDown() throws Exception { + try { + if (webServer != null) { + webServer.stop(); + } + } finally { + try { + if (proxyServer != null) { + proxyServer.abort(); } - } - - /** - * Sets up the HttpProxyServer instance for a test. This method initializes the proxyServer and proxyPort method variables, and should - * be called before making any requests through the proxy server. - * - * The proxy cannot be created in an @Before method because the filtersSource must be initialized by each test before the proxy is - * created. - * - * @param filtersSource HTTP filters source - */ - private void setUpHttpProxyServer(HttpFiltersSource filtersSource) { - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .start(); - - final InetSocketAddress isa = new InetSocketAddress("127.0.0.1", proxyServer.getListenAddress().getPort()); - while (true) { - try (Socket sock = new Socket()) { - sock.connect(isa); - break; - } catch (final IOException e) { - // Keep trying. - } - - try { - Thread.sleep(50); - } catch (InterruptedException e) { - throw new RuntimeException("Interrupted while verifying proxy is connectable"); - } + } finally { + if (mockServer != null) { + mockServer.stop(); } + } } - - @Test - public void testFiltering() throws Exception { - final AtomicInteger shouldFilterCalls = new AtomicInteger(0); - final AtomicInteger filterResponseCalls = new AtomicInteger(0); - final AtomicInteger fullHttpRequestsReceived = new AtomicInteger(0); - final AtomicInteger fullHttpResponsesReceived = new AtomicInteger(0); - final Queue associatedRequests = - new LinkedList<>(); - - final AtomicInteger requestCount = new AtomicInteger(0); - final AtomicLongArray proxyToServerRequestSendingNanos = new AtomicLongArray(new long[] { -1, -1, -1, -1, -1 }); - final AtomicLongArray proxyToServerRequestSentNanos = new AtomicLongArray(new long[] { -1, -1, -1,-1, -1 }); - final AtomicLongArray serverToProxyResponseReceivingNanos = new AtomicLongArray(new long[] { -1, -1,-1, -1, -1 }); - final AtomicLongArray serverToProxyResponseReceivedNanos = new AtomicLongArray(new long[] { -1, -1,-1, -1, -1 }); - final AtomicLongArray proxyToServerConnectionQueuedNanos = new AtomicLongArray(new long[] { -1, -1,-1, -1, -1 }); - final AtomicLongArray proxyToServerResolutionStartedNanos = new AtomicLongArray(new long[] { -1, -1,-1, -1, -1 }); - final AtomicLongArray proxyToServerResolutionSucceededNanos = new AtomicLongArray(new long[] { -1,-1, -1, -1, -1 }); - final AtomicLongArray proxyToServerResolutionFailedNanos = new AtomicLongArray(new long[] { -1,-1, -1, -1, -1 }); - final AtomicLongArray proxyToServerConnectionStartedNanos = new AtomicLongArray(new long[] { -1, -1,-1, -1, -1 }); - final AtomicLongArray proxyToServerConnectionSSLHandshakeStartedNanos = new AtomicLongArray(new long[] {-1, -1, -1, -1, -1 }); - final AtomicLongArray proxyToServerConnectionFailedNanos = new AtomicLongArray(new long[] { -1, -1,-1, -1, -1 }); - final AtomicLongArray proxyToServerConnectionSucceededNanos = new AtomicLongArray(new long[] { -1,-1, -1, -1, -1 }); - final AtomicLongArray serverToProxyResponseTimedOutNanos = new AtomicLongArray(new long[] { -1,-1, -1, -1, -1 }); - - final AtomicReference serverCtxReference = new AtomicReference<>(); - - final String url1 = "http://localhost:" + webServerPort + "/"; - final String url2 = "http://localhost:" + webServerPort + "/testing"; - final String url3 = "http://localhost:" + webServerPort + "/testing2"; - final String url4 = "http://localhost:" + webServerPort + "/testing3"; - final String url5 = "http://localhost:" + webServerPort + "/testing4"; - - final HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - public HttpFilters filterRequest(HttpRequest originalRequest) { - shouldFilterCalls.incrementAndGet(); - associatedRequests.add(originalRequest); - - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest( - HttpObject httpObject) { - fullHttpRequestsReceived.incrementAndGet(); - if (httpObject instanceof HttpRequest) { - HttpRequest httpRequest = (HttpRequest) httpObject; - if (httpRequest.uri().equals(url2)) { - return new DefaultFullHttpResponse( - HttpVersion.HTTP_1_1, - HttpResponseStatus.FORBIDDEN); - } - } - return null; - } - - @Override - public HttpResponse proxyToServerRequest( - HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - HttpRequest httpRequest = (HttpRequest) httpObject; - if (httpRequest.uri().equals("/testing2")) { - return new DefaultFullHttpResponse( - HttpVersion.HTTP_1_1, - HttpResponseStatus.FORBIDDEN); - } - } - return null; - } - - @Override - public void proxyToServerRequestSending() { - proxyToServerRequestSendingNanos.set(requestCount.get(), now()); - } - - @Override - public void proxyToServerRequestSent() { - proxyToServerRequestSentNanos.set(requestCount.get(), now()); - } - - public HttpObject serverToProxyResponse( - HttpObject httpObject) { - if (originalRequest.uri().contains("testing3")) { - return new DefaultFullHttpResponse( - HttpVersion.HTTP_1_1, - HttpResponseStatus.FORBIDDEN); - } - filterResponseCalls.incrementAndGet(); - if (httpObject instanceof FullHttpResponse) { - fullHttpResponsesReceived.incrementAndGet(); - } - if (httpObject instanceof HttpResponse) { - ((HttpResponse) httpObject).headers().set( - "Header-Pre", "1"); - } - return httpObject; - } - - @Override - public void serverToProxyResponseTimedOut() { - serverToProxyResponseTimedOutNanos.set(requestCount.get(), now()); - } - - @Override - public void serverToProxyResponseReceiving() { - serverToProxyResponseReceivingNanos.set(requestCount.get(), now()); - } - - @Override - public void serverToProxyResponseReceived() { - serverToProxyResponseReceivedNanos.set(requestCount.get(), now()); - } - - public HttpObject proxyToClientResponse( - HttpObject httpObject) { - if (originalRequest.uri().contains("testing4")) { - return new DefaultFullHttpResponse( - HttpVersion.HTTP_1_1, - HttpResponseStatus.FORBIDDEN); - } - if (httpObject instanceof HttpResponse) { - ((HttpResponse) httpObject).headers().set( - "Header-Post", "2"); - } - return httpObject; - } - - @Override - public void proxyToServerConnectionQueued() { - proxyToServerConnectionQueuedNanos.set(requestCount.get(), now()); - } - - @Override - public InetSocketAddress proxyToServerResolutionStarted( - String resolvingServerHostAndPort) { - proxyToServerResolutionStartedNanos.set(requestCount.get(), now()); - return super.proxyToServerResolutionStarted(resolvingServerHostAndPort); - } - - @Override - public void proxyToServerResolutionFailed(String hostAndPort) { - proxyToServerResolutionFailedNanos.set(requestCount.get(), now()); - } - - @Override - public void proxyToServerResolutionSucceeded( - String serverHostAndPort, - InetSocketAddress resolvedRemoteAddress) { - proxyToServerResolutionSucceededNanos.set(requestCount.get(), now()); - } - - @Override - public void proxyToServerConnectionStarted() { - proxyToServerConnectionStartedNanos.set(requestCount.get(), now()); - } - - @Override - public void proxyToServerConnectionSSLHandshakeStarted() { - proxyToServerConnectionSSLHandshakeStartedNanos.set(requestCount.get(), now()); - } - - @Override - public void proxyToServerConnectionFailed() { - proxyToServerConnectionFailedNanos.set(requestCount.get(), now()); - } - - @Override - public void proxyToServerConnectionSucceeded(ChannelHandlerContext ctx) { - proxyToServerConnectionSucceededNanos.set(requestCount.get(), now()); - serverCtxReference.set(ctx); - } - - }; - } - - public int getMaximumRequestBufferSizeInBytes() { - return 1024 * 1024; - } - - public int getMaximumResponseBufferSizeInBytes() { - return 1024 * 1024; - } - }; - - setUpHttpProxyServer(filtersSource); - - org.apache.http.HttpResponse response1 = HttpClientUtil.performHttpGet(url1, proxyServer); - // sleep for a short amount of time, to allow the filter methods to be invoked - Thread.sleep(500); - assertEquals( - "Response should have included the custom header from our pre filter", - "1", response1.getFirstHeader("Header-Pre").getValue()); - assertEquals( - "Response should have included the custom header from our post filter", - "2", response1.getFirstHeader("Header-Post").getValue()); - - assertEquals(1, associatedRequests.size()); - assertEquals(1, shouldFilterCalls.get()); - assertEquals(1, fullHttpRequestsReceived.get()); - assertEquals(1, fullHttpResponsesReceived.get()); - assertEquals(1, filterResponseCalls.get()); - - int i = requestCount.get(); - assertThat(proxyToServerConnectionQueuedNanos.get(i), lessThan(proxyToServerResolutionStartedNanos.get(i))); - assertThat(proxyToServerResolutionStartedNanos.get(i), lessThan(proxyToServerResolutionSucceededNanos.get(i))); - assertThat(proxyToServerResolutionSucceededNanos.get(i), lessThan(proxyToServerConnectionStartedNanos.get(i))); - assertEquals(-1, proxyToServerConnectionSSLHandshakeStartedNanos.get(i)); - assertEquals(-1, proxyToServerConnectionFailedNanos.get(i)); - assertEquals(-1, proxyToServerResolutionFailedNanos.get(i)); - assertEquals(-1, serverToProxyResponseTimedOutNanos.get(i)); - assertThat(proxyToServerConnectionStartedNanos.get(i), lessThan(proxyToServerConnectionSucceededNanos.get(i))); - assertThat(proxyToServerConnectionSucceededNanos.get(i), lessThan(proxyToServerRequestSendingNanos.get(i))); - assertThat(proxyToServerRequestSendingNanos.get(i), lessThan(proxyToServerRequestSentNanos.get(i))); - assertThat(proxyToServerRequestSentNanos.get(i), lessThan(serverToProxyResponseReceivingNanos.get(i))); - assertThat(serverToProxyResponseReceivingNanos.get(i), lessThan(serverToProxyResponseReceivedNanos.get(i))); - - // We just open a second connection here since reusing the original - // connection is inconsistent. - requestCount.incrementAndGet(); - org.apache.http.HttpResponse response2 = HttpClientUtil.performHttpGet(url2, proxyServer); - Thread.sleep(500); - - assertEquals(403, response2.getStatusLine().getStatusCode()); - - assertEquals(2, associatedRequests.size()); - assertEquals(2, shouldFilterCalls.get()); - assertEquals(2, fullHttpRequestsReceived.get()); - assertEquals(1, fullHttpResponsesReceived.get()); - assertEquals(1, filterResponseCalls.get()); - - requestCount.incrementAndGet(); - org.apache.http.HttpResponse response3 = HttpClientUtil.performHttpGet(url3, proxyServer); - Thread.sleep(500); - - assertEquals(403, response3.getStatusLine().getStatusCode()); - - assertEquals(3, associatedRequests.size()); - assertEquals(3, shouldFilterCalls.get()); - assertEquals(3, fullHttpRequestsReceived.get()); - assertEquals(1, fullHttpResponsesReceived.get()); - assertEquals(1, filterResponseCalls.get()); - - i = requestCount.get(); - assertThat(proxyToServerConnectionQueuedNanos.get(i), lessThan(proxyToServerResolutionStartedNanos.get(i))); - assertThat(proxyToServerResolutionStartedNanos.get(i), lessThan(proxyToServerResolutionSucceededNanos.get(i))); - assertEquals(-1, proxyToServerConnectionStartedNanos.get(i)); - assertEquals(-1, proxyToServerConnectionSSLHandshakeStartedNanos.get(i)); - assertEquals(-1, proxyToServerConnectionFailedNanos.get(i)); - assertEquals(-1, proxyToServerConnectionSucceededNanos.get(i)); - assertEquals(-1, proxyToServerRequestSendingNanos.get(i)); - assertEquals(-1, proxyToServerRequestSentNanos.get(i)); - assertEquals(-1, serverToProxyResponseReceivingNanos.get(i)); - assertEquals(-1, serverToProxyResponseReceivedNanos.get(i)); - assertEquals(-1, proxyToServerResolutionFailedNanos.get(i)); - assertEquals(-1, serverToProxyResponseTimedOutNanos.get(i)); - - final HttpRequest first = associatedRequests.remove(); - final HttpRequest second = associatedRequests.remove(); - final HttpRequest third = associatedRequests.remove(); - - // Make sure the requests in the filter calls were the requests they - // actually should have been. - assertEquals(url1, first.uri()); - assertEquals(url2, second.uri()); - assertEquals(url3, third.uri()); - - requestCount.incrementAndGet(); - org.apache.http.HttpResponse response4 = HttpClientUtil.performHttpGet(url4, proxyServer); - Thread.sleep(500); - - i = requestCount.get(); - assertThat(proxyToServerConnectionQueuedNanos.get(i), lessThan(proxyToServerResolutionStartedNanos.get(i))); - assertThat(proxyToServerResolutionStartedNanos.get(i), lessThan(proxyToServerResolutionSucceededNanos.get(i))); - assertThat(proxyToServerResolutionSucceededNanos.get(i), lessThan(proxyToServerConnectionStartedNanos.get(i))); - assertEquals(-1, proxyToServerConnectionSSLHandshakeStartedNanos.get(i)); - assertEquals(-1, proxyToServerConnectionFailedNanos.get(i)); - assertEquals(-1, proxyToServerResolutionFailedNanos.get(i)); - assertEquals(-1, serverToProxyResponseTimedOutNanos.get(i)); - assertThat(proxyToServerConnectionStartedNanos.get(i), lessThan(proxyToServerConnectionSucceededNanos.get(i))); - assertThat(proxyToServerConnectionSucceededNanos.get(i), lessThan(proxyToServerRequestSendingNanos.get(i))); - assertThat(proxyToServerRequestSendingNanos.get(i), lessThan(proxyToServerRequestSentNanos.get(i))); - assertThat(proxyToServerRequestSentNanos.get(i), lessThan(serverToProxyResponseReceivingNanos.get(i))); - assertThat(serverToProxyResponseReceivingNanos.get(i), lessThan(serverToProxyResponseReceivedNanos.get(i))); - - requestCount.incrementAndGet(); - org.apache.http.HttpResponse response5 = HttpClientUtil.performHttpGet(url5, proxyServer); - - assertEquals(403, response4.getStatusLine().getStatusCode()); - assertEquals(403, response5.getStatusLine().getStatusCode()); - - assertNotNull("Server channel context from proxyToServerConnectionSucceeded() should not be null", serverCtxReference.get()); - InetSocketAddress remoteAddress = (InetSocketAddress) serverCtxReference.get().channel().remoteAddress(); - assertNotNull("Server's remoteAddress from proxyToServerConnectionSucceeded() should not be null", remoteAddress); - // make sure we're getting the right remote address (and therefore the right server channel context) in the - // proxyToServerConnectionSucceeded() filter method - assertEquals("Server's remoteAddress should connect to localhost", "localhost", remoteAddress.getHostName()); - assertEquals("Server's port should match the web server port", webServerPort, remoteAddress.getPort()); - - webServer.stop(); + } + + /** + * Sets up the HttpProxyServer instance for a test. This method initializes the proxyServer and + * proxyPort method variables, and should be called before making any requests through the proxy + * server. + * + *

The proxy cannot be created in @BeforeEach method because the filtersSource must be + * initialized by each test before the proxy is created. + * + * @param filtersSource HTTP filters source + */ + private void setUpHttpProxyServer(@NonNull HttpFiltersSource filtersSource) { + proxyServer = + DefaultHttpProxyServer.bootstrap().withPort(0).withFiltersSource(filtersSource).start(); + + final InetSocketAddress isa = + new InetSocketAddress("127.0.0.1", proxyServer.getListenAddress().getPort()); + while (true) { + try (Socket sock = new Socket()) { + sock.connect(isa); + break; + } catch (final IOException e) { + // Keep trying. + } + + try { + Thread.sleep(50); + } catch (InterruptedException e) { + throw new RuntimeException("Interrupted while verifying proxy is connectable"); + } } - - @Test - public void testResolutionStartedFilterReturnsUnresolvedAddress() throws Exception { - final AtomicBoolean resolutionSucceeded = new AtomicBoolean(false); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public InetSocketAddress proxyToServerResolutionStarted(String resolvingServerHostAndPort) { - return InetSocketAddress.createUnresolved("localhost", webServerPort); - } - - @Override - public void proxyToServerResolutionSucceeded(String serverHostAndPort, InetSocketAddress resolvedRemoteAddress) { - assertFalse("expected to receive a resolved InetSocketAddress", resolvedRemoteAddress.isUnresolved()); - resolutionSucceeded.set(true); - } - }; - } + } + + @Test + public void testFiltering() throws Exception { + final AtomicInteger shouldFilterCalls = new AtomicInteger(0); + final AtomicInteger filterResponseCalls = new AtomicInteger(0); + final AtomicInteger fullHttpRequestsReceived = new AtomicInteger(0); + final AtomicInteger fullHttpResponsesReceived = new AtomicInteger(0); + final Queue associatedRequests = new LinkedList<>(); + + final AtomicInteger requestCount = new AtomicInteger(0); + final AtomicLongArray proxyToServerRequestSendingNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerRequestSentNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray serverToProxyResponseReceivingNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray serverToProxyResponseReceivedNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerConnectionQueuedNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerResolutionStartedNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerResolutionSucceededNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerResolutionFailedNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerConnectionStartedNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerConnectionSSLHandshakeStartedNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerConnectionFailedNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray proxyToServerConnectionSucceededNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + final AtomicLongArray serverToProxyResponseTimedOutNanos = + new AtomicLongArray(new long[] {-1, -1, -1, -1, -1}); + + final AtomicReference serverCtxReference = new AtomicReference<>(); + + final String url1 = "http://localhost:" + webServerPort + "/"; + final String url2 = "http://localhost:" + webServerPort + "/testing"; + final String url3 = "http://localhost:" + webServerPort + "/testing2"; + final String url4 = "http://localhost:" + webServerPort + "/testing3"; + final String url5 = "http://localhost:" + webServerPort + "/testing4"; + + final HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + shouldFilterCalls.incrementAndGet(); + associatedRequests.add(originalRequest); + + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + fullHttpRequestsReceived.incrementAndGet(); + if (httpObject instanceof HttpRequest httpRequest) { + if (httpRequest.uri().equals(url2)) { + return new DefaultFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.FORBIDDEN); + } + } + return null; + } + + @Nullable + @Override + public HttpResponse proxyToServerRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest httpRequest) { + if ("/testing2".equals(httpRequest.uri())) { + return new DefaultFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.FORBIDDEN); + } + } + return null; + } + + @Override + public void proxyToServerRequestSending() { + proxyToServerRequestSendingNanos.set(requestCount.get(), now()); + } + + @Override + public void proxyToServerRequestSent() { + proxyToServerRequestSentNanos.set(requestCount.get(), now()); + } + + @NonNull + public HttpObject serverToProxyResponse(@NonNull HttpObject httpObject) { + if (originalRequest.uri().contains("testing3")) { + return new DefaultFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.FORBIDDEN); + } + filterResponseCalls.incrementAndGet(); + if (httpObject instanceof FullHttpResponse) { + fullHttpResponsesReceived.incrementAndGet(); + } + if (httpObject instanceof HttpResponse) { + ((HttpResponse) httpObject).headers().set("Header-Pre", "1"); + } + return httpObject; + } + + @Override + public void serverToProxyResponseTimedOut() { + serverToProxyResponseTimedOutNanos.set(requestCount.get(), now()); + } + + @Override + public void serverToProxyResponseReceiving() { + serverToProxyResponseReceivingNanos.set(requestCount.get(), now()); + } + + @Override + public void serverToProxyResponseReceived() { + serverToProxyResponseReceivedNanos.set(requestCount.get(), now()); + } + + @NonNull + public HttpObject proxyToClientResponse(@NonNull HttpObject httpObject) { + if (originalRequest.uri().contains("testing4")) { + return new DefaultFullHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.FORBIDDEN); + } + if (httpObject instanceof HttpResponse) { + ((HttpResponse) httpObject).headers().set("Header-Post", "2"); + } + return httpObject; + } + + @Override + public void proxyToServerConnectionQueued() { + proxyToServerConnectionQueuedNanos.set(requestCount.get(), now()); + } + + @Nullable + @Override + public InetSocketAddress proxyToServerResolutionStarted( + @NonNull String resolvingServerHostAndPort) { + proxyToServerResolutionStartedNanos.set(requestCount.get(), now()); + return super.proxyToServerResolutionStarted(resolvingServerHostAndPort); + } + + @Override + public void proxyToServerResolutionFailed(@NonNull String hostAndPort) { + proxyToServerResolutionFailedNanos.set(requestCount.get(), now()); + } + + @Override + public void proxyToServerResolutionSucceeded( + @NonNull String serverHostAndPort, + @NonNull InetSocketAddress resolvedRemoteAddress) { + proxyToServerResolutionSucceededNanos.set(requestCount.get(), now()); + } + + @Override + public void proxyToServerConnectionStarted() { + proxyToServerConnectionStartedNanos.set(requestCount.get(), now()); + } + + @Override + public void proxyToServerConnectionSSLHandshakeStarted() { + proxyToServerConnectionSSLHandshakeStartedNanos.set(requestCount.get(), now()); + } + + @Override + public void proxyToServerConnectionFailed() { + proxyToServerConnectionFailedNanos.set(requestCount.get(), now()); + } + + @Override + public void proxyToServerConnectionSucceeded(@NonNull ChannelHandlerContext ctx) { + proxyToServerConnectionSucceededNanos.set(requestCount.get(), now()); + serverCtxReference.set(ctx); + } + }; + } + + public int getMaximumRequestBufferSizeInBytes() { + return 1024 * 1024; + } + + public int getMaximumResponseBufferSizeInBytes() { + return 1024 * 1024; + } }; - setUpHttpProxyServer(filtersSource); + setUpHttpProxyServer(filtersSource); + + org.apache.http.HttpResponse response1 = performHttpGet(url1, proxyServer); + // sleep for a short amount of time, to allow the filter methods to be invoked + Thread.sleep(500); + assertThat(response1.getFirstHeader("Header-Pre").getValue()) + .as("Response should have included the custom header from our pre filter") + .isEqualTo("1"); + assertThat(response1.getFirstHeader("Header-Post").getValue()) + .as("Response should have included the custom header from our post filter") + .isEqualTo("2"); + + assertThat(associatedRequests).hasSize(1); + assertThat(shouldFilterCalls.get()).isEqualTo(1); + assertThat(fullHttpRequestsReceived.get()).isEqualTo(1); + assertThat(fullHttpResponsesReceived.get()).isEqualTo(1); + assertThat(filterResponseCalls.get()).isEqualTo(1); + + int i = requestCount.get(); + assertThat(proxyToServerConnectionQueuedNanos.get(i)) + .isLessThan(proxyToServerResolutionStartedNanos.get(i)); + assertThat(proxyToServerResolutionStartedNanos.get(i)) + .isLessThan(proxyToServerResolutionSucceededNanos.get(i)); + assertThat(proxyToServerResolutionSucceededNanos.get(i)) + .isLessThan(proxyToServerConnectionStartedNanos.get(i)); + assertThat(proxyToServerConnectionSSLHandshakeStartedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerConnectionFailedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerResolutionFailedNanos.get(i)).isEqualTo(-1); + assertThat(serverToProxyResponseTimedOutNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerConnectionStartedNanos.get(i)) + .isLessThan(proxyToServerConnectionSucceededNanos.get(i)); + assertThat(proxyToServerConnectionSucceededNanos.get(i)) + .isLessThan(proxyToServerRequestSendingNanos.get(i)); + assertThat(proxyToServerRequestSendingNanos.get(i)) + .isLessThan(proxyToServerRequestSentNanos.get(i)); + assertThat(proxyToServerRequestSentNanos.get(i)) + .isLessThan(serverToProxyResponseReceivingNanos.get(i)); + assertThat(serverToProxyResponseReceivingNanos.get(i)) + .isLessThan(serverToProxyResponseReceivedNanos.get(i)); + + // We just open a second connection here since reusing the original + // connection is inconsistent. + requestCount.incrementAndGet(); + org.apache.http.HttpResponse response2 = performHttpGet(url2, proxyServer); + Thread.sleep(500); + + assertThat(response2.getStatusLine().getStatusCode()).isEqualTo(403); + + assertThat(associatedRequests).hasSize(2); + assertThat(shouldFilterCalls.get()).isEqualTo(2); + assertThat(fullHttpRequestsReceived.get()).isEqualTo(2); + assertThat(fullHttpResponsesReceived.get()).isEqualTo(1); + assertThat(filterResponseCalls.get()).isEqualTo(1); + + requestCount.incrementAndGet(); + org.apache.http.HttpResponse response3 = performHttpGet(url3, proxyServer); + Thread.sleep(500); + + assertThat(response3.getStatusLine().getStatusCode()).isEqualTo(403); + + assertThat(associatedRequests).hasSize(3); + assertThat(shouldFilterCalls.get()).isEqualTo(3); + assertThat(fullHttpRequestsReceived.get()).isEqualTo(3); + assertThat(fullHttpResponsesReceived.get()).isEqualTo(1); + assertThat(filterResponseCalls.get()).isEqualTo(1); + + i = requestCount.get(); + assertThat(proxyToServerConnectionQueuedNanos.get(i)) + .isLessThan(proxyToServerResolutionStartedNanos.get(i)); + assertThat(proxyToServerResolutionStartedNanos.get(i)) + .isLessThan(proxyToServerResolutionSucceededNanos.get(i)); + assertThat(proxyToServerConnectionStartedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerConnectionSSLHandshakeStartedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerConnectionFailedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerConnectionSucceededNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerRequestSendingNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerRequestSentNanos.get(i)).isEqualTo(-1); + assertThat(serverToProxyResponseReceivingNanos.get(i)).isEqualTo(-1); + assertThat(serverToProxyResponseReceivedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerResolutionFailedNanos.get(i)).isEqualTo(-1); + assertThat(serverToProxyResponseTimedOutNanos.get(i)).isEqualTo(-1); + + final HttpRequest first = associatedRequests.remove(); + final HttpRequest second = associatedRequests.remove(); + final HttpRequest third = associatedRequests.remove(); + + // Make sure the requests in the filter calls were the requests they + // actually should have been. + assertThat(first.uri()).isEqualTo(url1); + assertThat(second.uri()).isEqualTo(url2); + assertThat(third.uri()).isEqualTo(url3); + + requestCount.incrementAndGet(); + org.apache.http.HttpResponse response4 = performHttpGet(url4, proxyServer); + Thread.sleep(500); + + i = requestCount.get(); + assertThat(proxyToServerConnectionQueuedNanos.get(i)) + .isLessThan(proxyToServerResolutionStartedNanos.get(i)); + assertThat(proxyToServerResolutionStartedNanos.get(i)) + .isLessThan(proxyToServerResolutionSucceededNanos.get(i)); + assertThat(proxyToServerResolutionSucceededNanos.get(i)) + .isLessThan(proxyToServerConnectionStartedNanos.get(i)); + assertThat(proxyToServerConnectionSSLHandshakeStartedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerConnectionFailedNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerResolutionFailedNanos.get(i)).isEqualTo(-1); + assertThat(serverToProxyResponseTimedOutNanos.get(i)).isEqualTo(-1); + assertThat(proxyToServerConnectionStartedNanos.get(i)) + .isLessThan(proxyToServerConnectionSucceededNanos.get(i)); + assertThat(proxyToServerConnectionSucceededNanos.get(i)) + .isLessThan(proxyToServerRequestSendingNanos.get(i)); + assertThat(proxyToServerRequestSendingNanos.get(i)) + .isLessThan(proxyToServerRequestSentNanos.get(i)); + assertThat(proxyToServerRequestSentNanos.get(i)) + .isLessThan(serverToProxyResponseReceivingNanos.get(i)); + assertThat(serverToProxyResponseReceivingNanos.get(i)) + .isLessThan(serverToProxyResponseReceivedNanos.get(i)); + + requestCount.incrementAndGet(); + org.apache.http.HttpResponse response5 = performHttpGet(url5, proxyServer); + + assertThat(response4.getStatusLine().getStatusCode()).isEqualTo(403); + assertThat(response5.getStatusLine().getStatusCode()).isEqualTo(403); + + assertThat(serverCtxReference.get()) + .as("Server channel context from proxyToServerConnectionSucceeded() should not be null") + .isNotNull(); + + InetSocketAddress remoteAddress = + (InetSocketAddress) serverCtxReference.get().channel().remoteAddress(); + assertThat(remoteAddress) + .as("Server's remoteAddress from proxyToServerConnectionSucceeded() should not be null") + .isNotNull(); + // make sure we're getting the right remote address (and therefore the right + // server channel context) in the + // proxyToServerConnectionSucceeded() filter method + assertThat(remoteAddress.getHostName()) + .as("Server's remoteAddress should connect to localhost") + .isEqualTo("localhost"); + assertThat(remoteAddress.getPort()) + .as("Server's port should match the web server port") + .isEqualTo(webServerPort); + + webServer.stop(); + } + + @Test + public void testResolutionStartedFilterReturnsUnresolvedAddress() throws Exception { + final AtomicBoolean resolutionSucceeded = new AtomicBoolean(false); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Override + public InetSocketAddress proxyToServerResolutionStarted( + @NonNull String resolvingServerHostAndPort) { + return InetSocketAddress.createUnresolved("localhost", webServerPort); + } + + @Override + public void proxyToServerResolutionSucceeded( + @NonNull String serverHostAndPort, + @NonNull InetSocketAddress resolvedRemoteAddress) { + assertThat(resolvedRemoteAddress.isUnresolved()) + .as("expected to receive a resolved InetSocketAddress") + .isFalse(); + resolutionSucceeded.set(true); + } + }; + } + }; - HttpClientUtil.performHttpGet("http://localhost:" + webServerPort + "/", proxyServer); - Thread.sleep(500); + setUpHttpProxyServer(filtersSource); - assertTrue("proxyToServerResolutionSucceeded method was not called", resolutionSucceeded.get()); - } + performLocalHttpGet(webServerPort, "/", proxyServer); + Thread.sleep(500); - @Test - public void testResolutionFailedCalledAfterDnsFailure() throws Exception { - final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); + assertThat(resolutionSucceeded.get()) + .as("proxyToServerResolutionSucceeded method was not called") + .isTrue(); + } - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return filter; - } - }; + @Test + public void testResolutionFailedCalledAfterDnsFailure() throws Exception { + final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); - HostResolver mockFailingResolver = mock(HostResolver.class); - when(mockFailingResolver.resolve("www.doesnotexist", 80)).thenThrow(new UnknownHostException("www.doesnotexist")); - - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .withServerResolver(mockFailingResolver) - .start(); - - HttpClientUtil.performHttpGet("http://www.doesnotexist/some-resource", proxyServer); - Thread.sleep(500); - - // verify that the filters related to this functionality were correctly invoked/not invoked as appropriate, but also verify that - // other filters were invoked/not invoked as expected - assertFalse("proxyToServerResolutionSucceeded method was called but should not have been", filter.isProxyToServerResolutionSucceededInvoked()); - assertTrue("proxyToServerResolutionFailed method was not called", filter.isProxyToServerResolutionFailedInvoked()); - - assertTrue("Expected filter method to be called", filter.isClientToProxyRequestInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionQueuedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerResolutionStartedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToClientResponseInvoked()); - - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionStartedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionFailedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionSucceededInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSendingInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSentInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivingInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseTimedOutInvoked()); - } - - @Test - public void testConnectionFailedCalledAfterConnectionFailure() throws Exception { - final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return filter; - } + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return filter; + } }; - setUpHttpProxyServer(filtersSource); - - // port 0 is not connectable - HttpClientUtil.performHttpGet("http://localhost:0/some-resource", proxyServer); - Thread.sleep(500); - - // verify that the filters related to this functionality were correctly invoked/not invoked as appropriate, but also verify that - // other filters were invoked/not invoked as expected - assertFalse("proxyToServerConnectionSucceeded should not be called when connection fails", filter.isProxyToServerConnectionSucceededInvoked()); - assertTrue("proxyToServerConnectionFailed should be called when connection fails", filter.isProxyToServerConnectionFailedInvoked()); - - assertTrue("Expected filter method to be called", filter.isClientToProxyRequestInvoked()); - // proxyToServerRequest is invoked before the connection is made, so it should be hit - assertTrue("Expected filter method to be called", filter.isProxyToServerRequestInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionQueuedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionStartedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerResolutionStartedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerResolutionSucceededInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToClientResponseInvoked()); - - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSendingInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSentInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionFailedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivingInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseTimedOutInvoked()); - } - - /** - * Verifies the proper filters are invoked when an attempt to connect to an unencrypted upstream chained proxy fails. - */ - @Test - public void testFiltersAfterUnencryptedConnectionToUpstreamProxyFails() throws Exception { - final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return filter; - } + HostResolver mockFailingResolver = mock(); + when(mockFailingResolver.resolve("www.does-not-exist", 80)) + .thenThrow(new UnknownHostException("www.does-not-exist")); + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource(filtersSource) + .withServerResolver(mockFailingResolver) + .start(); + + performHttpGet("http://www.does-not-exist/some-resource", proxyServer); + Thread.sleep(500); + + // verify that the filters related to this functionality were correctly + // invoked/not invoked as appropriate, but also verify that + // other filters were invoked/not invoked as expected + assertThat(filter.isProxyToServerResolutionSucceededInvoked()) + .as("proxyToServerResolutionSucceeded method was called but should not have been") + .isFalse(); + assertThat(filter.isProxyToServerResolutionFailedInvoked()) + .as("proxyToServerResolutionFailed method was not called") + .isTrue(); + + assertThat(filter.isClientToProxyRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionQueuedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerResolutionStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToClientResponseInvoked()) + .as("Expected filter method to be called") + .isTrue(); + + assertThat(filter.isProxyToServerConnectionStartedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerAllowMitmInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerConnectionFailedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerConnectionSucceededInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSendingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSentInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseTimedOutInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + } + + @Test + public void testConnectionFailedCalledAfterConnectionFailure() throws Exception { + final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return filter; + } }; - // set up the proxy that the HTTP client will connect to - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .withChainProxyManager((httpRequest, chainedProxies, clientDetails) -> chainedProxies.add(new ChainedProxyAdapter() { - @Override - public InetSocketAddress getChainedProxyAddress() { - // port 0 is unconnectable - return new InetSocketAddress("127.0.0.1", 0); - } - })) - .start(); - - // the server doesn't have to exist, since the connection to the chained proxy will fail - HttpClientUtil.performHttpGet("http://localhost:1234/some-resource", proxyServer); - Thread.sleep(500); - - // verify that the filters related to this functionality were correctly invoked/not invoked as appropriate, but also verify that - // other filters were invoked/not invoked as expected - assertFalse("proxyToServerConnectionSucceeded should not be called when connection to chained proxy fails", filter.isProxyToServerConnectionSucceededInvoked()); - assertTrue("proxyToServerConnectionFailed should be called when connection to chained proxy fails", filter.isProxyToServerConnectionFailedInvoked()); - - assertTrue("Expected filter method to be called", filter.isClientToProxyRequestInvoked()); - // proxyToServerRequest is invoked before the connection is made, so it should be hit - assertTrue("Expected filter method to be called", filter.isProxyToServerRequestInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionQueuedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionStartedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToClientResponseInvoked()); - - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionStartedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionSucceededInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSendingInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSentInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionFailedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivingInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseTimedOutInvoked()); - } - - /** - * Verifies the proper filters are invoked when an attempt to connect to an upstream chained proxy over SSL fails. - * (The proxyToServerConnectionFailed() filter method is particularly important.) - */ - @Test - public void testFiltersAfterSSLConnectionToUpstreamProxyFails() throws Exception { - // create an upstream chained proxy using the same SSL engine as the chained proxy tests - final HttpProxyServer chainedProxy = DefaultHttpProxyServer.bootstrap() - .withName("ChainedProxy") - .withPort(0) - .withSslEngineSource(new SelfSignedSslEngineSource("chain_proxy_keystore_1.jks")) - .start(); - - final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return filter; - } + setUpHttpProxyServer(filtersSource); + + // port 0 is not connectable + performHttpGet("http://localhost:0/some-resource", proxyServer); + Thread.sleep(500); + + // verify that the filters related to this functionality were correctly + // invoked/not invoked as appropriate, but also verify that + // other filters were invoked/not invoked as expected + assertThat(filter.isProxyToServerConnectionSucceededInvoked()) + .as("proxyToServerConnectionSucceeded should not be called when connection fails") + .isFalse(); + assertThat(filter.isProxyToServerConnectionFailedInvoked()) + .as("proxyToServerConnectionFailed should be called when connection fails") + .isTrue(); + + assertThat(filter.isClientToProxyRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + // proxyToServerRequest is invoked before the connection is made, so it should + // be hit + assertThat(filter.isProxyToServerRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionQueuedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerResolutionStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerResolutionSucceededInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToClientResponseInvoked()) + .as("Expected filter method to be called") + .isTrue(); + + assertThat(filter.isProxyToServerAllowMitmInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSendingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSentInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionFailedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseTimedOutInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + } + + /** + * Verifies the proper filters are invoked when an attempt to connect to an unencrypted upstream + * chained proxy fails. + */ + @Test + public void testFiltersAfterUnencryptedConnectionToUpstreamProxyFails() throws Exception { + final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return filter; + } }; - // set up the proxy that the HTTP client will connect to - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .withChainProxyManager((httpRequest, chainedProxies, clientDetails) -> chainedProxies.add(new ChainedProxyAdapter() { - @Override - public InetSocketAddress getChainedProxyAddress() { - return chainedProxy.getListenAddress(); - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - // use the same "bad" keystore as BadServerAuthenticationTCPChainedProxyTest - return new SelfSignedSslEngineSource("chain_proxy_keystore_2.jks").newSslEngine(); - } - })) - .start(); - - // the server doesn't have to exist, since the connection to the chained proxy will fail - HttpClientUtil.performHttpGet("http://localhost:1234/some-resource", proxyServer); - Thread.sleep(500); - - // verify that the filters related to this functionality were correctly invoked/not invoked as appropriate, but also verify that - // other filters were invoked/not invoked as expected - assertFalse("proxyToServerConnectionSucceeded should not be called when connection to chained proxy fails", filter.isProxyToServerConnectionSucceededInvoked()); - assertTrue("proxyToServerConnectionFailed should be called when connection to chained proxy fails", filter.isProxyToServerConnectionFailedInvoked()); - - assertTrue("Expected filter method to be called", filter.isClientToProxyRequestInvoked()); - // proxyToServerRequest is invoked before the connection is made, so it should be hit - assertTrue("Expected filter method to be called", filter.isProxyToServerRequestInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionQueuedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionStartedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToClientResponseInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()); - - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionStartedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionSucceededInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSendingInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerRequestSentInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionFailedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivingInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivedInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseTimedOutInvoked()); - } - - @Test - public void testResponseTimedOutInvokedAfterServerTimeout() throws Exception { - mockServer.when(request() - .withMethod("GET") - .withPath("/servertimeout"), - Times.once()) - .respond(response() - .withStatusCode(200) - .withDelay(TimeUnit.SECONDS, 10) - .withBody("success")); - - final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return filter; - } + // set up the proxy that the HTTP client will connect to + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource(filtersSource) + .withChainProxyManager( + (httpRequest, chainedProxies, clientDetails) -> + chainedProxies.add( + new ChainedProxyAdapter() { + @Override + public InetSocketAddress getChainedProxyAddress() { + // port 0 is unconnectable + return new InetSocketAddress("127.0.0.1", 0); + } + })) + .start(); + + // the server doesn't have to exist, since the connection to the chained proxy + // will fail + performHttpGet("http://localhost:1234/some-resource", proxyServer); + Thread.sleep(500); + + // verify that the filters related to this functionality were correctly + // invoked/not invoked as appropriate, but also verify that + // other filters were invoked/not invoked as expected + assertThat(filter.isProxyToServerConnectionSucceededInvoked()) + .as( + "proxyToServerConnectionSucceeded should not be called when connection to chained proxy fails") + .isFalse(); + assertThat(filter.isProxyToServerConnectionFailedInvoked()) + .as("proxyToServerConnectionFailed should be called when connection to chained proxy fails") + .isTrue(); + + assertThat(filter.isClientToProxyRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + // proxyToServerRequest is invoked before the connection is made, so it should + // be hit + assertThat(filter.isProxyToServerRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionQueuedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToClientResponseInvoked()) + .as("Expected filter method to be called") + .isTrue(); + + assertThat(filter.isProxyToServerAllowMitmInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionStartedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionSucceededInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSendingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSentInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionFailedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseTimedOutInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + } + + /** + * Verifies the proper filters are invoked when an attempt to connect to an upstream chained proxy + * over SSL fails. (The proxyToServerConnectionFailed() filter method is particularly important.) + */ + @Test + public void testFiltersAfterSSLConnectionToUpstreamProxyFails() throws Exception { + // create an upstream chained proxy using the same SSL engine as the chained + // proxy tests + final HttpProxyServer chainedProxy = + DefaultHttpProxyServer.bootstrap() + .withName("ChainedProxy") + .withPort(0) + .withSslEngineSource(new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks")) + .start(); + + final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return filter; + } }; - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .withIdleConnectionTimeout(3) - .start(); - - org.apache.http.HttpResponse httpResponse = HttpClientUtil.performHttpGet("http://localhost:" + mockServerPort + "/servertimeout", proxyServer); - assertEquals("Expected to receive an HTTP 504 Gateway Timeout from proxy", 504, httpResponse.getStatusLine().getStatusCode()); - - Thread.sleep(500); - - // verify that the filters related to this functionality were correctly invoked/not invoked as appropriate, but also verify that - // other filters were invoked/not invoked as expected - assertTrue("Expected filter method to be called", filter.isServerToProxyResponseTimedOutInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivingInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseInvoked()); - assertFalse("Expected filter method to not be called", filter.isServerToProxyResponseReceivedInvoked()); - - assertTrue("Expected filter method to be called", filter.isClientToProxyRequestInvoked()); - // proxyToServerRequest is invoked before the connection is made, so it should be hit - assertTrue("Expected filter method to be called", filter.isProxyToServerRequestInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionQueuedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionStartedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerResolutionStartedInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerResolutionSucceededInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerRequestSendingInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerRequestSentInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToServerConnectionSucceededInvoked()); - assertTrue("Expected filter method to be called", filter.isProxyToClientResponseInvoked()); - - assertFalse("Expected filter method to not be called", filter.isProxyToServerResolutionFailedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionFailedInvoked()); - assertFalse("Expected filter method to not be called", filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()); - } - - @Test - public void testRequestSentInvokedAfterLastHttpContentSent() throws Exception { - final AtomicBoolean lastHttpContentProcessed = new AtomicBoolean(false); - final AtomicBoolean requestSentCallbackInvoked = new AtomicBoolean(false); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse proxyToServerRequest(HttpObject httpObject) { - if (httpObject instanceof LastHttpContent) { - assertFalse("requestSentCallback should not be invoked until the LastHttpContent is processed", requestSentCallbackInvoked.get()); - - lastHttpContentProcessed.set(true); - } - - return null; - } - - @Override - public void proxyToServerRequestSent() { - // proxyToServerRequestSent should only be invoked after the entire request, including payload, has been sent to the server - assertTrue("proxyToServerRequestSent callback invoked before LastHttpContent was received from the client and sent to the server", lastHttpContentProcessed.get()); - - requestSentCallbackInvoked.set(true); - } - }; - } + // set up the proxy that the HTTP client will connect to + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource(filtersSource) + .withChainProxyManager( + (httpRequest, chainedProxies, clientDetails) -> + chainedProxies.add( + new ChainedProxyAdapter() { + @Override + public InetSocketAddress getChainedProxyAddress() { + return chainedProxy.getListenAddress(); + } + + @Override + public boolean requiresEncryption() { + return true; + } + + @Override + public SSLEngine newSslEngine() { + // use the same "bad" keystore as + // BadServerAuthenticationTCPChainedProxyTest + return new SelfSignedSslEngineSource( + "target/chain_proxy_keystore_2.jks") + .newSslEngine(); + } + })) + .start(); + + // the server doesn't have to exist, since the connection to the chained proxy + // will fail + performHttpGet("http://localhost:1234/some-resource", proxyServer); + Thread.sleep(500); + + // verify that the filters related to this functionality were correctly + // invoked/not invoked as appropriate, but also verify that + // other filters were invoked/not invoked as expected + assertThat(filter.isProxyToServerConnectionSucceededInvoked()) + .as( + "proxyToServerConnectionSucceeded should not be called when connection to chained proxy fails") + .isFalse(); + assertThat(filter.isProxyToServerConnectionFailedInvoked()) + .as("proxyToServerConnectionFailed should be called when connection to chained proxy fails") + .isTrue(); + + assertThat(filter.isClientToProxyRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + // proxyToServerRequest is invoked before the connection is made, so it should + // be hit + assertThat(filter.isProxyToServerRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionQueuedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToClientResponseInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + + assertThat(filter.isProxyToServerAllowMitmInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionStartedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionSucceededInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSendingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerRequestSentInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionFailedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseTimedOutInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + } + + @Test + @Timeout(20) + public void testResponseTimedOutInvokedAfterServerTimeout() throws Exception { + mockServer.stubFor( + get(urlEqualTo("/servertimeout")) + .willReturn(aResponse().withStatus(200).withFixedDelay(10000).withBody("success"))); + + final HttpFiltersMethodInvokedAdapter filter = new HttpFiltersMethodInvokedAdapter(); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return filter; + } }; - setUpHttpProxyServer(filtersSource); - - // test with a POST request with a payload. post a large amount of data, to force chunked content. - HttpClientUtil.performHttpPost("http://localhost:" + webServerPort + "/", 50000, proxyServer); - Thread.sleep(500); - - assertTrue("proxyToServerRequest callback was not invoked for LastHttpContent for chunked POST", lastHttpContentProcessed.get()); - assertTrue("proxyToServerRequestSent callback was not invoked for chunked POST", requestSentCallbackInvoked.get()); + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource(filtersSource) + .withIdleConnectionTimeout(3) + .start(); + + org.apache.http.HttpResponse httpResponse = + performLocalHttpGet(mockServerPort, "/servertimeout", proxyServer); + assertThat(httpResponse.getStatusLine().getStatusCode()) + .as("Expected to receive an HTTP 504 Gateway Timeout from proxy") + .isEqualTo(504); + + Thread.sleep(500); + + // verify that the filters related to this functionality were correctly + // invoked/not invoked as appropriate, but also verify that + // other filters were invoked/not invoked as expected + assertThat(filter.isServerToProxyResponseTimedOutInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isServerToProxyResponseReceivingInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isServerToProxyResponseReceivedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + + assertThat(filter.isClientToProxyRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + // proxyToServerRequest is invoked before the connection is made, so it should + // be hit + assertThat(filter.isProxyToServerRequestInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionQueuedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerResolutionStartedInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerResolutionSucceededInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerRequestSendingInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerRequestSentInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToServerConnectionSucceededInvoked()) + .as("Expected filter method to be called") + .isTrue(); + assertThat(filter.isProxyToClientResponseInvoked()) + .as("Expected filter method to be called") + .isTrue(); + + assertThat(filter.isProxyToServerAllowMitmInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerResolutionFailedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerConnectionFailedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + assertThat(filter.isProxyToServerConnectionSSLHandshakeStartedInvoked()) + .as("Expected filter method to not be called") + .isFalse(); + } + + @Test + public void testRequestSentInvokedAfterLastHttpContentSent() throws Exception { + final AtomicBoolean lastHttpContentProcessed = new AtomicBoolean(false); + final AtomicBoolean requestSentCallbackInvoked = new AtomicBoolean(false); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse proxyToServerRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof LastHttpContent) { + assertThat(requestSentCallbackInvoked.get()) + .as( + "requestSentCallback should not be invoked until the LastHttpContent is processed") + .isFalse(); + + lastHttpContentProcessed.set(true); + } - // test with a non-payload-bearing GET request. - lastHttpContentProcessed.set(false); - requestSentCallbackInvoked.set(false); + return null; + } + + @Override + public void proxyToServerRequestSent() { + // proxyToServerRequestSent should only be invoked after the entire request, + // including payload, has been sent to the server + assertThat(lastHttpContentProcessed.get()) + .as( + "proxyToServerRequestSent callback invoked before LastHttpContent was received from the client and sent to the server") + .isTrue(); + + requestSentCallbackInvoked.set(true); + } + }; + } + }; - HttpClientUtil.performHttpGet("http://localhost:" + webServerPort + "/", proxyServer); - Thread.sleep(500); + setUpHttpProxyServer(filtersSource); + + // test with a POST request with a payload. post a large amount of data, to + // force chunked content. + HttpClientUtil.performHttpPost("http://localhost:" + webServerPort + "/", 50000, proxyServer); + Thread.sleep(500); + + assertThat(lastHttpContentProcessed.get()) + .as("proxyToServerRequest callback was not invoked for LastHttpContent for chunked POST") + .isTrue(); + assertThat(requestSentCallbackInvoked.get()) + .as("proxyToServerRequestSent callback was not invoked for chunked POST") + .isTrue(); + + // test with a non-payload-bearing GET request. + lastHttpContentProcessed.set(false); + requestSentCallbackInvoked.set(false); + + performLocalHttpGet(webServerPort, "/", proxyServer); + Thread.sleep(500); + + assertThat(lastHttpContentProcessed.get()) + .as("proxyToServerRequest callback was not invoked for LastHttpContent for GET") + .isTrue(); + assertThat(requestSentCallbackInvoked.get()) + .as("proxyToServerRequestSent callback was not invoked for GET") + .isTrue(); + } + + /** + * Verifies that the proxy properly handles a null HttpFilters instance, as allowed in the {@link + * HttpFiltersSource#filterRequest(HttpRequest, ChannelHandlerContext)} documentation. + */ + @Test + public void testNullHttpFilterSource() throws Exception { + mockServer.stubFor( + get(urlEqualTo("/testNullHttpFilterSource")) + .willReturn(aResponse().withStatus(200).withBody("success"))); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @Nullable + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return null; + } + }; - assertTrue("proxyToServerRequest callback was not invoked for LastHttpContent for GET", lastHttpContentProcessed.get()); - assertTrue("proxyToServerRequestSent callback was not invoked for GET", requestSentCallbackInvoked.get()); + setUpHttpProxyServer(filtersSource); + + org.apache.http.HttpResponse httpResponse = + performLocalHttpGet(mockServerPort, "/testNullHttpFilterSource", proxyServer); + Thread.sleep(500); + + assertThat(httpResponse.getStatusLine().getStatusCode()) + .as("Expected to receive an HTTP 200 from proxy") + .isEqualTo(200); + } + + private long now() { + // using nanoseconds instead of milliseconds, since it is extremely unlikely + // that any two callbacks would be invoked in the same nanosecond, + // even on very fast hardware + return System.nanoTime(); + } + + /** + * HttpFilters instance that monitors HttpFilters methods and tracks which methods have been + * invoked. + */ + private static class HttpFiltersMethodInvokedAdapter implements HttpFilters { + private final AtomicBoolean proxyToServerConnectionFailed = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerConnectionSucceeded = new AtomicBoolean(false); + private final AtomicBoolean clientToProxyRequest = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerAllowMitm = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerRequest = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerRequestSending = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerRequestSent = new AtomicBoolean(false); + private final AtomicBoolean serverToProxyResponse = new AtomicBoolean(false); + private final AtomicBoolean serverToProxyResponseReceiving = new AtomicBoolean(false); + private final AtomicBoolean serverToProxyResponseReceived = new AtomicBoolean(false); + private final AtomicBoolean proxyToClientResponse = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerConnectionStarted = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerConnectionQueued = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerResolutionStarted = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerResolutionFailed = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerResolutionSucceeded = new AtomicBoolean(false); + private final AtomicBoolean proxyToServerConnectionSSLHandshakeStarted = + new AtomicBoolean(false); + private final AtomicBoolean serverToProxyResponseTimedOut = new AtomicBoolean(false); + + public boolean isProxyToServerConnectionFailedInvoked() { + return proxyToServerConnectionFailed.get(); } - /** - * Verifies that the proxy properly handles a null HttpFilters instance, as allowed in the - * {@link HttpFiltersSource#filterRequest(HttpRequest, ChannelHandlerContext)} documentation. - */ - @Test - public void testNullHttpFilterSource() throws Exception { - mockServer.when(request() - .withMethod("GET") - .withPath("/testNullHttpFilterSource"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("success")); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return null; - } - }; + public boolean isProxyToServerConnectionSucceededInvoked() { + return proxyToServerConnectionSucceeded.get(); + } - setUpHttpProxyServer(filtersSource); - - org.apache.http.HttpResponse httpResponse = HttpClientUtil.performHttpGet("http://localhost:" + mockServerPort + "/testNullHttpFilterSource", proxyServer); - Thread.sleep(500); - - assertEquals("Expected to receive an HTTP 200 from proxy", 200, httpResponse.getStatusLine().getStatusCode()); - } - - private long now() { - // using nanoseconds instead of milliseconds, since it is extremely unlikely that any two callbacks would be invoked in the same nanosecond, - // even on very fast hardware - return System.nanoTime(); - } - - /** - * HttpFilters instance that monitors HttpFilters methods and tracks which methods have been invoked. - */ - private static class HttpFiltersMethodInvokedAdapter implements HttpFilters { - private final AtomicBoolean proxyToServerConnectionFailed = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerConnectionSucceeded = new AtomicBoolean(false); - private final AtomicBoolean clientToProxyRequest = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerRequest = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerRequestSending = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerRequestSent = new AtomicBoolean(false); - private final AtomicBoolean serverToProxyResponse = new AtomicBoolean(false); - private final AtomicBoolean serverToProxyResponseReceiving = new AtomicBoolean(false); - private final AtomicBoolean serverToProxyResponseReceived = new AtomicBoolean(false); - private final AtomicBoolean proxyToClientResponse = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerConnectionStarted = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerConnectionQueued = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerResolutionStarted = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerResolutionFailed = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerResolutionSucceeded = new AtomicBoolean(false); - private final AtomicBoolean proxyToServerConnectionSSLHandshakeStarted = new AtomicBoolean(false); - private final AtomicBoolean serverToProxyResponseTimedOut = new AtomicBoolean(false); - - public boolean isProxyToServerConnectionFailedInvoked() { - return proxyToServerConnectionFailed.get(); - } + public boolean isClientToProxyRequestInvoked() { + return clientToProxyRequest.get(); + } - public boolean isProxyToServerConnectionSucceededInvoked() { - return proxyToServerConnectionSucceeded.get(); - } + public boolean isProxyToServerAllowMitmInvoked() { + return proxyToServerAllowMitm.get(); + } - public boolean isClientToProxyRequestInvoked() { - return clientToProxyRequest.get(); - } + public boolean isProxyToServerRequestInvoked() { + return proxyToServerRequest.get(); + } - public boolean isProxyToServerRequestInvoked() { - return proxyToServerRequest.get(); - } + public boolean isProxyToServerRequestSendingInvoked() { + return proxyToServerRequestSending.get(); + } - public boolean isProxyToServerRequestSendingInvoked() { - return proxyToServerRequestSending.get(); - } + public boolean isProxyToServerRequestSentInvoked() { + return proxyToServerRequestSent.get(); + } - public boolean isProxyToServerRequestSentInvoked() { - return proxyToServerRequestSent.get(); - } + public boolean isServerToProxyResponseInvoked() { + return serverToProxyResponse.get(); + } - public boolean isServerToProxyResponseInvoked() { - return serverToProxyResponse.get(); - } + public boolean isServerToProxyResponseReceivingInvoked() { + return serverToProxyResponseReceiving.get(); + } - public boolean isServerToProxyResponseReceivingInvoked() { - return serverToProxyResponseReceiving.get(); - } + public boolean isServerToProxyResponseReceivedInvoked() { + return serverToProxyResponseReceived.get(); + } - public boolean isServerToProxyResponseReceivedInvoked() { - return serverToProxyResponseReceived.get(); - } + public boolean isProxyToClientResponseInvoked() { + return proxyToClientResponse.get(); + } - public boolean isProxyToClientResponseInvoked() { - return proxyToClientResponse.get(); - } + public boolean isProxyToServerConnectionStartedInvoked() { + return proxyToServerConnectionStarted.get(); + } - public boolean isProxyToServerConnectionStartedInvoked() { - return proxyToServerConnectionStarted.get(); - } + public boolean isProxyToServerConnectionQueuedInvoked() { + return proxyToServerConnectionQueued.get(); + } - public boolean isProxyToServerConnectionQueuedInvoked() { - return proxyToServerConnectionQueued.get(); - } + public boolean isProxyToServerResolutionStartedInvoked() { + return proxyToServerResolutionStarted.get(); + } - public boolean isProxyToServerResolutionStartedInvoked() { - return proxyToServerResolutionStarted.get(); - } + public boolean isProxyToServerResolutionFailedInvoked() { + return proxyToServerResolutionFailed.get(); + } - public boolean isProxyToServerResolutionFailedInvoked() { - return proxyToServerResolutionFailed.get(); - } + public boolean isProxyToServerResolutionSucceededInvoked() { + return proxyToServerResolutionSucceeded.get(); + } - public boolean isProxyToServerResolutionSucceededInvoked() { - return proxyToServerResolutionSucceeded.get(); - } + public boolean isProxyToServerConnectionSSLHandshakeStartedInvoked() { + return proxyToServerConnectionSSLHandshakeStarted.get(); + } - public boolean isProxyToServerConnectionSSLHandshakeStartedInvoked() { - return proxyToServerConnectionSSLHandshakeStarted.get(); - } + public boolean isServerToProxyResponseTimedOutInvoked() { + return serverToProxyResponseTimedOut.get(); + } - public boolean isServerToProxyResponseTimedOutInvoked() { - return serverToProxyResponseTimedOut.get(); - } + @Override + public void proxyToServerConnectionFailed() { + proxyToServerConnectionFailed.set(true); + } - @Override - public void proxyToServerConnectionFailed() { - proxyToServerConnectionFailed.set(true); - } + @Override + public void proxyToServerConnectionSucceeded(@NonNull ChannelHandlerContext serverCtx) { + proxyToServerConnectionSucceeded.set(true); + } - @Override - public void proxyToServerConnectionSucceeded(ChannelHandlerContext serverCtx) { - proxyToServerConnectionSucceeded.set(true); - } + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + clientToProxyRequest.set(true); + return null; + } - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - clientToProxyRequest.set(true); - return null; - } + @Nullable + @Override + public HttpResponse proxyToServerRequest(@NonNull HttpObject httpObject) { + proxyToServerRequest.set(true); + return null; + } - @Override - public HttpResponse proxyToServerRequest(HttpObject httpObject) { - proxyToServerRequest.set(true); - return null; - } + @Override + public void proxyToServerRequestSending() { + proxyToServerRequestSending.set(true); + } - @Override - public void proxyToServerRequestSending() { - proxyToServerRequestSending.set(true); - } + @Override + public void proxyToServerRequestSent() { + proxyToServerRequestSent.set(true); + } - @Override - public void proxyToServerRequestSent() { - proxyToServerRequestSent.set(true); - } + @Nullable + @Override + public HttpObject serverToProxyResponse(@NonNull HttpObject httpObject) { + serverToProxyResponse.set(true); + return httpObject; + } - @Override - public HttpObject serverToProxyResponse(HttpObject httpObject) { - serverToProxyResponse.set(true); - return httpObject; - } + @Override + public void serverToProxyResponseTimedOut() { + serverToProxyResponseTimedOut.set(true); + } - @Override - public void serverToProxyResponseTimedOut() { - serverToProxyResponseTimedOut.set(true); - } + @Override + public void serverToProxyResponseReceiving() { + serverToProxyResponseReceiving.set(true); + } - @Override - public void serverToProxyResponseReceiving() { - serverToProxyResponseReceiving.set(true); - } + @Override + public void serverToProxyResponseReceived() { + serverToProxyResponseReceived.set(true); + } - @Override - public void serverToProxyResponseReceived() { - serverToProxyResponseReceived.set(true); - } + @Override + public HttpObject proxyToClientResponse(@NonNull HttpObject httpObject) { + proxyToClientResponse.set(true); + return httpObject; + } - @Override - public HttpObject proxyToClientResponse(HttpObject httpObject) { - proxyToClientResponse.set(true); - return httpObject; - } + @Override + public void proxyToServerConnectionQueued() { + proxyToServerConnectionQueued.set(true); + } - @Override - public void proxyToServerConnectionQueued() { - proxyToServerConnectionQueued.set(true); - } + @Override + public InetSocketAddress proxyToServerResolutionStarted( + @NonNull String resolvingServerHostAndPort) { + proxyToServerResolutionStarted.set(true); + return null; + } - @Override - public InetSocketAddress proxyToServerResolutionStarted(String resolvingServerHostAndPort) { - proxyToServerResolutionStarted.set(true); - return null; - } + @Override + public void proxyToServerResolutionFailed(@NonNull String hostAndPort) { + proxyToServerResolutionFailed.set(true); + } - @Override - public void proxyToServerResolutionFailed(String hostAndPort) { - proxyToServerResolutionFailed.set(true); - } + @Override + public void proxyToServerResolutionSucceeded( + @NonNull String serverHostAndPort, @NonNull InetSocketAddress resolvedRemoteAddress) { + proxyToServerResolutionSucceeded.set(true); + } - @Override - public void proxyToServerResolutionSucceeded(String serverHostAndPort, InetSocketAddress resolvedRemoteAddress) { - proxyToServerResolutionSucceeded.set(true); - } + @Override + public void proxyToServerConnectionStarted() { + proxyToServerConnectionStarted.set(true); + } - @Override - public void proxyToServerConnectionStarted() { - proxyToServerConnectionStarted.set(true); - } + @Override + public void proxyToServerConnectionSSLHandshakeStarted() { + proxyToServerConnectionSSLHandshakeStarted.set(true); + } - @Override - public void proxyToServerConnectionSSLHandshakeStarted() { - proxyToServerConnectionSSLHandshakeStarted.set(true); - } + @Override + public boolean proxyToServerAllowMitm() { + clientToProxyRequest.set(true); + return true; } + } } diff --git a/src/test/java/org/littleshoot/proxy/HttpStreamingFilterTest.java b/src/test/java/org/littleshoot/proxy/HttpStreamingFilterTest.java index c38c3fb1..0ffdf0bb 100644 --- a/src/test/java/org/littleshoot/proxy/HttpStreamingFilterTest.java +++ b/src/test/java/org/littleshoot/proxy/HttpStreamingFilterTest.java @@ -1,107 +1,108 @@ package org.littleshoot.proxy; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.createProxiedHttpClient; + import io.netty.handler.codec.http.HttpObject; import io.netty.handler.codec.http.HttpRequest; import io.netty.handler.codec.http.HttpResponse; +import java.util.Arrays; +import java.util.concurrent.atomic.AtomicInteger; import org.apache.http.HttpHost; import org.apache.http.client.methods.HttpPost; import org.apache.http.entity.ByteArrayEntity; import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.util.EntityUtils; import org.eclipse.jetty.server.Server; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import java.util.concurrent.atomic.AtomicInteger; - -import static org.hamcrest.MatcherAssert.assertThat; -import static org.hamcrest.Matchers.greaterThanOrEqualTo; -import static org.junit.Assert.assertEquals; - -public class HttpStreamingFilterTest { - private Server webServer; - private int webServerPort = -1; - private HttpProxyServer proxyServer; - - private final AtomicInteger numberOfInitialRequestsFiltered = new AtomicInteger( - 0); - private final AtomicInteger numberOfSubsequentChunksFiltered = new AtomicInteger( - 0); - - @Before - public void setUp() { - numberOfInitialRequestsFiltered.set(0); - numberOfSubsequentChunksFiltered.set(0); - - webServer = TestUtils.startWebServer(true); - webServerPort = TestUtils.findLocalHttpPort(webServer); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(new HttpFiltersSourceAdapter() { - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest( - HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - numberOfInitialRequestsFiltered - .incrementAndGet(); - } else { - numberOfSubsequentChunksFiltered - .incrementAndGet(); - } - return null; - } - }; - } +public final class HttpStreamingFilterTest { + private static final String DEFAULT_JKS_KEYSTORE_PATH = "target/littleproxy_keystore.jks"; + private Server webServer; + private int webServerPort = -1; + private HttpProxyServer proxyServer; + + private final AtomicInteger numberOfInitialRequestsFiltered = new AtomicInteger(0); + private final AtomicInteger numberOfSubsequentChunksFiltered = new AtomicInteger(0); + + @BeforeEach + void setUp() { + numberOfInitialRequestsFiltered.set(0); + numberOfSubsequentChunksFiltered.set(0); + + webServer = TestUtils.startWebServer(true, DEFAULT_JKS_KEYSTORE_PATH); + webServerPort = TestUtils.findLocalHttpPort(webServer); + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource( + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest) { + numberOfInitialRequestsFiltered.incrementAndGet(); + } else { + numberOfSubsequentChunksFiltered.incrementAndGet(); + } + return null; + } + }; + } }) - .start(); - } - - @After - public void tearDown() throws Exception { - try { - if (proxyServer != null) { - proxyServer.abort(); - } - } finally { - if (webServer != null) { - webServer.stop(); - } - } + .start(); + } + + @AfterEach + void tearDown() throws Exception { + try { + if (proxyServer != null) { + proxyServer.abort(); + } + } finally { + if (webServer != null) { + webServer.stop(); + } } + } - @Test - public void testFiltering() throws Exception { - // Set up some large data to make sure we get chunked encoding on post - byte[] largeData = new byte[20000]; - for (int i = 0; i < largeData.length; i++) { - largeData[i] = 1; - } + @Test + public void testFiltering() throws Exception { + // Set up some large data to make sure we get chunked encoding on post + byte[] largeData = new byte[20000]; + Arrays.fill(largeData, (byte) 1); - final HttpPost request = new HttpPost("/"); - request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); + final HttpPost request = new HttpPost("/"); + request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); - final ByteArrayEntity entity = new ByteArrayEntity(largeData); - entity.setChunked(true); - request.setEntity(entity); + final ByteArrayEntity entity = new ByteArrayEntity(largeData); + entity.setChunked(true); + request.setEntity(entity); - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient( - proxyServer.getListenAddress().getPort()); + try (CloseableHttpClient httpClient = + createProxiedHttpClient(proxyServer.getListenAddress().getPort())) { - final org.apache.http.HttpResponse response = httpClient.execute( - new HttpHost("127.0.0.1", - webServerPort), request); + final org.apache.http.HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", webServerPort), request); - assertEquals("Received 20000 bytes\n", - EntityUtils.toString(response.getEntity())); + assertThat(EntityUtils.toString(response.getEntity())).isEqualTo("Received 20000 bytes\n"); - assertEquals("Filter should have seen only 1 HttpRequest", 1, - numberOfInitialRequestsFiltered.get()); - assertThat("Filter should have seen 1 or more chunks", - numberOfSubsequentChunksFiltered.get(), greaterThanOrEqualTo(1)); + assertThat(numberOfInitialRequestsFiltered.get()) + .as("Filter should have seen only 1 HttpRequest") + .isEqualTo(1); + assertThat(numberOfSubsequentChunksFiltered.get()) + .as("Filter should have seen 1 or more chunks") + .isGreaterThanOrEqualTo(1); } + } } diff --git a/src/test/java/org/littleshoot/proxy/IdleTest.java b/src/test/java/org/littleshoot/proxy/IdleTest.java index bfcc638f..a261bdc0 100644 --- a/src/test/java/org/littleshoot/proxy/IdleTest.java +++ b/src/test/java/org/littleshoot/proxy/IdleTest.java @@ -1,100 +1,104 @@ package org.littleshoot.proxy; -import org.eclipse.jetty.server.Server; -import org.junit.After; -import org.junit.Before; -import org.junit.Ignore; -import org.junit.Test; -import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.disableOnMac; +import static org.littleshoot.proxy.TestUtils.requireUnix; import java.net.InetSocketAddress; import java.net.Proxy; import java.net.URL; - -import static org.hamcrest.Matchers.lessThan; -import static org.junit.Assert.assertThat; -import static org.junit.Assume.assumeFalse; -import static org.junit.Assume.assumeTrue; +import org.eclipse.jetty.server.Server; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Disabled; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.test.EnableThreadDump; /** - * Note - this test only works on UNIX systems because it checks file descriptor - * counts. + * Note - this test only works on UNIX systems because it checks file descriptor counts. * - * It also fails on macOS (tested on 10.14 Mojave). It works on Ubuntu, and presumably most other *nix systems. + *

It also fails on macOS (tested on 10.14 Mojave). It works on Ubuntu, and presumably most other + * *nix systems. */ -public class IdleTest { - private static final int NUMBER_OF_CONNECTIONS_TO_OPEN = 2000; - - private Server webServer; - private int webServerPort = -1; - private HttpProxyServer proxyServer; - - @Before - public void setup() throws Exception { - assumeTrue("Skipping due to non-Unix OS", TestUtils.isUnixManagementCapable()); - assumeFalse("Skipping for travis-ci build", "true".equals(System.getenv("TRAVIS"))); - assumeFalse("Skipping due to Mac OS", System.getProperty("os.name").toLowerCase().contains("mac")); - - webServer = new Server(0); - webServer.start(); - webServerPort = TestUtils.findLocalHttpPort(webServer); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - proxyServer.setIdleConnectionTimeout(10); - +@EnableThreadDump +public final class IdleTest { + private static final int NUMBER_OF_CONNECTIONS_TO_OPEN = 2000; + + private Server webServer; + private int webServerPort = -1; + private HttpProxyServer proxyServer; + + @BeforeEach + void setup() throws Exception { + requireUnix(); + disableOnMac(); + + webServer = new Server(0); + webServer.start(); + webServerPort = TestUtils.findLocalHttpPort(webServer); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + proxyServer.setIdleConnectionTimeout(10); + } + + @AfterEach + void tearDown() throws Exception { + try { + if (webServer != null) { + webServer.stop(); + } + } finally { + if (proxyServer != null) { + proxyServer.abort(); + } } - - @After - public void tearDown() throws Exception { - try { - if (webServer != null) { - webServer.stop(); - } - } finally { - if (proxyServer != null) { - proxyServer.abort(); - } - } + } + + @Test + @Disabled( + "File descriptors vary too much on my laptop, other people saw this problem too: https://github.com/adamfisk/LittleProxy/pull/221") + public void testFileDescriptorCount() throws Exception { + System.out.println("------------------ Memory Usage At Beginning ------------------"); + long initialFileDescriptors = TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); + Proxy proxy = + new Proxy( + Proxy.Type.HTTP, + new InetSocketAddress("127.0.0.1", proxyServer.getListenAddress().getPort())); + for (int i = 0; i < NUMBER_OF_CONNECTIONS_TO_OPEN; i++) { + new URL("http://localhost:" + webServerPort).openConnection(proxy).connect(); } - @Test - @Ignore("File descriptors vary too much on my laptop, other people saw this problem too: https://github.com/adamfisk/LittleProxy/pull/221") - public void testFileDescriptorCount() throws Exception { - System.out - .println("------------------ Memory Usage At Beginning ------------------"); - long initialFileDescriptors = TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); - Proxy proxy = new Proxy(Proxy.Type.HTTP, new InetSocketAddress( - "127.0.0.1", proxyServer.getListenAddress().getPort())); - for (int i = 0; i < NUMBER_OF_CONNECTIONS_TO_OPEN; i++) { - new URL("http://localhost:" + webServerPort) - .openConnection(proxy).connect(); - } - - System.gc(); - System.out - .println("\n\n------------------ Memory Usage Before Idle Timeout ------------------"); - - long fileDescriptorsWhileConnectionsOpen = TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); - Thread.sleep(10000); - - System.gc(); - System.out - .println("\n\n------------------ Memory Usage After Idle Timeout ------------------"); - long fileDescriptorsAfterConnectionsClosed = TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); - - double fdDeltaToOpen = fileDescriptorsWhileConnectionsOpen - - initialFileDescriptors; - double fdDeltaToClosed = fileDescriptorsAfterConnectionsClosed - - initialFileDescriptors; - - double fdDeltaRatio = fdDeltaToClosed / fdDeltaToOpen; - assertThat( - "Number of file descriptors after close should be much closer to initial value than number of file descriptors while open (+ 1%).\n" - + "Initial file descriptors: " + initialFileDescriptors + "; file descriptors while connections open: " + fileDescriptorsWhileConnectionsOpen + "; " - + "file descriptors after connections closed: " + fileDescriptorsAfterConnectionsClosed + "\n" - + "Ratio of file descriptors after connections are closed to descriptors before connections were closed: " + fdDeltaRatio, - fdDeltaRatio, lessThan(0.01)); - } + System.gc(); + System.out.println( + "\n\n------------------ Memory Usage Before Idle Timeout ------------------"); + + long fileDescriptorsWhileConnectionsOpen = + TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); + Thread.sleep(10000); + + System.gc(); + System.out.println("\n\n------------------ Memory Usage After Idle Timeout ------------------"); + long fileDescriptorsAfterConnectionsClosed = + TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); + + double fdDeltaToOpen = fileDescriptorsWhileConnectionsOpen - initialFileDescriptors; + double fdDeltaToClosed = fileDescriptorsAfterConnectionsClosed - initialFileDescriptors; + + double fdDeltaRatio = fdDeltaToClosed / fdDeltaToOpen; + assertThat(fdDeltaRatio) + .as( + "Number of file descriptors after close should be much closer to initial value than number of file descriptors while open (+ 1%).\n" + + "Initial file descriptors: " + + initialFileDescriptors + + "; file descriptors while connections open: " + + fileDescriptorsWhileConnectionsOpen + + "; " + + "file descriptors after connections closed: " + + fileDescriptorsAfterConnectionsClosed + + "\n" + + "Ratio of file descriptors after connections are closed to descriptors before connections were closed: " + + fdDeltaRatio) + .isLessThan(0.01); + } } diff --git a/src/test/java/org/littleshoot/proxy/IdlingProxyTest.java b/src/test/java/org/littleshoot/proxy/IdlingProxyTest.java index 1132c3a2..817950a1 100644 --- a/src/test/java/org/littleshoot/proxy/IdlingProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/IdlingProxyTest.java @@ -1,26 +1,21 @@ package org.littleshoot.proxy; -import org.junit.Test; +import static org.assertj.core.api.Assertions.assertThat; -import static org.junit.Assert.assertEquals; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; -/** - * Tests just a single basic proxy. - */ -public class IdlingProxyTest extends AbstractProxyTest { - @Override - protected void setUp() { - this.proxyServer = bootstrapProxy() - .withPort(0) - .withIdleConnectionTimeout(1) - .start(); - } - - @Test - public void testTimeout() throws Exception { - ResponseInfo response = httpGetWithApacheClient(webHost, "/hang", true, - false); - assertEquals("Received: " + response, 504, response.getStatusCode()); - } +/** Tests just a single basic proxy. */ +@Tag("slow-test") +public final class IdlingProxyTest extends AbstractProxyTest { + @Override + protected void setUp() { + proxyServer = bootstrapProxy().withPort(0).withIdleConnectionTimeout(1).start(); + } + @Test + public void testTimeout() { + ResponseInfo response = httpGetWithApacheClient(webHost, "/hang", true, false); + assertThat(response.getStatusCode()).as("Received: %s", response).isEqualTo(504); + } } diff --git a/src/test/java/org/littleshoot/proxy/InternalRedirectTest.java b/src/test/java/org/littleshoot/proxy/InternalRedirectTest.java new file mode 100644 index 00000000..00eb5a8a --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/InternalRedirectTest.java @@ -0,0 +1,229 @@ +package org.littleshoot.proxy; + +import static com.github.tomakehurst.wiremock.client.WireMock.aResponse; +import static com.github.tomakehurst.wiremock.client.WireMock.get; +import static com.github.tomakehurst.wiremock.client.WireMock.urlEqualTo; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static org.assertj.core.api.Assertions.assertThat; + +import com.github.tomakehurst.wiremock.WireMockServer; +import io.netty.handler.codec.http.DefaultFullHttpResponse; +import io.netty.handler.codec.http.HttpHeaderNames; +import io.netty.handler.codec.http.HttpObject; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpVersion; +import java.io.IOException; +import org.apache.http.HttpEntity; +import org.apache.http.client.methods.CloseableHttpResponse; +import org.apache.http.client.methods.HttpGet; +import org.apache.http.impl.client.CloseableHttpClient; +import org.apache.http.impl.client.HttpClients; +import org.apache.http.util.EntityUtils; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.Timeout; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; + +/** + * Integration test for GitHub issue #68 - demonstrating how to implement internal redirects. + * + *

The question asks how to handle a 302 redirect internally, where the client never sees the 302 + * response - instead, the proxy follows the redirect and returns the final response. + * + *

Solution: Use an external HTTP client (like Apache HttpClient) to make a new request to the + * redirect URL, then return that response as a short-circuit response to the client. + */ +class InternalRedirectTest { + + public static final int FREE_PORT = 0; + @Nullable private HttpProxyServer proxyServer = null; + @Nullable private WireMockServer wireMockServer = null; + + /** Sets up a WireMock server that returns a 302 redirect to /final. */ + @BeforeEach + void setUpRedirectServer() { + wireMockServer = new WireMockServer(options().dynamicPort()); + wireMockServer.start(); + + // Stub: return 302 redirect to /final + wireMockServer.stubFor( + get(urlEqualTo("/redirect")) + .willReturn(aResponse().withStatus(302).withHeader("Location", "/final"))); + + // Stub: return final response + wireMockServer.stubFor( + get(urlEqualTo("/final")) + .willReturn( + aResponse().withStatus(200).withBody("Final response from internal redirect!"))); + } + + @AfterEach + void tearDown() { + if (proxyServer != null) { + proxyServer.abort(); + } + if (wireMockServer != null) { + wireMockServer.stop(); + } + } + + /** + * This test demonstrates how to implement an internal redirect: 1. Detect a 302 redirect in + * serverToProxyResponse 2. Use an external HTTP client to make a new request to the redirect URL + * 3. Return that response as a short-circuit response to the client + * + *

The client will never see the 302 - they'll only see the final response. + */ + @Test + @Timeout(10) + void testInternalRedirectFollowsRedirectTransparently() throws IOException { + assert wireMockServer != null; + int wireMockPort = wireMockServer.port(); + + // Create a filter that follows redirects internally + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Override + @NullMarked + public HttpObject serverToProxyResponse(HttpObject httpObject) { + if (httpObject instanceof HttpResponse response) { + if (HttpResponseStatus.FOUND.equals(response.status()) + && response.headers().contains(HttpHeaderNames.LOCATION)) { + String location = response.headers().get(HttpHeaderNames.LOCATION); + return followRedirect(originalRequest, location); + } + } + return httpObject; + } + }; + } + }; + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(FREE_PORT) + .withFiltersSource(filtersSource) + .start(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + + // Make a request that will trigger a redirect - use custom client to read body + String body = ""; + int statusCode; + try (CloseableHttpClient httpClient = + HttpClients.custom() + .setProxy(new org.apache.http.HttpHost("127.0.0.1", proxyPort)) + .build()) { + HttpGet request = new HttpGet("http://127.0.0.1:" + wireMockPort + "/redirect"); + + try (CloseableHttpResponse resp = httpClient.execute(request)) { + statusCode = resp.getStatusLine().getStatusCode(); + HttpEntity entity = resp.getEntity(); + if (entity != null) { + body = EntityUtils.toString(entity); + } + } + } + + // The client should see the FINAL response (200 OK), NOT the redirect (302) + assertThat(statusCode) + .as("Client should receive the final response, not the redirect") + .isEqualTo(200); + assertThat(body).contains("Final response from internal redirect!"); + } + + /** + * Helper method that follows a redirect by making a new HTTP request and returning the response + * as a short-circuit response. + * + *

This is the key to solving issue #68 - you must use an external HTTP client since + * LittleProxy's filter mechanism doesn't provide a way to "make a new request" directly. + * + * @param originalRequest The original HTTP request from the client (to extract the host) + * @param location The Location header value from the 302 response (can be absolute or relative) + * @return A short-circuit response to send to the client + */ + private HttpResponse followRedirect(HttpRequest originalRequest, String location) { + // Extract the original host from the request + String originalHost = originalRequest.headers().get(HttpHeaderNames.HOST); + if (originalHost == null) { + originalHost = "127.0.0.1"; // Fallback + } + + // Resolve the redirect URL - handle both absolute and relative Location headers + String targetUrl = resolveRedirectUrl(location, originalHost); + + try (CloseableHttpClient httpClient = HttpClients.createDefault()) { + HttpGet httpGet = new HttpGet(targetUrl); + try (CloseableHttpResponse backendResponse = httpClient.execute(httpGet)) { + int statusCode = backendResponse.getStatusLine().getStatusCode(); + HttpResponseStatus nettyStatus = HttpResponseStatus.valueOf(statusCode); + + String body = ""; + HttpEntity entity = backendResponse.getEntity(); + if (entity != null) { + body = EntityUtils.toString(entity); + } + + // Return a short-circuit response with the final response body + DefaultFullHttpResponse response = + new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, nettyStatus); + response.content().writeBytes(body.getBytes()); + response.headers().set(HttpHeaderNames.CONTENT_LENGTH, body.length()); + response.headers().set(HttpHeaderNames.CONTENT_TYPE, "text/plain; charset=UTF-8"); + + return response; + } + } catch (IOException e) { + // Return an error response if the redirect fails + return new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.BAD_GATEWAY); + } + } + + /** + * Resolves a redirect Location header to a full URL. + * + *

Handles both:- Absolute URLs (...) - returned as-is - + * Relative URLs (/path) - combined with the original host + * + * @param location The Location header value + * @param originalHost The original request's Host header + * @return The full URL to request + */ + private String resolveRedirectUrl(String location, String originalHost) { + if (location.startsWith("http://") || location.startsWith("https://")) { + // Already an absolute URL + return location; + } + + // Parse the original host to separate hostname and port + String host = originalHost; + int port = 80; // Default HTTP port + + if (originalHost.contains(":")) { + String[] parts = originalHost.split(":"); + host = parts[0]; + try { + port = Integer.parseInt(parts[1]); + } catch (NumberFormatException e) { + port = 80; + } + } + + // Combine host and relative path + if (location.startsWith("/")) { + return "http://" + host + ":" + port + location; + } else { + return "http://" + host + ":" + port + "/" + location; + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/Issue71NonSslConnectTest.java b/src/test/java/org/littleshoot/proxy/Issue71NonSslConnectTest.java new file mode 100644 index 00000000..8c410dcc --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/Issue71NonSslConnectTest.java @@ -0,0 +1,198 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.fail; + +import java.io.IOException; +import java.net.Socket; +import java.nio.charset.StandardCharsets; +import org.eclipse.jetty.server.Server; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.extras.TestMitmManager; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * Test for issue #71: Incorrect assumption that all connections made via CONNECT requests in MITM + * mode should use SSL. + * + *

When LittleProxy handles a CONNECT request in MITM mode, it assumes that the connection to the + * next hop should always be made over an SSL/TLS connection. Although this is true in 99.9% of + * cases, it is not required by the HTTP spec and breaks real-world scenarios like websocket + * connections (ws:// scheme) that use CONNECT tunnels but don't use SSL. + */ +class Issue71NonSslConnectTest { + private static final Logger LOG = LoggerFactory.getLogger(Issue71NonSslConnectTest.class); + + private static final String DEFAULT_JKS_KEYSTORE_PATH = "target/littleproxy_keystore.jks"; + + private Server webServer; + private int webServerPort; + private HttpProxyServer proxyServer; + + @BeforeEach + void setUp() throws Exception { + webServer = TestUtils.startWebServer(true, DEFAULT_JKS_KEYSTORE_PATH); + + webServerPort = TestUtils.findLocalHttpPort(webServer); + if (webServerPort < 0) { + throw new RuntimeException( + "HTTP connector should already be open and listening, but port was " + webServerPort); + } + + LOG.info("Started webserver on http:{}", webServerPort); + } + + @AfterEach + void tearDown() throws Exception { + try { + if (proxyServer != null) { + LOG.info("Stop proxy server {}", proxyServer.getListenAddress()); + proxyServer.abort(); + } + } finally { + if (webServer != null) { + LOG.info("Stop webserver on http:{}, https:{}", webServerPort, webServerPort); + webServer.stop(); + } + } + } + + /** + * This test reproduces issue #71: When MITM is enabled and a CONNECT request is made to a non-SSL + * server (like what happens with ws:// websockets), the proxy incorrectly tries to establish an + * SSL connection to the server, which fails. + * + *

Scenario: 1. Client sends a CONNECT request to the proxy (e.g., for a websocket connection) + * 2. The proxy, in MITM mode, should establish a plain TCP tunnel to the target server 3. But the + * bug causes the proxy to add SSL to the upstream connection, which fails because the server + * doesn't speak SSL + * + *

This is exactly what happens with browsers and websockets - the browser sends a CONNECT + * request to the proxy, but the target websocket server doesn't use SSL. + */ + @Test + @Tag("slow-test") + void testConnectRequestToNonSslServerInMitmMode() throws Exception { + // Set up MITM proxy + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withManInTheMiddle(new TestMitmManager()) + .start(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + LOG.info("Started MITM proxy on port {}", proxyPort); + + // Send a CONNECT request to the non-SSL HTTP server + // This simulates what browsers do when connecting to ws:// (websocket) servers + try (Socket socket = new Socket("127.0.0.1", proxyPort)) { + socket.setSoTimeout(10000); + + // Send CONNECT request to plain HTTP server (not HTTPS) + String connectRequest = + "CONNECT 127.0.0.1:" + + webServerPort + + " HTTP/1.1\r\n" + + "Host: 127.0.0.1:" + + webServerPort + + "\r\n" + + "\r\n"; + + socket.getOutputStream().write(connectRequest.getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().flush(); + + // Read response + byte[] buffer = new byte[4096]; + int read; + StringBuilder response = new StringBuilder(); + try { + while ((read = socket.getInputStream().read(buffer)) != -1) { + response.append(new String(buffer, 0, read, StandardCharsets.US_ASCII)); + if (response.toString().contains("\r\n\r\n")) { + break; + } + } + } catch (IOException e) { + LOG.error("IOException reading response", e); + String exceptionMessage = e.getMessage(); + if (exceptionMessage != null + && (exceptionMessage.toLowerCase().contains("ssl") + || exceptionMessage.toLowerCase().contains("handshake") + || exceptionMessage.toLowerCase().contains("not an ssl"))) { + fail( + "Issue #71 reproduced: Proxy incorrectly attempted SSL for non-SSL destination. " + + "Exception: " + + e.getMessage()); + } + throw e; + } + + LOG.info("CONNECT response: {}", response); + + String responseStr = response.toString(); + + // Check if we got a Bad Gateway response - this happens when the SSL handshake fails + if (responseStr.contains("502 Bad Gateway") || responseStr.contains("Bad Gateway")) { + fail( + "Issue #71 reproduced: Proxy incorrectly attempted SSL for non-SSL destination. " + + "Got 502 Bad Gateway because the SSL handshake failed with a non-SSL server."); + } + + // If we get here, the CONNECT was successful + assertThat(responseStr) + .as("Should receive successful CONNECT response (200)") + .contains("200"); + + // Now send a plain HTTP request through the tunnel to verify end-to-end data flow. + // This exposes Bug 1 (NPE in MitmEncryptClientChannel when disableSslForNonTls=true) + // and Bug 2 (client channel incorrectly encrypted for MITM while server is plain text). + String httpRequest = + "GET / HTTP/1.1\r\n" + "Host: 127.0.0.1:" + webServerPort + "\r\n" + "\r\n"; + socket.getOutputStream().write(httpRequest.getBytes(StandardCharsets.US_ASCII)); + socket.getOutputStream().flush(); + + StringBuilder tunnelResponse = new StringBuilder(); + try { + while ((read = socket.getInputStream().read(buffer)) != -1) { + tunnelResponse.append(new String(buffer, 0, read, StandardCharsets.US_ASCII)); + if (tunnelResponse.toString().contains("\r\n\r\n")) { + break; + } + } + } catch (IOException e) { + LOG.error("IOException reading tunneled response", e); + String exceptionMessage = e.getMessage(); + if (exceptionMessage != null + && (exceptionMessage.toLowerCase().contains("ssl") + || exceptionMessage.toLowerCase().contains("handshake") + || exceptionMessage.toLowerCase().contains("connection reset"))) { + fail( + "Bug #71 still present: Tunnel data flow failed because proxy encrypted client channel " + + "for MITM while server is plain text. Exception: " + + e.getMessage()); + } + throw e; + } + + LOG.info("Tunneled HTTP response: {}", tunnelResponse); + + String tunnelResponseStr = tunnelResponse.toString(); + assertThat(tunnelResponseStr) + .as( + "Tunneled HTTP request should return a valid response from the non-SSL server, " + + "not a 502 Bad Gateway or SSL error") + .doesNotContain("502") + .doesNotContain("Bad Gateway"); + + // Verify we got actual HTTP content from the upstream server + assertThat(tunnelResponseStr) + .as("Tunneled response should contain HTTP status line") + .contains("HTTP/"); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/KeepAliveTest.java b/src/test/java/org/littleshoot/proxy/KeepAliveTest.java index a50eb3f8..8999ad7a 100644 --- a/src/test/java/org/littleshoot/proxy/KeepAliveTest.java +++ b/src/test/java/org/littleshoot/proxy/KeepAliveTest.java @@ -1,342 +1,408 @@ package org.littleshoot.proxy; -import io.netty.handler.codec.http.*; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; -import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import org.littleshoot.proxy.test.SocketClientUtil; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; -import org.mockserver.model.ConnectionOptions; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static org.assertj.core.api.Assertions.assertThat; +import com.github.tomakehurst.wiremock.WireMockServer; +import io.netty.handler.codec.http.*; import java.io.IOException; import java.net.Socket; +import java.time.Duration; import java.util.Locale; -import java.util.concurrent.TimeUnit; - -import static org.hamcrest.Matchers.*; -import static org.junit.Assert.*; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; - -/** - * This class tests the proxy's keep alive/connection closure behavior. - */ -public class KeepAliveTest { - private HttpProxyServer proxyServer; - - private ClientAndServer mockServer; - private int mockServerPort; - - private Socket socket; - - @Before - public void setUp() { - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); - socket = null; - proxyServer = null; - } - - @After - public void tearDown() throws Exception { - try { - if (proxyServer != null) { - proxyServer.abort(); - } - } finally { - try { - if (mockServer != null) { - mockServer.stop(); - } - } finally { - if (socket != null) { - socket.close(); - } - } +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.junit.jupiter.api.*; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.test.EnableThreadDump; +import org.littleshoot.proxy.test.SocketClientUtil; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** This class tests the proxy's keep alive/connection closure behavior. */ +@Tag("slow-test") +@NullMarked +@Timeout(20) +@EnableThreadDump +public final class KeepAliveTest { + private static final Logger log = LoggerFactory.getLogger(KeepAliveTest.class); + @Nullable private HttpProxyServer proxyServer; + private WireMockServer mockServer; + private int mockServerPort; + @Nullable private Socket socket; + + @BeforeEach + void setUp() { + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + log.info("Mock server port: {} (started: {})", mockServerPort, mockServer.isRunning()); + } + + @AfterEach + @SuppressWarnings("ConstantValue") + void tearDown() throws Exception { + try { + if (proxyServer != null) { + proxyServer.abort(); + } + } finally { + try { + if (mockServer != null) { + mockServer.stop(); } - } - - /** - * Tests that the proxy does not close the connection after a successful HTTP 1.1 GET request and response. - */ - @Test - public void testHttp11DoesNotCloseConnectionByDefault() throws IOException, InterruptedException { - mockServer.when(request() - .withMethod("GET") - .withPath("/success"), - Times.exactly(2)) - .respond(response() - .withStatusCode(200) - .withBody("success")); - - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - this.socket = SocketClientUtil.getSocketToProxyServer(proxyServer); - - // construct the basic request: METHOD + URI + HTTP version + CRLF (to indicate the end of the request) - String successfulGet = "GET http://localhost:" + mockServerPort + "/success HTTP/1.1\r\n" - + "\r\n"; - - // send the same request twice over the same connection - for (int i = 1; i <= 2; i++) { - SocketClientUtil.writeStringToSocket(successfulGet, socket); - - // wait a bit to allow the proxy server to respond - Thread.sleep(750); - - String response = SocketClientUtil.readStringFromSocket(socket); - - assertThat("Expected to receive an HTTP 200 from the server (iteration: " + i + ")", response, startsWith("HTTP/1.1 200 OK")); - assertThat("Unexpected message body (iteration: " + i + ")", response, endsWith("success")); + } finally { + if (socket != null) { + socket.close(); } - - assertTrue("Expected connection to proxy server to be open and readable", SocketClientUtil.isSocketReadyToRead(socket)); - assertTrue("Expected connection to proxy server to be open and writable", SocketClientUtil.isSocketReadyToWrite(socket)); + } } - - /** - * Tests that the proxy keeps the connection to the client open after a server disconnect, even when the server is using - * connection closure to indicate the end of a message. - */ - @Test - public void testProxyKeepsClientConnectionOpenAfterServerDisconnect() throws IOException, InterruptedException { - mockServer.when(request() - .withMethod("GET") - .withPath("/success"), - Times.exactly(2)) - .respond(response() - .withStatusCode(200) - .withBody("success") - .withConnectionOptions(new ConnectionOptions() - .withKeepAliveOverride(false) - .withSuppressContentLengthHeader(true) - .withCloseSocket(true))); - - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - this.socket = SocketClientUtil.getSocketToProxyServer(proxyServer); - - // construct the basic request: METHOD + URI + HTTP version + CRLF (to indicate the end of the request) - String successfulGet = "GET http://localhost:" + mockServerPort + "/success HTTP/1.1\r\n" - + "\r\n"; - - // send the same request twice over the same connection - for (int i = 1; i <= 2; i++) { - SocketClientUtil.writeStringToSocket(successfulGet, socket); - - // wait a bit to allow the proxy server to respond - Thread.sleep(750); - - String response = SocketClientUtil.readStringFromSocket(socket); - - assertThat("Expected to receive an HTTP 200 from the server (iteration: " + i + ")", response, startsWith("HTTP/1.1 200 OK")); - // the proxy will set the Transfer-Encoding to chunked since the server is using connection closure to indicate the end of the message. - // (matching capitalized or lowercase Transfer-Encoding, since Netty 4.1+ uses lowercase header names) - assertThat("Expected proxy to set Transfer-Encoding to chunked", response.toLowerCase(Locale.US), containsString("transfer-encoding: chunked")); - // the Transfer-Encoding is chunked, so the body text will be followed by a 0 and 2 CRLFs - assertThat("Unexpected message body (iteration: " + i + ")", response, containsString("success")); - } - - assertTrue("Expected connection to proxy server to be open and readable", SocketClientUtil.isSocketReadyToRead(socket)); - assertTrue("Expected connection to proxy server to be open and writable", SocketClientUtil.isSocketReadyToWrite(socket)); + } + + /** + * Tests that the proxy does not close the connection after a successful HTTP 1.1 GET request and + * response. + */ + @Test + public void testHttp11DoesNotCloseConnectionByDefault() throws IOException, InterruptedException { + mockServer.stubFor( + get(urlEqualTo("/success")) + .willReturn( + aResponse().withStatus(200).withBody("success").withHeader("Content-Length", "7"))); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + log.info("Started proxy server {}", proxyServer.getListenAddress()); + socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + + // construct the basic request: METHOD + URI + HTTP version + CRLF (to indicate + // the end of the request) + String successfulGet = + "GET http://localhost:" + + mockServerPort + + "/success HTTP/1.1\r\n" + + "Host: localhost:" + + mockServerPort + + "\r\n" + + "\r\n"; + + // send the same request twice over the same connection + for (int i = 1; i <= 2; i++) { + log.debug("#{} Sending to socket {} packet '{}'...", i, socket.getLocalPort(), successfulGet); + SocketClientUtil.writeStringToSocket(successfulGet, socket); + + // wait a bit to allow the proxy server to respond + Thread.sleep(750); + + log.debug("#{} Reading from socket {}...", i, socket.getLocalPort()); + String response = SocketClientUtil.readStringFromSocket(socket); + + assertThat(response) + .as("Expected to receive an HTTP 200 from the server (iteration: %s)", i) + .startsWith("HTTP/1.1 200 OK"); + assertThat(response).as("Unexpected message body (iteration: %s)", i).endsWith("success"); } - /** - * Tests that the proxy does not close the connection after a 502 Bad Gateway response. - */ - @Test - public void testBadGatewayDoesNotCloseConnection() throws IOException, InterruptedException { - mockServer.when(request() - .withMethod("GET") - .withPath("/success"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("success")); - - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - - socket = SocketClientUtil.getSocketToProxyServer(proxyServer); - - String badGatewayGet = "GET http://localhost:0/success HTTP/1.1\r\n" - + "\r\n"; - - // send the same request twice over the same connection - for (int i = 1; i <= 2; i++) { - SocketClientUtil.writeStringToSocket(badGatewayGet, socket); - - // wait a bit to allow the proxy server to respond - Thread.sleep(1500); - - String response = SocketClientUtil.readStringFromSocket(socket); - - assertThat("Expected to receive an HTTP 200 from the server (iteration: " + i + ")", response, startsWith("HTTP/1.1 502 Bad Gateway")); - } - - assertTrue("Expected connection to proxy server to be open and readable", SocketClientUtil.isSocketReadyToRead(socket)); - assertTrue("Expected connection to proxy server to be open and writable", SocketClientUtil.isSocketReadyToWrite(socket)); + assertThat(SocketClientUtil.isSocketReadyToRead(socket)) + .as("Expected connection to proxy server to be open and readable") + .isTrue(); + assertThat(SocketClientUtil.isSocketReadyToWrite(socket)) + .as("Expected connection to proxy server to be open and writable") + .isTrue(); + } + + /** + * Tests that the proxy keeps the connection to the client open after a server disconnect, even + * when the server is using connection closure to indicate the end of a message. + */ + @Test + public void testProxyKeepsClientConnectionOpenAfterServerDisconnect() + throws IOException, InterruptedException { + mockServer.stubFor( + get(urlEqualTo("/success")) + .willReturn( + aResponse().withStatus(200).withBody("success").withHeader("Connection", "close"))); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + log.info("Started proxy server {}", proxyServer.getListenAddress()); + socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + + // construct the basic request: METHOD + URI + HTTP version + CRLF (to indicate + // the end of the request) + String successfulGet = + "GET http://localhost:" + + mockServerPort + + "/success HTTP/1.1\r\n" + + "Host: localhost:" + + mockServerPort + + "\r\n" + + "\r\n"; + + // send the same request twice over the same connection + for (int i = 1; i <= 2; i++) { + log.debug("#{} Sending to socket {} packet '{}'...", i, socket.getLocalPort(), successfulGet); + SocketClientUtil.writeStringToSocket(successfulGet, socket); + + // wait a bit to allow the proxy server to respond + Thread.sleep(750); + + log.debug("#{} Reading from socket {}...", i, socket.getLocalPort()); + String response = SocketClientUtil.readStringFromSocket(socket); + + assertThat(response) + .as("Expected to receive an HTTP 200 from the server (iteration: %s)", i) + .startsWith("HTTP/1.1 200 OK"); + // the proxy will set the Transfer-Encoding to chunked since the server is using + // connection closure to indicate the end of the message. + // (matching capitalized or lowercase Transfer-Encoding, since Netty 4.1+ uses + // lowercase header names) + assertThat(response.toLowerCase(Locale.US)) + .as("Expected proxy to set Transfer-Encoding to chunked") + .contains("transfer-encoding: chunked"); + // the Transfer-Encoding is chunked, so the body text will be followed by a 0 + // and 2 CRLFs + assertThat(response).as("Unexpected message body (iteration: %s)", i).contains("success"); } - /** - * Tests that the proxy does not close the connection after a 504 Gateway Timeout response. - */ - @Test - public void testGatewayTimeoutDoesNotCloseConnection() throws IOException { - mockServer.when(request() - .withMethod("GET") - .withPath("/success"), - Times.exactly(2)) - .respond(response() - .withStatusCode(200) - .withDelay(TimeUnit.SECONDS, 10) - .withBody("success")); - - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withIdleConnectionTimeout(2) - .withPort(0) - .start(); - - socket = SocketClientUtil.getSocketToProxyServer(proxyServer); - - String successfulGet = "GET http://localhost:" + mockServerPort + "/success HTTP/1.1\r\n" - + "\r\n"; - - // send the same request twice over the same connection - for (int i = 1; i <= 2; i++) { - SocketClientUtil.writeStringToSocket(successfulGet, socket); - - String response = SocketClientUtil.readStringFromSocket(socket); - - // match the whole response to make sure that the it is not repeated - assertThat("The response is repeated:", response, is("HTTP/1.1 504 Gateway Timeout\r\n" + - "content-length: 15\r\n" + - "content-type: text/html; charset=utf-8\r\n" + - "\r\n" + - "Gateway Timeout")); - } + assertThat(SocketClientUtil.isSocketReadyToRead(socket)) + .as("Expected connection to proxy server to be open and readable") + .isTrue(); + assertThat(SocketClientUtil.isSocketReadyToWrite(socket)) + .as("Expected connection to proxy server to be open and writable") + .isTrue(); + } + + /** Tests that the proxy does not close the connection after a 502 Bad Gateway response. */ + @Test + @Timeout(25) + public void testBadGatewayDoesNotCloseConnection() throws IOException, InterruptedException { + mockServer.stubFor( + get(urlEqualTo("/success")) + .willReturn( + aResponse().withStatus(200).withBody("success").withHeader("Content-Length", "7"))); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + log.info("Started proxy server {}", proxyServer.getListenAddress()); + socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + + String badGatewayGet = "GET http://localhost:0/success HTTP/1.1\r\n" + "\r\n"; + + // send the same request twice over the same connection + for (int i = 1; i <= 2; i++) { + log.debug("#{} Sending to socket {} packet '{}'...", i, socket.getLocalPort(), badGatewayGet); + SocketClientUtil.writeStringToSocket(badGatewayGet, socket); + + // wait a bit to allow the proxy server to respond + Thread.sleep(1500); + + log.debug("#{} Reading from socket {}...", i, socket.getLocalPort()); + String response = SocketClientUtil.readStringFromSocket(socket); + + assertThat(response) + .as("Expected to receive an HTTP 200 from the server (iteration: %s)", i) + .startsWith("HTTP/1.1 502 Bad Gateway"); + } - assertTrue("Expected connection to proxy server to be open and readable", SocketClientUtil.isSocketReadyToRead(socket)); - assertTrue("Expected connection to proxy server to be open and writable", SocketClientUtil.isSocketReadyToWrite(socket)); + assertThat(SocketClientUtil.isSocketReadyToRead(socket)) + .as("Expected connection to proxy server to be open and readable") + .isTrue(); + assertThat(SocketClientUtil.isSocketReadyToWrite(socket)) + .as("Expected connection to proxy server to be open and writable") + .isTrue(); + } + + /** Tests that the proxy does not close the connection after a 504 Gateway Timeout response. */ + @Test + @Timeout(25) + public void testGatewayTimeoutDoesNotCloseConnection() throws IOException { + mockServer.stubFor( + get(urlEqualTo("/success")) + .willReturn(aResponse().withStatus(200).withFixedDelay(10000).withBody("success"))); + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withIdleConnectionTimeout(Duration.ofSeconds(2)) + .withPort(0) + .start(); + log.info("Started proxy server {}", proxyServer.getListenAddress()); + socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + + String successfulGet = + "GET http://localhost:" + + mockServerPort + + "/success HTTP/1.1\r\n" + + "Host: localhost:" + + mockServerPort + + "\r\n" + + "\r\n"; + + // send the same request twice over the same connection + for (int i = 1; i <= 2; i++) { + log.debug("#{} Sending to socket {} packet '{}'...", i, socket.getLocalPort(), successfulGet); + SocketClientUtil.writeStringToSocket(successfulGet, socket); + + log.debug("#{} Reading from socket {}...", i, socket.getLocalPort()); + String response = SocketClientUtil.readStringFromSocket(socket); + + // match the whole response to make sure that it's not repeated + assertThat(response) + .as("The response is repeated:") + .isEqualTo( + "HTTP/1.1 504 Gateway Timeout\r\n" + + "content-length: 15\r\n" + + "content-type: text/html; charset=utf-8\r\n" + + "\r\n" + + "Gateway Timeout"); } - /** - * Tests that the proxy does not close the connection by default after a short-circuit response. - */ - @Test - public void testShortCircuitResponseDoesNotCloseConnectionByDefault() throws IOException, InterruptedException { - mockServer.when(request() - .withMethod("GET") - .withPath("/success"), - Times.exactly(1)) - .respond(response() - .withStatusCode(500) - .withBody("this response should never be sent")); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - HttpResponse shortCircuitResponse = new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - HttpUtil.setContentLength(shortCircuitResponse, 0); - return shortCircuitResponse; - } else { - return null; - } - } - }; - } + assertThat(SocketClientUtil.isSocketReadyToRead(socket)) + .as("Expected connection to proxy server to be open and readable") + .isTrue(); + assertThat(SocketClientUtil.isSocketReadyToWrite(socket)) + .as("Expected connection to proxy server to be open and writable") + .isTrue(); + } + + /** + * Tests that the proxy does not close the connection by default after a short-circuit response. + */ + @Test + public void testShortCircuitResponseDoesNotCloseConnectionByDefault() + throws IOException, InterruptedException { + mockServer.stubFor( + get(urlEqualTo("/success")) + .willReturn( + aResponse().withStatus(500).withBody("this response should never be sent"))); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest) { + HttpResponse shortCircuitResponse = + new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); + HttpUtil.setContentLength(shortCircuitResponse, 0); + return shortCircuitResponse; + } else { + return null; + } + } + }; + } }; - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .start(); - - socket = SocketClientUtil.getSocketToProxyServer(proxyServer); - - String successfulGet = "GET http://localhost:" + mockServerPort + "/success HTTP/1.1\r\n" - + "\r\n"; - - // send the same request twice over the same connection - for (int i = 1; i <= 2; i++) { - SocketClientUtil.writeStringToSocket(successfulGet, socket); - - // wait a bit to allow the proxy server to respond - Thread.sleep(750); - - String response = SocketClientUtil.readStringFromSocket(socket); - - assertThat("Expected to receive an HTTP 200 from the server (iteration: " + i + ")", response, startsWith("HTTP/1.1 200 OK")); - } - - assertTrue("Expected connection to proxy server to be open and readable", SocketClientUtil.isSocketReadyToRead(socket)); - assertTrue("Expected connection to proxy server to be open and writable", SocketClientUtil.isSocketReadyToWrite(socket)); + proxyServer = + DefaultHttpProxyServer.bootstrap().withPort(0).withFiltersSource(filtersSource).start(); + log.info("Started proxy server {}", proxyServer.getListenAddress()); + socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + + String successfulGet = + "GET http://localhost:" + + mockServerPort + + "/success HTTP/1.1\r\n" + + "Host: localhost:" + + mockServerPort + + "\r\n" + + "\r\n"; + + // send the same request twice over the same connection + for (int i = 1; i <= 2; i++) { + log.debug("#{} Sending to socket {} packet '{}'...", i, socket.getLocalPort(), successfulGet); + SocketClientUtil.writeStringToSocket(successfulGet, socket); + + // wait a bit to allow the proxy server to respond + Thread.sleep(750); + + log.debug("#{} Reading from socket {}...", i, socket.getLocalPort()); + String response = SocketClientUtil.readStringFromSocket(socket); + + assertThat(response) + .as("Expected to receive an HTTP 200 from the server (iteration: %s)", i) + .startsWith("HTTP/1.1 200 OK"); } - /** - * Tests that the proxy will close the connection after a short circuit response if the short circuit response - * contains a Connection: close header. - */ - @Test - public void testShortCircuitResponseCanCloseConnection() throws IOException, InterruptedException { - mockServer.when(request() - .withMethod("GET") - .withPath("/success"), - Times.exactly(1)) - .respond(response() - .withStatusCode(500) - .withBody("this response should never be sent")); - - HttpFiltersSource filtersSource = new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest(HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - HttpResponse shortCircuitResponse = new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - HttpUtil.setContentLength(shortCircuitResponse, 0); - HttpUtil.setKeepAlive(shortCircuitResponse, false); - return shortCircuitResponse; - } else { - return null; - } - } - }; - } + assertThat(SocketClientUtil.isSocketReadyToRead(socket)) + .as("Expected connection to proxy server to be open and readable") + .isTrue(); + assertThat(SocketClientUtil.isSocketReadyToWrite(socket)) + .as("Expected connection to proxy server to be open and writable") + .isTrue(); + } + + /** + * Tests that the proxy will close the connection after a short circuit response if the short + * circuit response contains a Connection: close header. + */ + @Test + public void testShortCircuitResponseCanCloseConnection() + throws IOException, InterruptedException { + mockServer.stubFor( + get(urlEqualTo("/success")) + .willReturn( + aResponse().withStatus(500).withBody("this response should never be sent"))); + + HttpFiltersSource filtersSource = + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest) { + HttpResponse shortCircuitResponse = + new DefaultFullHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); + HttpUtil.setContentLength(shortCircuitResponse, 0); + HttpUtil.setKeepAlive(shortCircuitResponse, false); + return shortCircuitResponse; + } else { + return null; + } + } + }; + } }; - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(filtersSource) - .start(); - - socket = SocketClientUtil.getSocketToProxyServer(proxyServer); - - String successfulGet = "GET http://localhost:" + mockServerPort + "/success HTTP/1.1\r\n" - + "\r\n"; - - // only send this request once, since we expect the short circuit response to close the connection - SocketClientUtil.writeStringToSocket(successfulGet, socket); - - // wait a bit to allow the proxy server to respond - Thread.sleep(750); - - String response = SocketClientUtil.readStringFromSocket(socket); - - assertThat("Expected to receive an HTTP 200 from the server", response, startsWith("HTTP/1.1 200 OK")); - - assertFalse("Expected connection to proxy server to be closed", SocketClientUtil.isSocketReadyToRead(socket)); - assertFalse("Expected connection to proxy server to be closed", SocketClientUtil.isSocketReadyToWrite(socket)); - } + proxyServer = + DefaultHttpProxyServer.bootstrap().withPort(0).withFiltersSource(filtersSource).start(); + log.info("Started proxy server {}", proxyServer.getListenAddress()); + socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + + String successfulGet = + "GET http://localhost:" + + mockServerPort + + "/success HTTP/1.1\r\n" + + "Host: localhost:" + + mockServerPort + + "\r\n" + + "\r\n"; + + // only send this request once, since we expect the short circuit response to + // close the connection + log.debug("Sending to socket {} packet '{}'...", socket.getLocalPort(), successfulGet); + SocketClientUtil.writeStringToSocket(successfulGet, socket); + + // wait a bit to allow the proxy server to respond + Thread.sleep(750); + + log.debug("Reading from socket {}...", socket.getLocalPort()); + String response = SocketClientUtil.readStringFromSocket(socket); + + assertThat(response) + .as("Expected to receive an HTTP 200 from the server") + .startsWith("HTTP/1.1 200 OK"); + + assertThat(SocketClientUtil.isSocketReadyToRead(socket)) + .as("Expected connection to proxy server to be closed") + .isFalse(); + assertThat(SocketClientUtil.isSocketReadyToWrite(socket)) + .as("Expected connection to proxy server to be closed") + .isFalse(); + } } - diff --git a/src/test/java/org/littleshoot/proxy/LauncherTest.java b/src/test/java/org/littleshoot/proxy/LauncherTest.java new file mode 100644 index 00000000..b80c32c5 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/LauncherTest.java @@ -0,0 +1,551 @@ +package org.littleshoot.proxy; + +import static java.lang.System.nanoTime; +import static java.util.concurrent.TimeUnit.SECONDS; +import static org.assertj.core.api.Assertions.assertThat; +import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; +import static org.junit.jupiter.api.Assertions.assertSame; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +/** Unit tests for the Launcher class, specifically testing the start method. */ +class LauncherTest { + private Launcher launcher; + private final int port = 0; + + @BeforeEach + void setUp() { + launcher = new Launcher(); + } + + @AfterEach + void tearDown() { + launcher.stop(); + } + + /** + * Test that the start method handles the help option correctly. Should print help and exit + * gracefully without starting the server. + */ + @Test + void testStartWithHelpOption() { + // Given + String[] args = {"--help"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles valid port option. */ + @Test + void testStartWithValidPortOption() { + + // Given + String[] args = {"--port", "" + port}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid port option gracefully. */ + @Test + void testStartWithInvalidPortOption() { + // Given + String[] args = {"--port", "invalid"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles MITM option. */ + @Test + void testStartWithMitmOption() { + // Given + String[] args = { + "--port", + "" + port, + "--mitm", + "--ssl_clients_keystore_path", + "target/testStartWithMitmOption_keystore.jks" + }; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles DNSSEC option. */ + @Test + void testStartWithDnssecOption() { + // Given + String[] args = {"--port", "" + port, "--dnssec", "true"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid DNSSEC option gracefully. */ + @Test + void testStartWithInvalidDnssecOption() { + // Given + String[] args = {"--dnssec", "invalid"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles config file option. */ + @Test + void testStartWithConfigOption() { + // Given + String[] args = {"--port", "" + port, "--config", "src/test/resources/littleproxy.properties"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + @Test + void testStartWithConfigOptionWithAnAbsentPropertiesFile() { + // Given + String[] args = {"--port", "" + port, "--config", "src/test/resources/notfound.properties"}; + + // When/Then - should not throw exception + assertThrows(IllegalArgumentException.class, () -> launcher.start(args)); + } + + @Test + void testStartWithConfigOptionWithADirectoryFile() { + // Given + String[] args = {"--port", "" + port, "--config", "src/test/resources"}; + + // When/Then - should not throw exception + assertThrows(IllegalArgumentException.class, () -> launcher.start(args)); + } + + /** Test that the start method handles throttling options. */ + @Test + void testStartWithThrottlingOptions() { + // Given + String[] args = { + "--port", + "" + port, + "--throttle_read_bytes_per_second", + "1000", + "--throttle_write_bytes_per_second", + "2000" + }; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles activity log format option. */ + @Test + void testStartWithActivityLogFormat() { + // Given + String[] args = {"--port", "" + port, "--activity_log_format", "CLF"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid activity log format gracefully. */ + @Test + void testStartWithInvalidActivityLogFormat() { + // Given + String[] args = {"--activity_log_format", "INVALID"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method works with no arguments (default configuration). */ + @Test + void testStartWithNoArgs() { + // Given + String[] args = {"--port", "" + port}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles extra/unrecognized arguments by throwing exception. */ + @Test + void testStartWithExtraArguments() { + // Given + String[] args = {"extra", "arguments"}; + + // When/Then - should throw exception for unrecognized arguments + assertThrows(IllegalArgumentException.class, () -> launcher.start(args)); + } + + /** + * Test that the start method handles server mode (though it will hang, so we test in separate + * thread). + */ + @Test + void testStartWithServerOption() throws InterruptedException { + // Given + String[] args = {"--server"}; + + // When - run in separate thread to avoid hanging + Thread testThread = new Thread(() -> launcher.start(args)); + testThread.start(); + waitForServerToStart(); + assertThat(launcher.isRunning()).isTrue(); + + testThread.interrupt(); + testThread.join(1000); + + // Then - thread should terminate + waitForServerToStop(); + assertThat(launcher.isRunning()).isFalse(); + assertSame(Thread.State.TERMINATED, testThread.getState()); + } + + /** Test that the start method handles NIC option. */ + @Test + void testStartWithNicOption() { + // Given + String[] args = {"--port", "" + port, "--nic", "localhost"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles name option. */ + @Test + void testStartWithNameOption() { + // Given + String[] args = {"--port", "" + port, "--name", "TestProxy"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles address option. */ + @Test + void testStartWithAddressOption() { + // Given + String[] args = {"--port", "" + port, "--address", "127.0.0.1:" + port}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles proxy alias option. */ + @Test + void testStartWithProxyAliasOption() { + // Given + String[] args = {"--port", "" + port, "--proxy_alias", "test-alias"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles allow local only option. */ + @Test + void testStartWithAllowLocalOnlyOption() { + // Given + String[] args = {"--port", "" + port, "--allow_local_only", "true"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid allow local only option gracefully. */ + @Test + void testStartWithInvalidAllowLocalOnlyOption() { + // Given + String[] args = {"--port", "" + port, "--allow_local_only", "invalid"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles authenticate SSL clients option. */ + @Test + void testStartWithAuthenticateSslClientsOption() { + // Given + String[] args = { + "--port", + "" + port, + "--authenticate_ssl_clients", + "true", + "--ssl_clients_keystore_path", + "target/testStartWithAuthenticateSslClientsOption_keystore.jks" + }; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles SSL clients trust all servers option. */ + @Test + void testStartWithSslClientsTrustAllServersOption() { + // Given + String[] args = { + "--port", + "" + port, + "--authenticate_ssl_clients", + "true", + "--ssl_clients_trust_all_servers", + "true", + "--ssl_clients_keystore_path", + "target/testStartWithSslClientsTrustAllServersOption_keystore.jks" + }; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles SSL clients send certs option. */ + @Test + void testStartWithSslClientsSendCertsOption() { + // Given + String[] args = { + "--port", + "" + port, + "--authenticate_ssl_clients", + "true", + "--ssl_clients_send_certs", + "true", + "--ssl_clients_keystore_path", + "target/testStartWithSslClientsSendCertsOption_keystore.jks" + }; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles transparent option. */ + @Test + void testStartWithTransparentOption() { + // Given + String[] args = {"--port", "" + port, "--transparent", "true"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid transparent option gracefully. */ + @Test + void testStartWithInvalidTransparentOption() { + // Given + String[] args = {"--port", "" + port, "--transparent", "invalid"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles allow requests to origin server option. */ + @Test + void testStartWithAllowRequestsToOriginServerOption() { + // Given + String[] args = {"--port", "" + port, "--allow_requests_to_origin_server", "true"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** + * Test that the start method handles invalid allow requests to origin server option gracefully. + */ + @Test + void testStartWithInvalidAllowRequestsToOriginServerOption() { + // Given + String[] args = {"--port", "" + port, "--allow_requests_to_origin_server", "invalid"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles allow proxy protocol option. */ + @Test + void testStartWithAllowProxyProtocolOption() { + // Given + String[] args = {"--port", "" + port, "--allow_proxy_protocol", "true"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid allow proxy protocol option gracefully. */ + @Test + void testStartWithInvalidAllowProxyProtocolOption() { + // Given + String[] args = {"--port", "" + port, "--allow_proxy_protocol", "invalid"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles send proxy protocol option. */ + @Test + void testStartWithSendProxyProtocolOption() { + // Given + String[] args = {"--port", "" + port, "--send_proxy_protocol", "true"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid send proxy protocol option gracefully. */ + @Test + void testStartWithInvalidSendProxyProtocolOption() { + // Given + String[] args = {"--port", "" + port, "--send_proxy_protocol", "invalid"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles client to proxy worker threads option. */ + @Test + void testStartWithClientToProxyWorkerThreadsOption() { + // Given + String[] args = {"--port", "" + port, "--client_to_proxy_worker_threads", "4"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid client to proxy worker threads option. */ + @Test + void testStartWithInvalidClientToProxyWorkerThreadsOption() { + // Given + String[] args = {"--client_to_proxy_worker_threads", "invalid"}; + + // When/Then - should throw exception for invalid numeric value + assertThrows(NumberFormatException.class, () -> launcher.start(args)); + } + + /** Test that the start method handles proxy to server worker threads option. */ + @Test + void testStartWithProxyToServerWorkerThreadsOption() { + // Given + String[] args = {"--port", "" + port, "--proxy_to_server_worker_threads", "4"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid proxy to server worker threads option. */ + @Test + void testStartWithInvalidProxyToServerWorkerThreadsOption() { + // Given + String[] args = {"--proxy_to_server_worker_threads", "invalid"}; + + // When/Then - should throw exception for invalid numeric value + assertThrows(NumberFormatException.class, () -> launcher.start(args)); + } + + /** Test that the start method handles acceptor threads option. */ + @Test + void testStartWithAcceptorThreadsOption() { + // Given + String[] args = {"--port", "" + port, "--acceptor_threads", "2"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid acceptor threads option. */ + @Test + void testStartWithInvalidAcceptorThreadsOption() { + // Given + String[] args = {"--acceptor_threads", "invalid"}; + + // When/Then - should throw exception for invalid numeric value + assertThrows(NumberFormatException.class, () -> launcher.start(args)); + } + + /** Test that the start method handles log config option. */ + @Test + void testStartWithLogConfigOption() { + // Given + String[] args = {"--port", "" + port, "--log_config", "src/test/resources/log4j.xml"}; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles invalid log config option. */ + @Test + void testStartWithInvalidLogConfigOption() { + // Given + String[] args = {"--port", "" + port, "--log_config", "nonexistent.xml"}; + + // When/Then - should not throw exception and handle gracefully + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles multiple options combination. */ + @Test + void testStartWithMultipleOptions() { + // Given + String[] args = { + "--port", + "" + port, + "--mitm", + "--dnssec", + "true", + "--name", + "TestProxy", + "--allow_local_only", + "true", + "--ssl_clients_keystore_path", + "target/testStartWithMultipleOptions_keystore.jks" + }; + + // When/Then - should not throw exception + assertDoesNotThrow(() -> launcher.start(args)); + } + + /** Test that the start method handles missing required values for options. */ + @Test + void testStartWithMissingRequiredValues() { + // Given - port option without value + String[] args = {"--port"}; + + // When/Then - should throw exception for missing required value + assertThrows(IllegalArgumentException.class, () -> launcher.start(args)); + } + + /** Test that the start method handles invalid port values. */ + @Test + void testStartWithInvalidPortValues() { + // Given - port with negative value + String[] args = {"--port", "-1"}; + + // When/Then - should throw exception for invalid port value + assertThrows(IllegalArgumentException.class, () -> launcher.start(args)); + } + + /** Test that the start method handles port value that is too large. */ + @Test + void testStartWithPortValueTooLarge() { + // Given - port with very large value + String[] args = {"--port", "999999"}; + + // When/Then - should throw exception for port out of range + assertThrows(IllegalArgumentException.class, () -> launcher.start(args)); + } + + private void waitForServerToStart() throws InterruptedException { + waitForServer(true); + } + + private void waitForServerToStop() throws InterruptedException { + waitForServer(false); + } + + private void waitForServer(boolean alive) throws InterruptedException { + long timeout = SECONDS.toNanos(10); + for (long start = nanoTime(); launcher.isRunning() != alive && nanoTime() - start < timeout; ) { + Thread.sleep(10); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/MITMUsernamePasswordAuthenticatingProxyTest.java b/src/test/java/org/littleshoot/proxy/MITMUsernamePasswordAuthenticatingProxyTest.java index 48369cdd..c2072a9d 100644 --- a/src/test/java/org/littleshoot/proxy/MITMUsernamePasswordAuthenticatingProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MITMUsernamePasswordAuthenticatingProxyTest.java @@ -1,25 +1,22 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedMitmManager; +import org.littleshoot.proxy.extras.TestMitmManager; -/** - * Tests a single proxy that requires username/password authentication and that - * uses MITM. - */ -public class MITMUsernamePasswordAuthenticatingProxyTest extends - UsernamePasswordAuthenticatingProxyTest - implements ProxyAuthenticator { - @Override - protected void setUp() { - this.proxyServer = bootstrapProxy() - .withPort(0) - .withProxyAuthenticator(this) - .withManInTheMiddle(new SelfSignedMitmManager()) - .start(); - } +/** Tests a single proxy that requires username/password authentication and that uses MITM. */ +public class MITMUsernamePasswordAuthenticatingProxyTest + extends UsernamePasswordAuthenticatingProxyTest implements ProxyAuthenticator { + @Override + protected void setUp() { + proxyServer = + bootstrapProxy() + .withPort(0) + .withProxyAuthenticator(this) + .withManInTheMiddle(new TestMitmManager()) + .start(); + } - @Override - protected boolean isMITM() { - return true; - } + @Override + protected boolean isMITM() { + return true; + } } diff --git a/src/test/java/org/littleshoot/proxy/MessageTerminationTest.java b/src/test/java/org/littleshoot/proxy/MessageTerminationTest.java index a7f05737..312d9a56 100644 --- a/src/test/java/org/littleshoot/proxy/MessageTerminationTest.java +++ b/src/test/java/org/littleshoot/proxy/MessageTerminationTest.java @@ -1,201 +1,201 @@ package org.littleshoot.proxy; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.createProxiedHttpClient; + +import com.github.tomakehurst.wiremock.WireMockServer; +import java.util.Objects; import org.apache.http.Header; import org.apache.http.HttpResponse; import org.apache.http.client.HttpClient; import org.apache.http.client.methods.HttpGet; import org.apache.http.client.methods.HttpHead; import org.apache.http.util.EntityUtils; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; -import org.mockserver.model.ConnectionOptions; - -import static org.hamcrest.Matchers.emptyArray; -import static org.hamcrest.Matchers.greaterThanOrEqualTo; -import static org.junit.Assert.*; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; - -public class MessageTerminationTest { - private ClientAndServer mockServer; - private int mockServerPort; - private HttpProxyServer proxyServer; - - @Before - public void setUp() { - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); - } - - @After - public void tearDown() { - if (mockServer != null) { - mockServer.stop(); - } - - if (proxyServer != null) { - proxyServer.abort(); - } - } - @Test - public void testResponseWithoutTerminationIsChunked() throws Exception { - // set up the server so that it indicates the end of the response by closing the connection. the proxy - // should automatically add the Transfer-Encoding: chunked header when sending to the client. - mockServer.when(request() - .withMethod("GET") - .withPath("/"), - Times.unlimited()) - .respond(response() - .withStatusCode(200) - .withBody("Success!") - .withConnectionOptions(new ConnectionOptions() - .withCloseSocket(true) - .withSuppressConnectionHeader(true) - .withSuppressContentLengthHeader(true)) - ); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - int proxyServerPort = proxyServer.getListenAddress().getPort(); - - HttpClient httpClient = TestUtils.createProxiedHttpClient(proxyServerPort); - HttpResponse response = httpClient.execute(new HttpGet("http://127.0.0.1:" + mockServerPort + "/")); - - assertEquals("Expected to receive a 200 from the server", 200, response.getStatusLine().getStatusCode()); - - // verify the Transfer-Encoding header was added - Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); - assertThat("Expected to see a Transfer-Encoding header", transferEncodingHeaders.length, greaterThanOrEqualTo(1)); - String transferEncoding = transferEncodingHeaders[0].getValue(); - assertEquals("Expected Transfer-Encoding to be chunked", "chunked", transferEncoding); - - String bodyString = EntityUtils.toString(response.getEntity(), "ISO-8859-1"); - response.getEntity().getContent().close(); - - assertEquals("Success!", bodyString); +public final class MessageTerminationTest { + private WireMockServer mockServer; + private int mockServerPort; + private HttpProxyServer proxyServer; + + @BeforeEach + void setUp() { + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + } + + @AfterEach + void tearDown() { + if (mockServer != null) { + mockServer.stop(); } - @Test - public void testResponseWithContentLengthNotModified() throws Exception { - // the proxy should not modify the response since it contains a Content-Length header. - mockServer.when(request() - .withMethod("GET") - .withPath("/"), - Times.unlimited()) - .respond(response() - .withStatusCode(200) - .withBody("Success!") - .withConnectionOptions(new ConnectionOptions() - .withCloseSocket(true) - .withSuppressConnectionHeader(true)) - ); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - int proxyServerPort = proxyServer.getListenAddress().getPort(); - - HttpClient httpClient = TestUtils.createProxiedHttpClient(proxyServerPort); - HttpResponse response = httpClient.execute(new HttpGet("http://127.0.0.1:" + mockServerPort + "/")); - - assertEquals("Expected to receive a 200 from the server", 200, response.getStatusLine().getStatusCode()); - - // verify the Transfer-Encoding header was NOT added - Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); - assertThat("Did not expect to see a Transfer-Encoding header", transferEncodingHeaders, emptyArray()); - - String bodyString = EntityUtils.toString(response.getEntity(), "ISO-8859-1"); - response.getEntity().getContent().close(); - - assertEquals("Success!", bodyString); + if (proxyServer != null) { + proxyServer.abort(); } - - @Test - public void testFilterAddsContentLength() throws Exception { - // when a filter with buffering is added to the filter chain, the aggregated FullHttpResponse should - // automatically have a Content-Length header - mockServer.when(request() - .withMethod("GET") - .withPath("/"), - Times.unlimited()) - .respond(response() - .withStatusCode(200) - .withBody("Success!") - .withConnectionOptions(new ConnectionOptions() - .withCloseSocket(true) - .withSuppressConnectionHeader(true) - .withSuppressContentLengthHeader(true)) - ); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withFiltersSource(new HttpFiltersSourceAdapter() { - @Override - public int getMaximumResponseBufferSizeInBytes() { - return 100000; - } + } + + @Test + public void testResponseWithoutTerminationIsChunked() throws Exception { + // set up the server so that it indicates the end of the response by closing the + // connection. the proxy + // should automatically add the Transfer-Encoding: chunked header when sending + // to the client. + mockServer.stubFor( + get(urlEqualTo("/")) + .willReturn( + aResponse() + .withStatus(200) + .withBody("Success!") + .withHeader("Connection", "close"))); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + int proxyServerPort = proxyServer.getListenAddress().getPort(); + + HttpClient httpClient = createProxiedHttpClient(proxyServerPort); + HttpResponse response = + httpClient.execute(new HttpGet("http://127.0.0.1:" + mockServerPort + "/")); + + assertThat(Objects.requireNonNull(response.getStatusLine()).getStatusCode()) + .as("Expected to receive a 200 from the server") + .isEqualTo(200); + + // verify the Transfer-Encoding header was added + Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); + assertThat(transferEncodingHeaders) + .as("Expected to see a Transfer-Encoding header") + .isNotEmpty(); + String transferEncoding = Objects.requireNonNull(transferEncodingHeaders[0].getValue()); + assertThat(transferEncoding) + .as("Expected Transfer-Encoding to be chunked") + .isEqualTo("chunked"); + + String bodyString = EntityUtils.toString(response.getEntity(), "ISO-8859-1"); + response.getEntity().getContent().close(); + + assertThat(bodyString).isEqualTo("Success!"); + } + + @Test + public void testResponseWithContentLengthNotModified() throws Exception { + // the proxy should not modify the response since it contains a Content-Length + // header. + mockServer.stubFor( + get(urlEqualTo("/")) + .willReturn( + aResponse() + .withStatus(200) + .withBody("Success!") + .withHeader("Connection", "close") + .withHeader("Content-Length", "8"))); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + int proxyServerPort = proxyServer.getListenAddress().getPort(); + + HttpClient httpClient = createProxiedHttpClient(proxyServerPort); + HttpResponse response = + httpClient.execute(new HttpGet("http://127.0.0.1:" + mockServerPort + "/")); + + assertThat(Objects.requireNonNull(response.getStatusLine()).getStatusCode()) + .as("Expected to receive a 200 from the server") + .isEqualTo(200); + + // verify the Transfer-Encoding header was NOT added + Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); + assertThat(transferEncodingHeaders) + .as("Did not expect to see a Transfer-Encoding header") + .isEmpty(); + + String bodyString = EntityUtils.toString(response.getEntity(), "ISO-8859-1"); + response.getEntity().getContent().close(); + + assertThat(bodyString).isEqualTo("Success!"); + } + + @Test + public void testFilterAddsContentLength() throws Exception { + // when a filter with buffering is added to the filter chain, the aggregated + // FullHttpResponse should + // automatically have a Content-Length header + mockServer.stubFor( + get(urlEqualTo("/")) + .willReturn( + aResponse() + .withStatus(200) + .withBody("Success!") + .withHeader("Connection", "close") + .withHeader("Content-Length", "8"))); + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withFiltersSource( + new HttpFiltersSourceAdapter() { + @Override + public int getMaximumResponseBufferSizeInBytes() { + return 100000; + } }) - .withPort(0) - .start(); - int proxyServerPort = proxyServer.getListenAddress().getPort(); - - - HttpClient httpClient = TestUtils.createProxiedHttpClient(proxyServerPort); - HttpResponse response = httpClient.execute(new HttpGet("http://127.0.0.1:" + mockServerPort + "/")); - - assertEquals("Expected to receive a 200 from the server", 200, response.getStatusLine().getStatusCode()); - - // verify the Transfer-Encoding header was NOT added - Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); - assertThat("Did not expect to see a Transfer-Encoding header", transferEncodingHeaders, emptyArray()); - - Header[] contentLengthHeaders = response.getHeaders("Content-Length"); - assertThat("Expected to see a Content-Length header", contentLengthHeaders.length, greaterThanOrEqualTo(1)); - - String bodyString = EntityUtils.toString(response.getEntity(), "ISO-8859-1"); - response.getEntity().getContent().close(); - - assertEquals("Success!", bodyString); - } - - @Test - public void testResponseToHEADNotModified() throws Exception { - // the proxy should not modify the response since it is an HTTP HEAD request - mockServer.when(request() - .withMethod("HEAD") - .withPath("/"), - Times.unlimited()) - .respond(response() - .withStatusCode(200) - .withConnectionOptions(new ConnectionOptions() - .withCloseSocket(false) - .withSuppressConnectionHeader(true) - .withSuppressContentLengthHeader(true)) - ); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - int proxyServerPort = proxyServer.getListenAddress().getPort(); - - HttpClient httpClient = TestUtils.createProxiedHttpClient(proxyServerPort); - HttpResponse response = httpClient.execute(new HttpHead("http://127.0.0.1:" + mockServerPort + "/")); - - assertEquals("Expected to receive a 200 from the server", 200, response.getStatusLine().getStatusCode()); - - // verify the Transfer-Encoding header was NOT added - Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); - assertThat("Did not expect to see a Transfer-Encoding header", transferEncodingHeaders, emptyArray()); - - // verify the Content-Length header was not added - Header[] contentLengthHeaders = response.getHeaders("Content-Length"); - assertThat("Did not expect to see a Content-Length header", contentLengthHeaders, emptyArray()); - - assertNull("Expected response to HEAD to have no entity body", response.getEntity()); - } + .withPort(0) + .start(); + int proxyServerPort = proxyServer.getListenAddress().getPort(); + + HttpClient httpClient = createProxiedHttpClient(proxyServerPort); + HttpResponse response = + httpClient.execute(new HttpGet("http://127.0.0.1:" + mockServerPort + "/")); + + assertThat(Objects.requireNonNull(response.getStatusLine()).getStatusCode()) + .as("Expected to receive a 200 from the server") + .isEqualTo(200); + + // verify the Transfer-Encoding header was NOT added + Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); + assertThat(transferEncodingHeaders) + .as("Did not expect to see a Transfer-Encoding header") + .isEmpty(); + + Header[] contentLengthHeaders = response.getHeaders("Content-Length"); + assertThat(contentLengthHeaders).as("Expected to see a Content-Length header").isNotEmpty(); + + String bodyString = EntityUtils.toString(response.getEntity(), "ISO-8859-1"); + response.getEntity().getContent().close(); + + assertThat(bodyString).isEqualTo("Success!"); + } + + @Test + public void testResponseToHEADNotModified() throws Exception { + // the proxy should not modify the response since it is an HTTP HEAD request + mockServer.stubFor(head(urlEqualTo("/")).willReturn(aResponse().withStatus(200))); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); + int proxyServerPort = proxyServer.getListenAddress().getPort(); + + HttpClient httpClient = createProxiedHttpClient(proxyServerPort); + HttpResponse response = + httpClient.execute(new HttpHead("http://127.0.0.1:" + mockServerPort + "/")); + + assertThat(Objects.requireNonNull(response.getStatusLine()).getStatusCode()) + .as("Expected to receive a 200 from the server") + .isEqualTo(200); + + // verify the Transfer-Encoding header was NOT added + Header[] transferEncodingHeaders = response.getHeaders("Transfer-Encoding"); + assertThat(transferEncodingHeaders) + .as("Did not expect to see a Transfer-Encoding header") + .isEmpty(); + + // verify the Content-Length header was not added + Header[] contentLengthHeaders = response.getHeaders("Content-Length"); + assertThat(contentLengthHeaders).as("Did not expect to see a Content-Length header").isEmpty(); + + assertThat(response.getEntity()) + .as("Expected response to HEAD to have no entity body") + .isNull(); + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmProxyTest.java index 68b4a000..c33230fd 100644 --- a/src/test/java/org/littleshoot/proxy/MitmProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MitmProxyTest.java @@ -1,157 +1,150 @@ package org.littleshoot.proxy; -import io.netty.handler.codec.http.*; -import org.littleshoot.proxy.extras.SelfSignedMitmManager; +import static java.nio.charset.StandardCharsets.UTF_8; +import static org.assertj.core.api.Assertions.assertThat; -import java.nio.charset.Charset; +import io.netty.handler.codec.http.*; import java.util.HashSet; import java.util.Set; - -import static org.hamcrest.Matchers.hasItem; -import static org.junit.Assert.assertEquals; -import static org.junit.Assert.assertThat; - -/** - * Tests just a single basic proxy running as a man in the middle. - */ -public class MitmProxyTest extends BaseProxyTest { - private Set requestPreMethodsSeen = new HashSet<>(); - private Set requestPostMethodsSeen = new HashSet<>(); - private StringBuilder responsePreBody = new StringBuilder(); - private StringBuilder responsePostBody = new StringBuilder(); - private Set responsePreOriginalRequestMethodsSeen = new HashSet<>(); - private Set responsePostOriginalRequestMethodsSeen = new HashSet<>(); - - @Override - protected void setUp() { - this.proxyServer = bootstrapProxy() - .withPort(0) - .withManInTheMiddle(new SelfSignedMitmManager()) - .withFiltersSource(new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest( - HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - requestPreMethodsSeen - .add(((HttpRequest) httpObject) - .method()); - } - return null; - } - - @Override - public HttpResponse proxyToServerRequest( - HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - requestPostMethodsSeen - .add(((HttpRequest) httpObject) - .method()); - } - return null; - } - - @Override - public HttpObject serverToProxyResponse( - HttpObject httpObject) { - if (httpObject instanceof HttpResponse) { - responsePreOriginalRequestMethodsSeen - .add(originalRequest.method()); - } else if (httpObject instanceof HttpContent) { - responsePreBody.append(((HttpContent) httpObject) - .content().toString( - Charset.forName("UTF-8"))); - } - return httpObject; - } - - @Override - public HttpObject proxyToClientResponse( - HttpObject httpObject) { - if (httpObject instanceof HttpResponse) { - responsePostOriginalRequestMethodsSeen - .add(originalRequest.method()); - } else if (httpObject instanceof HttpContent) { - responsePostBody.append(((HttpContent) httpObject) - .content().toString( - Charset.forName("UTF-8"))); - } - return httpObject; - } - }; - } +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.extras.TestMitmManager; + +/** Tests just a single basic proxy running as a man in the middle. */ +@NullMarked +public final class MitmProxyTest extends BaseProxyTest { + private final Set requestPreMethodsSeen = new HashSet<>(); + private final Set requestPostMethodsSeen = new HashSet<>(); + private final StringBuilder responsePreBody = new StringBuilder(); + private final StringBuilder responsePostBody = new StringBuilder(); + private final Set responsePreOriginalRequestMethodsSeen = new HashSet<>(); + private final Set responsePostOriginalRequestMethodsSeen = new HashSet<>(); + + @Override + protected void setUp() { + proxyServer = + bootstrapProxy() + .withPort(0) + .withManInTheMiddle(new TestMitmManager()) + .withFiltersSource( + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest) { + requestPreMethodsSeen.add(((HttpRequest) httpObject).method()); + } + return null; + } + + @Nullable + @Override + public HttpResponse proxyToServerRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest) { + requestPostMethodsSeen.add(((HttpRequest) httpObject).method()); + } + return null; + } + + @Override + public HttpObject serverToProxyResponse(HttpObject httpObject) { + if (httpObject instanceof HttpResponse) { + responsePreOriginalRequestMethodsSeen.add(originalRequest.method()); + } else if (httpObject instanceof HttpContent) { + responsePreBody.append( + ((HttpContent) httpObject).content().toString(UTF_8)); + } + return httpObject; + } + + @Override + public HttpObject proxyToClientResponse(HttpObject httpObject) { + if (httpObject instanceof HttpResponse) { + responsePostOriginalRequestMethodsSeen.add(originalRequest.method()); + } else if (httpObject instanceof HttpContent) { + responsePostBody.append( + ((HttpContent) httpObject).content().toString(UTF_8)); + } + return httpObject; + } + }; + } }) - .start(); - } - - @Override - protected boolean isMITM() { - return true; - } - - @Override - public void testSimpleGetRequest() throws Exception { - super.testSimpleGetRequest(); - assertMethodSeenInRequestFilters(HttpMethod.GET); - assertMethodSeenInResponseFilters(HttpMethod.GET); - assertResponseFromFiltersMatchesActualResponse(); - } - - @Override - public void testSimpleGetRequestOverHTTPS() throws Exception { - super.testSimpleGetRequestOverHTTPS(); - assertMethodSeenInRequestFilters(HttpMethod.CONNECT); - assertMethodSeenInRequestFilters(HttpMethod.GET); - assertMethodSeenInResponseFilters(HttpMethod.GET); - assertResponseFromFiltersMatchesActualResponse(); - } - - @Override - public void testSimplePostRequest() throws Exception { - super.testSimplePostRequest(); - assertMethodSeenInRequestFilters(HttpMethod.POST); - assertMethodSeenInResponseFilters(HttpMethod.POST); - assertResponseFromFiltersMatchesActualResponse(); - } - - @Override - public void testSimplePostRequestOverHTTPS() throws Exception { - super.testSimplePostRequestOverHTTPS(); - assertMethodSeenInRequestFilters(HttpMethod.CONNECT); - assertMethodSeenInRequestFilters(HttpMethod.POST); - assertMethodSeenInResponseFilters(HttpMethod.POST); - assertResponseFromFiltersMatchesActualResponse(); - } - - private void assertMethodSeenInRequestFilters(HttpMethod method) { - assertThat(method - + " should have been seen in clientToProxyRequest filter", - requestPreMethodsSeen, hasItem(method)); - assertThat(method - + " should have been seen in proxyToServerRequest filter", - requestPostMethodsSeen, hasItem(method)); - } - - private void assertMethodSeenInResponseFilters(HttpMethod method) { - assertThat( - method - + " should have been seen as the original requests's method in serverToProxyResponse filter", - responsePreOriginalRequestMethodsSeen, hasItem(method)); - assertThat( - method - + " should have been seen as the original requests's method in proxyToClientResponse filter", - responsePostOriginalRequestMethodsSeen, hasItem(method)); - } - - private void assertResponseFromFiltersMatchesActualResponse() { - assertEquals( - "Data received through HttpFilters.serverToProxyResponse should match response", - lastResponse, responsePreBody.toString()); - assertEquals( - "Data received through HttpFilters.proxyToClientResponse should match response", - lastResponse, responsePostBody.toString()); - } - + .start(); + } + + @Override + protected boolean isMITM() { + return true; + } + + @Override + public void testSimpleGetRequest() { + super.testSimpleGetRequest(); + assertMethodSeenInRequestFilters(HttpMethod.GET); + assertMethodSeenInResponseFilters(HttpMethod.GET); + assertResponseFromFiltersMatchesActualResponse(); + } + + @Override + public void testSimpleGetRequestOverHTTPS() { + super.testSimpleGetRequestOverHTTPS(); + assertMethodSeenInRequestFilters(HttpMethod.CONNECT); + assertMethodSeenInRequestFilters(HttpMethod.GET); + assertMethodSeenInResponseFilters(HttpMethod.GET); + assertResponseFromFiltersMatchesActualResponse(); + } + + @Override + public void testSimplePostRequest() { + super.testSimplePostRequest(); + assertMethodSeenInRequestFilters(HttpMethod.POST); + assertMethodSeenInResponseFilters(HttpMethod.POST); + assertResponseFromFiltersMatchesActualResponse(); + } + + @Override + public void testSimplePostRequestOverHTTPS() { + super.testSimplePostRequestOverHTTPS(); + assertMethodSeenInRequestFilters(HttpMethod.CONNECT); + assertMethodSeenInRequestFilters(HttpMethod.POST); + assertMethodSeenInResponseFilters(HttpMethod.POST); + assertResponseFromFiltersMatchesActualResponse(); + } + + private void assertMethodSeenInRequestFilters(HttpMethod method) { + assertThat(requestPreMethodsSeen) + .as(method + " should have been seen in clientToProxyRequest filter") + .contains(method); + assertThat(requestPostMethodsSeen) + .as(method + " should have been seen in proxyToServerRequest filter") + .contains(method); + } + + private void assertMethodSeenInResponseFilters(HttpMethod method) { + assertThat(responsePreOriginalRequestMethodsSeen) + .as( + method + + " should have been seen as the original request's method in serverToProxyResponse filter") + .contains(method); + assertThat(responsePostOriginalRequestMethodsSeen) + .as( + method + + " should have been seen as the original request's method in proxyToClientResponse filter") + .contains(method); + } + + private void assertResponseFromFiltersMatchesActualResponse() { + assertThat(lastResponse) + .as(responsePreBody.toString()) + .isEqualTo("Data received through HttpFilters.serverToProxyResponse should match response"); + assertThat(lastResponse) + .as(responsePostBody.toString()) + .isEqualTo("Data received through HttpFilters.proxyToClientResponse should match response"); + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmWithBadClientAuthenticationTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithBadClientAuthenticationTCPChainedProxyTest.java index 56e69935..52ec8caf 100644 --- a/src/test/java/org/littleshoot/proxy/MitmWithBadClientAuthenticationTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MitmWithBadClientAuthenticationTCPChainedProxyTest.java @@ -1,52 +1,38 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -/** - * Tests that clients are authenticated and that if they're missing certs, we - * get an error. - */ -public class MitmWithBadClientAuthenticationTCPChainedProxyTest extends - MitmWithChainedProxyTest { - private final SslEngineSource serverSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - private final SslEngineSource clientSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks", false, false); - - @Override - protected boolean expectBadGatewayForEverything() { +/** Tests that clients are authenticated and that if they're missing certs, we get an error. */ +public final class MitmWithBadClientAuthenticationTCPChainedProxyTest + extends MitmWithChainedProxyTest { + private final SslEngineSource serverSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks"); + private final SslEngineSource clientSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks", false, false); + + @Override + protected boolean expectBadGatewayForEverything() { + return true; + } + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(serverSslEngineSource); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { return true; - } - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(serverSslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return clientSslEngineSource.newSslEngine(); - } - }; - } + } + + @Override + public SSLEngine newSslEngine() { + return clientSslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmWithBadServerAuthenticationTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithBadServerAuthenticationTCPChainedProxyTest.java index 813f7b42..97f0dd1d 100644 --- a/src/test/java/org/littleshoot/proxy/MitmWithBadServerAuthenticationTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MitmWithBadServerAuthenticationTCPChainedProxyTest.java @@ -1,52 +1,38 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -/** - * Tests that servers are authenticated and that if they're missing certs, we - * get an error. - */ -public class MitmWithBadServerAuthenticationTCPChainedProxyTest extends - MitmWithChainedProxyTest { - protected final SslEngineSource serverSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - protected final SslEngineSource clientSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_2.jks"); - - @Override - protected boolean expectBadGatewayForEverything() { +/** Tests that servers are authenticated and that if they're missing certs, we get an error. */ +public final class MitmWithBadServerAuthenticationTCPChainedProxyTest + extends MitmWithChainedProxyTest { + private final SslEngineSource serverSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks"); + private final SslEngineSource clientSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_2.jks"); + + @Override + protected boolean expectBadGatewayForEverything() { + return true; + } + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(serverSslEngineSource); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { return true; - } - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(serverSslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return clientSslEngineSource.newSslEngine(); - } - }; - } + } + + @Override + public SSLEngine newSslEngine() { + return clientSslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmWithChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithChainedProxyTest.java index d1cca41f..c24ecb68 100644 --- a/src/test/java/org/littleshoot/proxy/MitmWithChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MitmWithChainedProxyTest.java @@ -1,178 +1,170 @@ package org.littleshoot.proxy; -import io.netty.handler.codec.http.*; -import org.littleshoot.proxy.extras.SelfSignedMitmManager; +import static java.nio.charset.StandardCharsets.UTF_8; +import static org.assertj.core.api.Assertions.assertThat; -import java.nio.charset.Charset; +import io.netty.handler.codec.http.*; import java.util.HashSet; import java.util.Set; +import org.jspecify.annotations.NonNull; +import org.jspecify.annotations.NullMarked; +import org.jspecify.annotations.Nullable; +import org.littleshoot.proxy.extras.TestMitmManager; -import static org.hamcrest.Matchers.hasItem; -import static org.junit.Assert.assertEquals; -import static org.junit.Assert.assertThat; - -/** - * Tests a proxy that runs as a MITM and which is chained with - * another proxy. - */ +/** Tests a proxy that runs as a MITM and which is chained with another proxy. */ +@NullMarked public class MitmWithChainedProxyTest extends BaseChainedProxyTest { - private Set requestPreMethodsSeen = new HashSet<>(); - private Set requestPostMethodsSeen = new HashSet<>(); - private StringBuilder responsePreBody = new StringBuilder(); - private StringBuilder responsePostBody = new StringBuilder(); - private Set responsePreOriginalRequestMethodsSeen = new HashSet<>(); - private Set responsePostOriginalRequestMethodsSeen = new HashSet<>(); - - @Override - protected void setUp() { - - REQUESTS_SENT_BY_DOWNSTREAM.set(0); - REQUESTS_RECEIVED_BY_UPSTREAM.set(0); - TRANSPORTS_USED.clear(); - this.upstreamProxy = upstreamProxy().start(); - - this.proxyServer = bootstrapProxy() - .withPort(0) - .withChainProxyManager(chainedProxyManager()) - .plusActivityTracker(DOWNSTREAM_TRACKER) - .withManInTheMiddle(new SelfSignedMitmManager()) - .withFiltersSource(new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public HttpResponse clientToProxyRequest( - HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - requestPreMethodsSeen - .add(((HttpRequest) httpObject) - .method()); - } - return null; - } - - @Override - public HttpResponse proxyToServerRequest( - HttpObject httpObject) { - if (httpObject instanceof HttpRequest) { - requestPostMethodsSeen - .add(((HttpRequest) httpObject) - .method()); - } - return null; - } - - @Override - public HttpObject serverToProxyResponse( - HttpObject httpObject) { - if (httpObject instanceof HttpResponse) { - responsePreOriginalRequestMethodsSeen - .add(originalRequest.method()); - } else if (httpObject instanceof HttpContent) { - responsePreBody.append(((HttpContent) httpObject) - .content().toString( - Charset.forName("UTF-8"))); - } - return httpObject; - } - - @Override - public HttpObject proxyToClientResponse( - HttpObject httpObject) { - if (httpObject instanceof HttpResponse) { - responsePostOriginalRequestMethodsSeen - .add(originalRequest.method()); - } else if (httpObject instanceof HttpContent) { - responsePostBody.append(((HttpContent) httpObject) - .content().toString( - Charset.forName("UTF-8"))); - } - return httpObject; - } - }; - } + private final Set requestPreMethodsSeen = new HashSet<>(); + private final Set requestPostMethodsSeen = new HashSet<>(); + private final StringBuilder responsePreBody = new StringBuilder(); + private final StringBuilder responsePostBody = new StringBuilder(); + private final Set responsePreOriginalRequestMethodsSeen = new HashSet<>(); + private final Set responsePostOriginalRequestMethodsSeen = new HashSet<>(); + + @Override + protected final void setUp() { + REQUESTS_SENT_BY_DOWNSTREAM.set(0); + REQUESTS_RECEIVED_BY_UPSTREAM.set(0); + TRANSPORTS_USED.clear(); + upstreamProxy = upstreamProxy().start(); + + proxyServer = + bootstrapProxy() + .withPort(0) + .withChainProxyManager(chainedProxyManager()) + .plusActivityTracker(DOWNSTREAM_TRACKER) + .withManInTheMiddle(new TestMitmManager()) + .withFiltersSource( + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Nullable + @Override + public HttpResponse clientToProxyRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest) { + requestPreMethodsSeen.add(((HttpRequest) httpObject).method()); + } + return null; + } + + @Nullable + @Override + public HttpResponse proxyToServerRequest(@NonNull HttpObject httpObject) { + if (httpObject instanceof HttpRequest) { + requestPostMethodsSeen.add(((HttpRequest) httpObject).method()); + } + return null; + } + + @Override + public HttpObject serverToProxyResponse(HttpObject httpObject) { + if (httpObject instanceof HttpResponse) { + responsePreOriginalRequestMethodsSeen.add(originalRequest.method()); + } else if (httpObject instanceof HttpContent) { + responsePreBody.append( + ((HttpContent) httpObject).content().toString(UTF_8)); + } + return httpObject; + } + + @Override + public HttpObject proxyToClientResponse(HttpObject httpObject) { + if (httpObject instanceof HttpResponse) { + responsePostOriginalRequestMethodsSeen.add(originalRequest.method()); + } else if (httpObject instanceof HttpContent) { + responsePostBody.append( + ((HttpContent) httpObject).content().toString(UTF_8)); + } + return httpObject; + } + }; + } }) - .start(); - } - - @Override - protected boolean isMITM() { - return true; - } - - @Override - public void testSimpleGetRequest() throws Exception { - super.testSimpleGetRequest(); - if (isChained() && !expectBadGatewayForEverything()) { - assertMethodSeenInRequestFilters(HttpMethod.GET); - assertMethodSeenInResponseFilters(HttpMethod.GET); - assertResponseFromFiltersMatchesActualResponse(); - } - } - - @Override - public void testSimpleGetRequestOverHTTPS() throws Exception { - super.testSimpleGetRequestOverHTTPS(); - if (isChained() && !expectBadGatewayForEverything()) { - assertMethodSeenInRequestFilters(HttpMethod.CONNECT); - assertMethodSeenInRequestFilters(HttpMethod.GET); - assertMethodSeenInResponseFilters(HttpMethod.GET); - assertResponseFromFiltersMatchesActualResponse(); - } + .start(); + } + + @Override + protected boolean isMITM() { + return true; + } + + @Override + public void testSimpleGetRequest() { + super.testSimpleGetRequest(); + if (isChained() && !expectBadGatewayForEverything()) { + assertMethodSeenInRequestFilters(HttpMethod.GET); + assertMethodSeenInResponseFilters(HttpMethod.GET); + assertResponseFromFiltersMatchesActualResponse(); } - - @Override - public void testSimplePostRequest() throws Exception { - super.testSimplePostRequest(); - if (isChained() && !expectBadGatewayForEverything()) { - assertMethodSeenInRequestFilters(HttpMethod.POST); - assertMethodSeenInResponseFilters(HttpMethod.POST); - assertResponseFromFiltersMatchesActualResponse(); - } + } + + @Override + public void testSimpleGetRequestOverHTTPS() { + super.testSimpleGetRequestOverHTTPS(); + if (isChained() && !expectBadGatewayForEverything()) { + assertMethodSeenInRequestFilters(HttpMethod.CONNECT); + assertMethodSeenInRequestFilters(HttpMethod.GET); + assertMethodSeenInResponseFilters(HttpMethod.GET); + assertResponseFromFiltersMatchesActualResponse(); } - - @Override - public void testSimplePostRequestOverHTTPS() throws Exception { - super.testSimplePostRequestOverHTTPS(); - if (isChained() && !expectBadGatewayForEverything()) { - assertMethodSeenInRequestFilters(HttpMethod.CONNECT); - assertMethodSeenInRequestFilters(HttpMethod.POST); - assertMethodSeenInResponseFilters(HttpMethod.POST); - assertResponseFromFiltersMatchesActualResponse(); - } + } + + @Override + public void testSimplePostRequest() { + super.testSimplePostRequest(); + if (isChained() && !expectBadGatewayForEverything()) { + assertMethodSeenInRequestFilters(HttpMethod.POST); + assertMethodSeenInResponseFilters(HttpMethod.POST); + assertResponseFromFiltersMatchesActualResponse(); } - - private void assertMethodSeenInRequestFilters(HttpMethod method) { - assertThat(method - + " should have been seen in clientToProxyRequest filter", - requestPreMethodsSeen, hasItem(method)); - assertThat(method - + " should have been seen in proxyToServerRequest filter", - requestPostMethodsSeen, hasItem(method)); - } - - private void assertMethodSeenInResponseFilters(HttpMethod method) { - assertThat( - method - + " should have been seen as the original requests's method in serverToProxyResponse filter", - responsePreOriginalRequestMethodsSeen, hasItem(method)); - assertThat( - method - + " should have been seen as the original requests's method in proxyToClientResponse filter", - responsePostOriginalRequestMethodsSeen, hasItem(method)); - } - - private void assertResponseFromFiltersMatchesActualResponse() { - assertEquals( - "Data received through HttpFilters.serverToProxyResponse should match response", - lastResponse, responsePreBody.toString()); - assertEquals( - "Data received through HttpFilters.proxyToClientResponse should match response", - lastResponse, responsePostBody.toString()); - } - - @Override - protected void tearDown() { - this.upstreamProxy.abort(); + } + + @Override + public void testSimplePostRequestOverHTTPS() { + super.testSimplePostRequestOverHTTPS(); + if (isChained() && !expectBadGatewayForEverything()) { + assertMethodSeenInRequestFilters(HttpMethod.CONNECT); + assertMethodSeenInRequestFilters(HttpMethod.POST); + assertMethodSeenInResponseFilters(HttpMethod.POST); + assertResponseFromFiltersMatchesActualResponse(); } + } + + private void assertMethodSeenInRequestFilters(HttpMethod method) { + assertThat(requestPreMethodsSeen) + .as("%s should have been seen in clientToProxyRequest filter", method) + .contains(method); + assertThat(requestPostMethodsSeen) + .as("%s should have been seen in proxyToServerRequest filter", method) + .contains(method); + } + + private void assertMethodSeenInResponseFilters(HttpMethod method) { + assertThat(responsePreOriginalRequestMethodsSeen) + .as( + method + + " should have been seen as the original request's method in serverToProxyResponse filter") + .contains(method); + assertThat(responsePostOriginalRequestMethodsSeen) + .as( + method + + " should have been seen as the original request's method in proxyToClientResponse filter") + .contains(method); + } + + private void assertResponseFromFiltersMatchesActualResponse() { + assertThat(responsePreBody.toString()) + .as("Data received through HttpFilters.serverToProxyResponse should match response") + .isEqualTo(lastResponse); + assertThat(responsePostBody.toString()) + .as("Data received through HttpFilters.proxyToClientResponse should match response") + .isEqualTo(lastResponse); + } + + @Override + protected final void tearDown() { + upstreamProxy.abort(); + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmWithClientAuthenticationNotRequiredTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithClientAuthenticationNotRequiredTCPChainedProxyTest.java index d6c1ace6..f7d60d47 100644 --- a/src/test/java/org/littleshoot/proxy/MitmWithClientAuthenticationNotRequiredTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MitmWithClientAuthenticationNotRequiredTCPChainedProxyTest.java @@ -1,48 +1,38 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; - -import static org.littleshoot.proxy.TransportProtocol.TCP; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; /** - * Tests that when client authentication is not required, it doesn't matter what - * certs the client sends. + * Tests that when client authentication is not required, it doesn't matter what certs the client + * sends. */ -public class MitmWithClientAuthenticationNotRequiredTCPChainedProxyTest extends - MitmWithChainedProxyTest { - private final SslEngineSource serverSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - private final SslEngineSource clientSslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks", false, false); - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(serverSslEngineSource) - .withAuthenticateSslClients(false); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return clientSslEngineSource.newSslEngine(); - } - }; - } +public final class MitmWithClientAuthenticationNotRequiredTCPChainedProxyTest + extends MitmWithChainedProxyTest { + private final SslEngineSource serverSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks"); + private final SslEngineSource clientSslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks", false, false); + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy() + .withSslEngineSource(serverSslEngineSource) + .withAuthenticateSslClients(false); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { + return true; + } + + @Override + public SSLEngine newSslEngine() { + return clientSslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmWithEncryptedTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithEncryptedTCPChainedProxyTest.java index 5251b477..43af03f3 100644 --- a/src/test/java/org/littleshoot/proxy/MitmWithEncryptedTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MitmWithEncryptedTCPChainedProxyTest.java @@ -1,39 +1,32 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; +import org.junit.jupiter.api.parallel.Execution; +import org.junit.jupiter.api.parallel.ExecutionMode; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -public class MitmWithEncryptedTCPChainedProxyTest extends MitmWithChainedProxyTest { - private final SslEngineSource sslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP) - .withSslEngineSource(sslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return sslEngineSource.newSslEngine(); - } - }; - } +@Execution(ExecutionMode.SAME_THREAD) +public final class MitmWithEncryptedTCPChainedProxyTest extends MitmWithChainedProxyTest { + private final SslEngineSource sslEngineSource = + new SelfSignedSslEngineSource("target/chain_proxy_keystore_1.jks"); + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(sslEngineSource); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { + return true; + } + + @Override + public SSLEngine newSslEngine() { + return sslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmWithEncryptedUDTChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithEncryptedUDTChainedProxyTest.java deleted file mode 100644 index b259a7f9..00000000 --- a/src/test/java/org/littleshoot/proxy/MitmWithEncryptedUDTChainedProxyTest.java +++ /dev/null @@ -1,39 +0,0 @@ -package org.littleshoot.proxy; - -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - -import javax.net.ssl.SSLEngine; - -import static org.littleshoot.proxy.TransportProtocol.UDT; - -public class MitmWithEncryptedUDTChainedProxyTest extends MitmWithChainedProxyTest { - private final SslEngineSource sslEngineSource = new SelfSignedSslEngineSource( - "chain_proxy_keystore_1.jks"); - - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(UDT) - .withSslEngineSource(sslEngineSource); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.UDT; - } - - @Override - public boolean requiresEncryption() { - return true; - } - - @Override - public SSLEngine newSslEngine() { - return sslEngineSource.newSslEngine(); - } - }; - } -} diff --git a/src/test/java/org/littleshoot/proxy/MitmWithPerRequestPoolTest.java b/src/test/java/org/littleshoot/proxy/MitmWithPerRequestPoolTest.java new file mode 100644 index 00000000..0fbda8e8 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/MitmWithPerRequestPoolTest.java @@ -0,0 +1,63 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.extras.TestMitmManager; + +/** + * Integration tests for per-request MITM pooling. When poolPerRequestInMitm is enabled, each HTTP + * request through the MITM tunnel independently acquires and releases a server connection from the + * shared pool, rather than using a dedicated per-session connection. + */ +@Tag("slow-test") +public class MitmWithPerRequestPoolTest extends AbstractProxyTest { + + @Override + protected void setUp() { + proxyServer = + bootstrapProxy() + .withPort(0) + .withManInTheMiddle(new TestMitmManager()) + .withSharedServerConnectionPool(true) + .withPoolSharedMitmConnections(true) + .withPoolPerRequestInMitm(true) + .start(); + } + + @Override + protected boolean isMITM() { + return true; + } + + @Test + void testGetRequestOverHTTPS() { + ResponseInfo response = httpGetWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true, false); + assertThat(response.getStatusCode()).isEqualTo(200); + } + + @Test + void testPostRequestOverHTTPS() { + ResponseInfo response = httpPostWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true); + assertThat(response.getStatusCode()).isEqualTo(200); + } + + @Test + void testMultipleRequestsOverHTTPS() { + for (int i = 0; i < 5; i++) { + ResponseInfo response = httpGetWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true, false); + assertThat(response.getStatusCode()).as("Request %d should succeed", i).isEqualTo(200); + } + } + + @Test + void testMultipleRequestsOverHTTPSFromDifferentClients() { + for (int i = 0; i < 3; i++) { + ResponseInfo response = httpGetWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true, false); + assertThat(response.getStatusCode()) + .as("Request from client %d should succeed", i) + .isEqualTo(200); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/MitmWithSharedPoolTest.java b/src/test/java/org/littleshoot/proxy/MitmWithSharedPoolTest.java new file mode 100644 index 00000000..122dab73 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/MitmWithSharedPoolTest.java @@ -0,0 +1,69 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.extras.TestMitmManager; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.impl.PoolMetrics; +import org.littleshoot.proxy.impl.ServerConnectionPool; + +/** + * Integration tests for MITM with shared server connection pool. Tests that upstream TLS + * connections are pooled and reused across MITM sessions when poolSharedMitmConnections is enabled. + */ +public class MitmWithSharedPoolTest extends AbstractProxyTest { + + @Override + protected void setUp() { + proxyServer = + bootstrapProxy() + .withPort(0) + .withManInTheMiddle(new TestMitmManager()) + .withSharedServerConnectionPool(true) + .withPoolSharedMitmConnections(true) + .start(); + } + + @Override + protected boolean isMITM() { + return true; + } + + @Test + void testMitmGetRequestOverHTTPS() { + ResponseInfo response = httpGetWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true, false); + assertThat(response.getStatusCode()).isEqualTo(200); + } + + @Test + void testMitmPostRequestOverHTTPS() { + ResponseInfo response = httpPostWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true); + assertThat(response.getStatusCode()).isEqualTo(200); + } + + @Test + void testMitmGetRequestOverHTTPSFromDifferentClients() { + for (int i = 0; i < 3; i++) { + ResponseInfo response = httpGetWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true, false); + assertThat(response.getStatusCode()) + .as("Request from client %d should succeed", i) + .isEqualTo(200); + } + } + + @Test + void testSharedPoolMetricsShowConnectionReuse() { + DefaultHttpProxyServer impl = (DefaultHttpProxyServer) proxyServer; + ServerConnectionPool pool = impl.getServerConnectionPool(); + + httpGetWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true, false); + PoolMetrics afterFirst = pool.getMetrics(); + assertThat(afterFirst.getTotalConnections()).isGreaterThanOrEqualTo(1); + assertThat(afterFirst.getBorrowCount()).isGreaterThanOrEqualTo(1); + + httpGetWithApacheClient(httpsWebHost, DEFAULT_RESOURCE, true, false); + PoolMetrics afterSecond = pool.getMetrics(); + assertThat(afterSecond.getBorrowCount()).isGreaterThan(afterFirst.getBorrowCount()); + } +} diff --git a/src/test/java/org/littleshoot/proxy/MitmWithUnencryptedTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithUnencryptedTCPChainedProxyTest.java index c2c92d81..0e468ed3 100644 --- a/src/test/java/org/littleshoot/proxy/MitmWithUnencryptedTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/MitmWithUnencryptedTCPChainedProxyTest.java @@ -1,26 +1,8 @@ package org.littleshoot.proxy; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -public class MitmWithUnencryptedTCPChainedProxyTest extends MitmWithChainedProxyTest { - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return false; - } - }; - } +public final class MitmWithUnencryptedTCPChainedProxyTest extends MitmWithChainedProxyTest { + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy(); + } } diff --git a/src/test/java/org/littleshoot/proxy/MitmWithUnencryptedUDTChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/MitmWithUnencryptedUDTChainedProxyTest.java deleted file mode 100644 index 18fb0b34..00000000 --- a/src/test/java/org/littleshoot/proxy/MitmWithUnencryptedUDTChainedProxyTest.java +++ /dev/null @@ -1,26 +0,0 @@ -package org.littleshoot.proxy; - -import static org.littleshoot.proxy.TransportProtocol.UDT; - -public class MitmWithUnencryptedUDTChainedProxyTest extends MitmWithChainedProxyTest { - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(UDT); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.UDT; - } - - @Override - public boolean requiresEncryption() { - return false; - } - }; - } -} diff --git a/src/test/java/org/littleshoot/proxy/NoChainedProxiesTest.java b/src/test/java/org/littleshoot/proxy/NoChainedProxiesTest.java index 11cf3ce1..5919f697 100644 --- a/src/test/java/org/littleshoot/proxy/NoChainedProxiesTest.java +++ b/src/test/java/org/littleshoot/proxy/NoChainedProxiesTest.java @@ -1,26 +1,25 @@ package org.littleshoot.proxy; -import org.junit.Test; +import org.junit.jupiter.api.Test; -/** - * Tests that when there are no chained proxies, we get a bad gateway. - */ -public class NoChainedProxiesTest extends AbstractProxyTest { - @Override - protected void setUp() { - this.proxyServer = bootstrapProxy() - .withPort(0) - .withChainProxyManager((httpRequest, chainedProxies, clientDetails) -> { - // Leave list empty +/** Tests that when there are no chained proxies, we get a bad gateway. */ +public final class NoChainedProxiesTest extends AbstractProxyTest { + @Override + protected void setUp() { + proxyServer = + bootstrapProxy() + .withPort(0) + .withChainProxyManager( + (httpRequest, chainedProxies, clientDetails) -> { + // Leave list empty }) - .withIdleConnectionTimeout(1) - .start(); - } + .withIdleConnectionTimeout(1) + .start(); + } - @Test - public void testNoChainedProxy() throws Exception { - ResponseInfo response = httpGetWithApacheClient(webHost, - DEFAULT_RESOURCE, true, false); - assertReceivedBadGateway(response); - } + @Test + public void testNoChainedProxy() { + ResponseInfo response = httpGetWithApacheClient(webHost, DEFAULT_RESOURCE, true, false); + assertReceivedBadGateway(response); + } } diff --git a/src/test/java/org/littleshoot/proxy/PerformanceServer.java b/src/test/java/org/littleshoot/proxy/PerformanceServer.java index 44add8db..f1260131 100644 --- a/src/test/java/org/littleshoot/proxy/PerformanceServer.java +++ b/src/test/java/org/littleshoot/proxy/PerformanceServer.java @@ -7,35 +7,31 @@ import org.eclipse.jetty.server.handler.HandlerList; import org.eclipse.jetty.server.handler.ResourceHandler; -/** - * This launches a Jetty-based HTTP server that serves static files from the - * perfsite folder. - */ +/** This launches a Jetty-based HTTP server that serves static files from the perfsite folder. */ public class PerformanceServer { - public void run(int port) throws Exception { - Server server = new Server(); - ServerConnector connector = new ServerConnector(server); - connector.setPort(port); - server.addConnector(connector); + public void run(int port) throws Exception { + Server server = new Server(); + ServerConnector connector = new ServerConnector(server); + connector.setPort(port); + server.addConnector(connector); - ResourceHandler resource_handler = new ResourceHandler(); - resource_handler.setDirectoriesListed(true); - resource_handler.setWelcomeFiles(new String[] { "index.html" }); + ResourceHandler resource_handler = new ResourceHandler(); + resource_handler.setDirectoriesListed(true); + resource_handler.setWelcomeFiles(new String[] {"index.html"}); - resource_handler.setResourceBase("./performance/site/"); + resource_handler.setResourceBase("./performance/site/"); - HandlerList handlers = new HandlerList(); - handlers.setHandlers(new Handler[] { resource_handler, - new DefaultHandler() }); - server.setHandler(handlers); + HandlerList handlers = new HandlerList(); + handlers.setHandlers(new Handler[] {resource_handler, new DefaultHandler()}); + server.setHandler(handlers); - server.start(); - System.out.println("Started performance file server at port: " + port); - server.join(); - } + server.start(); + System.out.println("Started performance file server at port: " + port); + server.join(); + } - public static void main(String[] args) throws Exception { - int port = args.length > 0 ? Integer.parseInt(args[0]) : 9000; - new PerformanceServer().run(port); - } + public static void main(String[] args) throws Exception { + int port = args.length > 0 ? Integer.parseInt(args[0]) : 9000; + new PerformanceServer().run(port); + } } diff --git a/src/test/java/org/littleshoot/proxy/ProxyHeadersTest.java b/src/test/java/org/littleshoot/proxy/ProxyHeadersTest.java index c25525ff..d18cf0ed 100644 --- a/src/test/java/org/littleshoot/proxy/ProxyHeadersTest.java +++ b/src/test/java/org/littleshoot/proxy/ProxyHeadersTest.java @@ -1,76 +1,99 @@ package org.littleshoot.proxy; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.createProxiedHttpClient; + +import com.github.tomakehurst.wiremock.WireMockServer; import org.apache.http.Header; import org.apache.http.HttpResponse; -import org.apache.http.client.HttpClient; import org.apache.http.client.methods.HttpGet; +import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.util.EntityUtils; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; -import org.mockserver.model.ConnectionOptions; -import static org.hamcrest.Matchers.emptyArray; -import static org.junit.Assert.assertThat; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; +/** Tests the proxy's handling and manipulation of headers. */ +public final class ProxyHeadersTest { + private HttpProxyServer proxyServer; -/** - * Tests the proxy's handling and manipulation of headers. - */ -public class ProxyHeadersTest { - private HttpProxyServer proxyServer; + private WireMockServer mockServer; + private int mockServerPort; - private ClientAndServer mockServer; - private int mockServerPort; + @BeforeEach + void setUp() { + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + } - @Before - public void setUp() { - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); + @AfterEach + void tearDown() { + try { + if (proxyServer != null) { + proxyServer.abort(); + } + } finally { + if (mockServer != null) { + mockServer.stop(); + } } + } - @After - public void tearDown() { - try { - if (proxyServer != null) { - proxyServer.abort(); - } - } finally { - if (mockServer != null) { - mockServer.stop(); - } - } - } + @Test + public void testProxyRemovesConnectionHeadersFromServer() throws Exception { + // the proxy should remove all Connection headers, since all values in the + // Connection header are hop-by-hop headers. + mockServer.stubFor( + get(urlEqualTo("/connectionheaders")) + .willReturn( + aResponse() + .withStatus(200) + .withBody("success") + .withHeader("Connection", "Dummy-Header") + .withHeader("Dummy-Header", "dummy-value"))); + + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - @Test - public void testProxyRemovesConnectionHeadersFromServer() throws Exception { - // the proxy should remove all Connection headers, since all values in the Connection header are hop-by-hop headers. - mockServer.when(request() - .withMethod("GET") - .withPath("/connectionheaders"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("success") - .withHeader("Connection", "Dummy-Header") - .withHeader("Dummy-Header", "dummy-value") - .withConnectionOptions(new ConnectionOptions() - .withSuppressConnectionHeader(true)) - ); + try (CloseableHttpClient httpClient = + createProxiedHttpClient(proxyServer.getListenAddress().getPort())) { + HttpResponse response = + httpClient.execute( + new HttpGet("http://localhost:" + mockServerPort + "/connectionheaders")); + EntityUtils.consume(response.getEntity()); - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); + Header[] dummyHeaders = response.getHeaders("Dummy-Header"); + assertThat(dummyHeaders) + .as("Expected proxy to remove the Dummy-Header specified in the Connection header") + .isEmpty(); + } + } + + @Test + public void testProxyRemovesHopByHopHeadersFromClient() throws Exception { + mockServer.stubFor( + get(urlEqualTo("/connectionheaders")) + .willReturn(aResponse().withStatus(200).withBody("success"))); - HttpClient httpClient = TestUtils.createProxiedHttpClient(proxyServer.getListenAddress().getPort()); - HttpResponse response = httpClient.execute(new HttpGet("http://localhost:" + mockServerPort + "/connectionheaders")); - EntityUtils.consume(response.getEntity()); + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - Header[] dummyHeaders = response.getHeaders("Dummy-Header"); - assertThat("Expected proxy to remove the Dummy-Header specified in the Connection header", dummyHeaders, emptyArray()); + try (CloseableHttpClient httpClient = + createProxiedHttpClient(proxyServer.getListenAddress().getPort())) { + HttpGet clientRequest = + new HttpGet("http://localhost:" + mockServerPort + "/connectionheaders"); + clientRequest.addHeader("Proxy-Authenticate", ""); + clientRequest.addHeader("Proxy-Authorization", ""); + HttpResponse response = httpClient.execute(clientRequest); + EntityUtils.consume(response.getEntity()); + assertThat(response.getStatusLine().getStatusCode()).isEqualTo(200); } + + mockServer.verify( + getRequestedFor(urlEqualTo("/connectionheaders")) + .withoutHeader("Proxy-Authenticate") + .withoutHeader("Proxy-Authorization")); + } } diff --git a/src/test/java/org/littleshoot/proxy/ResponseInfo.java b/src/test/java/org/littleshoot/proxy/ResponseInfo.java index 8ad0df24..efe6c330 100644 --- a/src/test/java/org/littleshoot/proxy/ResponseInfo.java +++ b/src/test/java/org/littleshoot/proxy/ResponseInfo.java @@ -1,53 +1,46 @@ package org.littleshoot.proxy; public class ResponseInfo { - private int statusCode; - private String body; - - public ResponseInfo(int statusCode, String body) { - super(); - this.statusCode = statusCode; - this.body = body; - } - - public int getStatusCode() { - return statusCode; - } - - public String getBody() { - return body; - } - - @Override - public int hashCode() { - final int prime = 31; - int result = 1; - result = prime * result + ((body == null) ? 0 : body.hashCode()); - result = prime * result + statusCode; - return result; - } - - @Override - public boolean equals(Object obj) { - if (this == obj) - return true; - if (obj == null) - return false; - if (getClass() != obj.getClass()) - return false; - ResponseInfo other = (ResponseInfo) obj; - if (body == null) { - if (other.body != null) - return false; - } else if (!body.equals(other.body)) - return false; - return statusCode == other.statusCode; - } - - @Override - public String toString() { - return "ResponseInfo [statusCode=" + statusCode + ", body=" + body - + "]"; - } - + private final int statusCode; + private final String body; + + public ResponseInfo(int statusCode, String body) { + super(); + this.statusCode = statusCode; + this.body = body; + } + + public int getStatusCode() { + return statusCode; + } + + public String getBody() { + return body; + } + + @Override + public int hashCode() { + final int prime = 31; + int result = 1; + result = prime * result + ((body == null) ? 0 : body.hashCode()); + result = prime * result + statusCode; + return result; + } + + @Override + public boolean equals(Object obj) { + if (this == obj) return true; + if (obj == null) return false; + if (getClass() != obj.getClass()) return false; + ResponseInfo other = (ResponseInfo) obj; + if (body == null) { + if (other.body != null) return false; + } else if (!body.equals(other.body)) return false; + return statusCode == other.statusCode; + } + + @Override + public String toString() { + return "ResponseInfo [statusCode=" + statusCode + ", body=" + body + "]"; + } } diff --git a/src/test/java/org/littleshoot/proxy/SelfSignedGeneratedSslEngineChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/SelfSignedGeneratedSslEngineChainedProxyTest.java new file mode 100644 index 00000000..f92f1fb7 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/SelfSignedGeneratedSslEngineChainedProxyTest.java @@ -0,0 +1,58 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.io.File; +import java.io.IOException; +import javax.net.ssl.SSLEngine; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; + +public final class SelfSignedGeneratedSslEngineChainedProxyTest extends BaseChainedProxyTest { + + @TempDir private File temporaryFolder; + + private SslEngineSource sslEngineSource; + + @Override + protected void setUp() throws IOException { + String keyStorePath = temporaryFolder.toPath().resolve("chain_proxy_keystore.jks").toString(); + sslEngineSource = + new SelfSignedSslEngineSource( + keyStorePath, false, true, "littleproxy", "Be Your Own Lantern"); + super.setUp(); + } + + @Test + public void testKeyStoreGeneratedAtProvidedPath() { + File keyStoreFile = temporaryFolder.toPath().resolve("chain_proxy_keystore.jks").toFile(); + assertThat(keyStoreFile.exists()).isTrue(); + } + + @Test + public void testCertExportedToKeyStoreDirectory() { + File certFile = temporaryFolder.toPath().resolve("littleproxy_cert").toFile(); + assertThat(certFile.exists()).isTrue(); + } + + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(sslEngineSource); + } + + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { + return true; + } + + @Override + public SSLEngine newSslEngine() { + return sslEngineSource.newSslEngine(); + } + }; + } +} diff --git a/src/test/java/org/littleshoot/proxy/SelfSignedSslEngineChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/SelfSignedSslEngineChainedProxyTest.java index b4be9a4b..ea4db0d9 100644 --- a/src/test/java/org/littleshoot/proxy/SelfSignedSslEngineChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/SelfSignedSslEngineChainedProxyTest.java @@ -1,31 +1,34 @@ package org.littleshoot.proxy; -import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; - import javax.net.ssl.SSLEngine; +import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -public class SelfSignedSslEngineChainedProxyTest extends BaseChainedProxyTest { - private final SslEngineSource sslEngineSource = new SelfSignedSslEngineSource("/certificate/chain_proxy_keystore.jks", - false, true, "littleproxy", "Be Your Own Lantern"); +public final class SelfSignedSslEngineChainedProxyTest extends BaseChainedProxyTest { + private final SslEngineSource sslEngineSource = + new SelfSignedSslEngineSource( + "/certificate/chain_proxy_keystore.jks", + false, + true, + "littleproxy", + "Be Your Own Lantern"); - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withSslEngineSource(sslEngineSource); - } + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy().withSslEngineSource(sslEngineSource); + } - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public boolean requiresEncryption() { - return true; - } + @Override + protected ChainedProxy newChainedProxy() { + return new BaseChainedProxy() { + @Override + public boolean requiresEncryption() { + return true; + } - @Override - public SSLEngine newSslEngine() { - return sslEngineSource.newSslEngine(); - } - }; - } + @Override + public SSLEngine newSslEngine() { + return sslEngineSource.newSslEngine(); + } + }; + } } diff --git a/src/test/java/org/littleshoot/proxy/ServerConnectionPoolTypeTest.java b/src/test/java/org/littleshoot/proxy/ServerConnectionPoolTypeTest.java new file mode 100644 index 00000000..ef38cdd3 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/ServerConnectionPoolTypeTest.java @@ -0,0 +1,93 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.impl.ConcurrentMapServerConnectionPool; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.impl.ServerConnectionPool; + +class ServerConnectionPoolTypeTest { + + @Test + void shouldCreateConcurrentMapPool() { + DefaultHttpProxyServer server = startServer(ServerConnectionPoolType.CONCURRENT_MAP, 3, 7); + try { + ServerConnectionPool pool = server.getServerConnectionPool(); + assertThat(pool).isInstanceOf(ConcurrentMapServerConnectionPool.class); + assertThat(pool.getMaxConnectionsPerHost()).isEqualTo(3); + assertThat(pool.getMaxConnections()).isEqualTo(7); + } finally { + server.abort(); + } + } + + @Test + void shouldReturnNullWhenPoolDisabled() { + DefaultHttpProxyServer server = + (DefaultHttpProxyServer) + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withSharedServerConnectionPool(false) + .start(); + try { + assertThat(server.getServerConnectionPool()).isNull(); + } finally { + server.abort(); + } + } + + @Test + void shouldReturnSameInstanceOnRepeatedCalls() { + DefaultHttpProxyServer server = startServer(ServerConnectionPoolType.CONCURRENT_MAP, 3, 7); + try { + ServerConnectionPool first = server.getServerConnectionPool(); + ServerConnectionPool second = server.getServerConnectionPool(); + assertThat(second).isSameAs(first); + } finally { + server.abort(); + } + } + + @Test + void shouldReturnSameInstanceUnderConcurrentAccess() throws Exception { + DefaultHttpProxyServer server = startServer(ServerConnectionPoolType.CONCURRENT_MAP, 3, 7); + try { + int threadCount = 10; + CountDownLatch latch = new CountDownLatch(threadCount); + AtomicReference[] refs = new AtomicReference[threadCount]; + for (int i = 0; i < threadCount; i++) { + refs[i] = new AtomicReference<>(); + int idx = i; + new Thread( + () -> { + refs[idx].set(server.getServerConnectionPool()); + latch.countDown(); + }) + .start(); + } + assertThat(latch.await(10, TimeUnit.SECONDS)).isTrue(); + ServerConnectionPool first = refs[0].get(); + for (int i = 1; i < threadCount; i++) { + assertThat(refs[i].get()).isSameAs(first); + } + } finally { + server.abort(); + } + } + + private static DefaultHttpProxyServer startServer( + ServerConnectionPoolType poolType, int maxConnectionsPerHost, int maxConnections) { + return (DefaultHttpProxyServer) + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withSharedServerConnectionPool(true) + .withServerConnectionPoolType(poolType) + .withMaxConnectionsPerHost(maxConnectionsPerHost) + .withMaxConnections(maxConnections) + .start(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/ServerGroupTest.java b/src/test/java/org/littleshoot/proxy/ServerGroupTest.java index 5e4d8eac..7fd6134e 100644 --- a/src/test/java/org/littleshoot/proxy/ServerGroupTest.java +++ b/src/test/java/org/littleshoot/proxy/ServerGroupTest.java @@ -1,143 +1,156 @@ package org.littleshoot.proxy; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static java.util.concurrent.Executors.newFixedThreadPool; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.test.HttpClientUtil.performLocalHttpGet; + +import com.github.tomakehurst.wiremock.WireMockServer; import io.netty.handler.codec.http.HttpObject; import io.netty.handler.codec.http.HttpRequest; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Future; +import java.util.concurrent.atomic.AtomicReference; import org.apache.http.HttpResponse; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; +import org.jspecify.annotations.NonNull; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; import org.littleshoot.proxy.impl.ThreadPoolConfiguration; -import org.littleshoot.proxy.test.HttpClientUtil; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; - -import java.util.concurrent.*; -import java.util.concurrent.atomic.AtomicReference; - -import static org.junit.Assert.assertEquals; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; - -public class ServerGroupTest { - private ClientAndServer mockServer; - private int mockServerPort; - - private HttpProxyServer proxyServer; - @Before - public void setUp() { - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); +public final class ServerGroupTest { + private WireMockServer mockServer; + private int mockServerPort; + + private HttpProxyServer proxyServer; + + @BeforeEach + void setUp() { + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + } + + @AfterEach + void tearDown() { + try { + if (mockServer != null) { + mockServer.stop(); + } + } finally { + if (proxyServer != null) { + proxyServer.abort(); + } } - - @After - public void tearDown() { - try { - if (mockServer != null) { - mockServer.stop(); - } - } finally { - if (proxyServer != null) { - proxyServer.abort(); - } - } - } - - @Test - public void testSingleWorkerThreadPoolConfiguration() throws ExecutionException, InterruptedException { - final String firstRequestPath = "/testSingleThreadFirstRequest"; - final String secondRequestPath = "/testSingleThreadSecondRequest"; - - // set up two server responses that will execute more or less simultaneously. the first request has a small - // delay, to reduce the chance that the first request will finish entirely before the second request is finished - // (and thus be somewhat more likely to be serviced by the same thread, even if the ThreadPoolConfiguration is - // not behaving properly). - mockServer.when(request() - .withMethod("GET") - .withPath(firstRequestPath), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("first") - .withDelay(TimeUnit.MILLISECONDS, 500) - ); - - mockServer.when(request() - .withMethod("GET") - .withPath(secondRequestPath), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withBody("second") - ); - - // save the names of the threads that execute the filter methods. filter methods are executed by the worker thread - // handling the request/response, so if there is only one worker thread, the filter methods should be executed - // by the same thread. - final AtomicReference firstClientThreadName = new AtomicReference<>(); - final AtomicReference secondClientThreadName = new AtomicReference<>(); - - final AtomicReference firstProxyThreadName = new AtomicReference<>(); - final AtomicReference secondProxyThreadName = new AtomicReference<>(); - - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withFiltersSource(new HttpFiltersSourceAdapter() { - @Override - public HttpFilters filterRequest(HttpRequest originalRequest) { - return new HttpFiltersAdapter(originalRequest) { - @Override - public io.netty.handler.codec.http.HttpResponse clientToProxyRequest(HttpObject httpObject) { - if (originalRequest.uri().endsWith(firstRequestPath)) { - firstClientThreadName.set(Thread.currentThread().getName()); - } else if (originalRequest.uri().endsWith(secondRequestPath)) { - secondClientThreadName.set(Thread.currentThread().getName()); - } - - return super.clientToProxyRequest(httpObject); - } - - @Override - public void serverToProxyResponseReceived() { - if (originalRequest.uri().endsWith(firstRequestPath)) { - firstProxyThreadName.set(Thread.currentThread().getName()); - } else if (originalRequest.uri().endsWith(secondRequestPath)) { - secondProxyThreadName.set(Thread.currentThread().getName()); - } - } - }; - } + } + + @Test + public void testSingleWorkerThreadPoolConfiguration() + throws ExecutionException, InterruptedException { + final String firstRequestPath = "/testSingleThreadFirstRequest"; + final String secondRequestPath = "/testSingleThreadSecondRequest"; + + // set up two server responses that will execute more or less simultaneously. + // the first request has a small + // delay, to reduce the chance that the first request will finish entirely + // before the second request is finished + // (and thus be somewhat more likely to be serviced by the same thread, even if + // the ThreadPoolConfiguration is + // not behaving properly). + mockServer.stubFor( + get(urlEqualTo(firstRequestPath)) + .willReturn(aResponse().withStatus(200).withBody("first").withFixedDelay(500))); + + mockServer.stubFor( + get(urlEqualTo(secondRequestPath)) + .willReturn(aResponse().withStatus(200).withBody("second"))); + + // save the names of the threads that execute the filter methods. filter methods + // are executed by the worker thread + // handling the request/response, so if there is only one worker thread, the + // filter methods should be executed + // by the same thread. + final AtomicReference firstClientThreadName = new AtomicReference<>(); + final AtomicReference secondClientThreadName = new AtomicReference<>(); + + final AtomicReference firstProxyThreadName = new AtomicReference<>(); + final AtomicReference secondProxyThreadName = new AtomicReference<>(); + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withFiltersSource( + new HttpFiltersSourceAdapter() { + @NonNull + @Override + public HttpFilters filterRequest(@NonNull HttpRequest originalRequest) { + return new HttpFiltersAdapter(originalRequest) { + @Override + public io.netty.handler.codec.http.HttpResponse clientToProxyRequest( + @NonNull HttpObject httpObject) { + if (originalRequest.uri().endsWith(firstRequestPath)) { + firstClientThreadName.set(Thread.currentThread().getName()); + } else if (originalRequest.uri().endsWith(secondRequestPath)) { + secondClientThreadName.set(Thread.currentThread().getName()); + } + + return super.clientToProxyRequest(httpObject); + } + + @Override + public void serverToProxyResponseReceived() { + if (originalRequest.uri().endsWith(firstRequestPath)) { + firstProxyThreadName.set(Thread.currentThread().getName()); + } else if (originalRequest.uri().endsWith(secondRequestPath)) { + secondProxyThreadName.set(Thread.currentThread().getName()); + } + } + }; + } }) - .withThreadPoolConfiguration(new ThreadPoolConfiguration() - .withAcceptorThreads(1) - .withClientToProxyWorkerThreads(1) - .withProxyToServerWorkerThreads(1)) - .start(); - - // execute both requests in parallel, to increase the chance of blocking due to the single-threaded ThreadPoolConfiguration - - Runnable firstRequest = () -> { - HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + mockServerPort + firstRequestPath, proxyServer); - assertEquals(200, response.getStatusLine().getStatusCode()); + .withThreadPoolConfiguration( + new ThreadPoolConfiguration() + .withAcceptorThreads(1) + .withClientToProxyWorkerThreads(1) + .withProxyToServerWorkerThreads(1)) + .start(); + + // execute both requests in parallel, to increase the chance of blocking due to + // the single-threaded ThreadPoolConfiguration + + Runnable firstRequest = + () -> { + HttpResponse response = + performLocalHttpGet(mockServerPort, firstRequestPath, proxyServer); + assertThat(response.getStatusLine().getStatusCode()).isEqualTo(200); }; - Runnable secondRequest = () -> { - HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + mockServerPort + secondRequestPath, proxyServer); - assertEquals(200, response.getStatusLine().getStatusCode()); + Runnable secondRequest = + () -> { + HttpResponse response = + performLocalHttpGet(mockServerPort, secondRequestPath, proxyServer); + assertThat(response.getStatusLine().getStatusCode()).isEqualTo(200); }; - ExecutorService executor = Executors.newFixedThreadPool(2); - Future firstFuture = executor.submit(firstRequest); - Future secondFuture = executor.submit(secondRequest); - - firstFuture.get(); - secondFuture.get(); + ExecutorService executor = newFixedThreadPool(2); + Future firstFuture = executor.submit(firstRequest); + Future secondFuture = executor.submit(secondRequest); - Thread.sleep(500); + firstFuture.get(); + secondFuture.get(); - assertEquals("Expected clientToProxy filter methods to be executed on the same thread for both requests", firstClientThreadName.get(), secondClientThreadName.get()); - assertEquals("Expected serverToProxy filter methods to be executed on the same thread for both requests", firstProxyThreadName.get(), secondProxyThreadName.get()); - } + Thread.sleep(500); + assertThat(secondClientThreadName.get()) + .as( + "Expected clientToProxy filter methods to be executed on the same thread for both requests") + .isEqualTo(firstClientThreadName.get()); + assertThat(secondProxyThreadName.get()) + .as( + "Expected serverToProxy filter methods to be executed on the same thread for both requests") + .isEqualTo(firstProxyThreadName.get()); + } } diff --git a/src/test/java/org/littleshoot/proxy/SharedConnectionPoolTest.java b/src/test/java/org/littleshoot/proxy/SharedConnectionPoolTest.java new file mode 100644 index 00000000..18373aa7 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/SharedConnectionPoolTest.java @@ -0,0 +1,93 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.impl.ConcurrentMapServerConnectionPool; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.impl.PoolMetrics; +import org.littleshoot.proxy.impl.ServerConnectionPool; + +/** + * Integration tests for the shared ServerConnectionPool feature. Tests that connection pooling + * works correctly when enabled. + */ +public class SharedConnectionPoolTest extends BaseProxyTest { + + @Override + protected void setUp() { + // Enable the shared server connection pool + proxyServer = bootstrapProxy().withSharedServerConnectionPool(true).withPort(0).start(); + } + + @Test + void testSimpleGetRequestWithPoolEnabled() { + // Basic test - verify the proxy works with pool enabled + ResponseInfo response = httpGetWithApacheClient(webHost, DEFAULT_RESOURCE, true, false); + assertThat(response.getStatusCode()).isEqualTo(200); + } + + @Test + void testSimplePostRequestWithPoolEnabled() { + // Test POST request works with pool enabled + ResponseInfo response = httpPostWithApacheClient(webHost, DEFAULT_RESOURCE, true); + assertThat(response.getStatusCode()).isEqualTo(200); + } + + @Test + void testMultipleRequestsToSameServerWithPoolEnabled() { + // Make multiple requests to the same server - with pool enabled, + // these should reuse connections + ServerConnectionPool pool = ((DefaultHttpProxyServer) proxyServer).getServerConnectionPool(); + PoolMetrics before = pool.getMetrics(); + for (int i = 0; i < 5; i++) { + ResponseInfo response = httpGetWithApacheClient(webHost, DEFAULT_RESOURCE, true, false); + assertThat(response.getStatusCode()).as("Request %d should succeed", i).isEqualTo(200); + } + PoolMetrics after = pool.getMetrics(); + assertThat(after.getTotalConnections() - before.getTotalConnections()) + .as("Should reuse a single connection across all 5 requests") + .isEqualTo(1); + } + + @Test + void testProxyWorksWithPoolEnabled() { + // Verify the proxy server is configured correctly + assertThat(proxyServer).isNotNull(); + // The pool should be created when enabled + assertThat(((DefaultHttpProxyServer) proxyServer).getServerConnectionPool()).isNotNull(); + } + + @Test + void testDefaultPoolTypeIsConcurrentMap() { + ServerConnectionPool pool = ((DefaultHttpProxyServer) proxyServer).getServerConnectionPool(); + assertThat(pool).isInstanceOf(ConcurrentMapServerConnectionPool.class); + } + + @Test + void testKeepAliveWithPoolEnabled() { + // Test that keep-alive works with the pool enabled + // This is important because the pool relies on connection reuse + ServerConnectionPool pool = ((DefaultHttpProxyServer) proxyServer).getServerConnectionPool(); + PoolMetrics before = pool.getMetrics(); + ResponseInfo response1 = httpGetWithApacheClient(webHost, DEFAULT_RESOURCE, true, false); + assertThat(response1.getStatusCode()).isEqualTo(200); + + // Make another request on the same connection + ResponseInfo response2 = httpGetWithApacheClient(webHost, DEFAULT_RESOURCE, true, false); + assertThat(response2.getStatusCode()).isEqualTo(200); + PoolMetrics after = pool.getMetrics(); + assertThat(after.getTotalConnections() - before.getTotalConnections()) + .as("Should reuse a single connection across both requests") + .isEqualTo(1); + } + + @Test + void testMaxConnectionsPerHostSetting() { + // Verify the pool has the correct max connections per host setting + ServerConnectionPool pool = ((DefaultHttpProxyServer) proxyServer).getServerConnectionPool(); + assertThat(pool).isNotNull(); + assertThat(pool.getMaxConnectionsPerHost()).isEqualTo(10); + assertThat(pool.getMaxConnections()).isEqualTo(200); + } +} diff --git a/src/test/java/org/littleshoot/proxy/SimpleProxyTest.java b/src/test/java/org/littleshoot/proxy/SimpleProxyTest.java index 3d0251ec..bc6271a6 100644 --- a/src/test/java/org/littleshoot/proxy/SimpleProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/SimpleProxyTest.java @@ -1,13 +1,9 @@ package org.littleshoot.proxy; -/** - * Tests just a single basic proxy. - */ -public class SimpleProxyTest extends BaseProxyTest { - @Override - protected void setUp() { - this.proxyServer = bootstrapProxy() - .withPort(0) - .start(); - } +/** Tests just a single basic proxy. */ +public final class SimpleProxyTest extends BaseProxyTest { + @Override + protected void setUp() { + proxyServer = bootstrapProxy().withPort(0).start(); + } } diff --git a/src/test/java/org/littleshoot/proxy/Socks4ChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/Socks4ChainedProxyTest.java index c202a22e..d79a2dd2 100644 --- a/src/test/java/org/littleshoot/proxy/Socks4ChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/Socks4ChainedProxyTest.java @@ -1,8 +1,8 @@ package org.littleshoot.proxy; -public class Socks4ChainedProxyTest extends BaseChainedSocksProxyTest { - @Override - protected ChainedProxyType getSocksProxyType() { - return ChainedProxyType.SOCKS4; - } +public final class Socks4ChainedProxyTest extends BaseChainedSocksProxyTest { + @Override + protected ChainedProxyType getSocksProxyType() { + return ChainedProxyType.SOCKS4; + } } diff --git a/src/test/java/org/littleshoot/proxy/Socks4ChainedProxyWithMissConfiguredSendProxyProtocolTest.java b/src/test/java/org/littleshoot/proxy/Socks4ChainedProxyWithMissConfiguredSendProxyProtocolTest.java new file mode 100644 index 00000000..035a2f2c --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/Socks4ChainedProxyWithMissConfiguredSendProxyProtocolTest.java @@ -0,0 +1,25 @@ +package org.littleshoot.proxy; + +import org.junit.jupiter.api.Tag; + +@Tag("slow-test") +public final class Socks4ChainedProxyWithMissConfiguredSendProxyProtocolTest + extends BaseChainedSocksProxyTest { + @Override + protected ChainedProxyType getSocksProxyType() { + return ChainedProxyType.SOCKS4; + } + + @Override + protected void setUp() throws Exception { + initializeSocksServer(); + proxyServer = + bootstrapProxy() + .withName("Downstream") + .withPort(0) + .withChainProxyManager(chainedProxyManager()) + // misconfigured option + .withSendProxyProtocol(true) + .start(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyAuthenticationTest.java b/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyAuthenticationTest.java new file mode 100644 index 00000000..21ea6c3c --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyAuthenticationTest.java @@ -0,0 +1,278 @@ +package org.littleshoot.proxy; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.buildHttpClient; + +import io.netty.bootstrap.ServerBootstrap; +import io.netty.channel.*; +import io.netty.channel.nio.NioIoHandler; +import io.netty.channel.socket.nio.NioServerSocketChannel; +import io.netty.handler.codec.socksx.SocksMessage; +import io.netty.handler.codec.socksx.v5.DefaultSocks5InitialResponse; +import io.netty.handler.codec.socksx.v5.DefaultSocks5PasswordAuthResponse; +import io.netty.handler.codec.socksx.v5.Socks5AuthMethod; +import io.netty.handler.codec.socksx.v5.Socks5CommandRequest; +import io.netty.handler.codec.socksx.v5.Socks5CommandRequestDecoder; +import io.netty.handler.codec.socksx.v5.Socks5PasswordAuthRequest; +import io.netty.handler.codec.socksx.v5.Socks5PasswordAuthRequestDecoder; +import io.netty.handler.codec.socksx.v5.Socks5PasswordAuthStatus; +import java.net.InetSocketAddress; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import org.apache.http.HttpHost; +import org.apache.http.HttpResponse; +import org.apache.http.client.methods.HttpGet; +import org.apache.http.impl.client.CloseableHttpClient; +import org.apache.http.util.EntityUtils; +import org.eclipse.jetty.server.Server; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * Integration test for SOCKS5 proxy authentication (issue #57). Tests that upstream SOCKS5 proxy + * authentication works correctly. + * + *

This test verifies that: 1. LittleProxy correctly sends username/password authentication to an + * upstream SOCKS5 proxy 2. The authentication flow completes successfully 3. The credentials are + * transmitted correctly + */ +class Socks5ChainedProxyAuthenticationTest { + + private static final String SOCKS_USERNAME = "testuser"; + private static final String SOCKS_PASSWORD = "testpass"; + private static final String DEFAULT_JKS_KEYSTORE_PATH = "target/littleproxy_keystore.jks"; + private static final Logger log = + LoggerFactory.getLogger(Socks5ChainedProxyAuthenticationTest.class); + private EventLoopGroup socksBossGroup; + private EventLoopGroup socksWorkerGroup; + private int socksPort; + private Channel socksServerChannel; + + private Server webServer; + private int webServerPort; + private HttpProxyServer proxyServer; + + // Track authentication attempts for verification + private final AtomicInteger authAttempts = new AtomicInteger(0); + private final AtomicInteger authSuccesses = new AtomicInteger(0); + private final AtomicInteger connectRequests = new AtomicInteger(0); + + // Latch to track when we've received all expected messages + private final CountDownLatch authCompleteLatch = new CountDownLatch(1); + + @BeforeEach + void setUp() throws Exception { + // Start a simple web server + webServer = TestUtils.startWebServer(true, DEFAULT_JKS_KEYSTORE_PATH); + webServerPort = TestUtils.findLocalHttpPort(webServer); + + // Start SOCKS5 server with authentication + initializeSocksServer(); + + // Start LittleProxy with chained SOCKS5 proxy + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withName("Downstream") + .withPort(0) + .withChainProxyManager(chainedProxyManagerWithAuth()) + .start(); + } + + @AfterEach + void tearDown() throws Exception { + if (proxyServer != null) { + proxyServer.abort(); + } + if (webServer != null) { + webServer.stop(); + } + if (socksServerChannel != null) { + socksServerChannel.close(); + } + if (socksBossGroup != null) { + socksBossGroup.shutdownGracefully(); + } + if (socksWorkerGroup != null) { + socksWorkerGroup.shutdownGracefully(); + } + } + + private void initializeSocksServer() throws Exception { + socksBossGroup = new MultiThreadIoEventLoopGroup(NioIoHandler.newFactory()); + socksWorkerGroup = new MultiThreadIoEventLoopGroup(NioIoHandler.newFactory()); + + ServerBootstrap bootstrap = new ServerBootstrap(); + bootstrap + .group(socksBossGroup, socksWorkerGroup) + .channel(NioServerSocketChannel.class) + .childHandler(new Socks5AuthServerInitializer()); + + ChannelFuture channelFuture = bootstrap.bind(0).sync(); + socksServerChannel = channelFuture.channel(); + socksPort = ((InetSocketAddress) socksServerChannel.localAddress()).getPort(); + log.info("SOCKS5 server with auth started on port {}", socksPort); + } + + private ChainedProxyManager chainedProxyManagerWithAuth() { + return (httpRequest, chainedProxies, clientDetails) -> + chainedProxies.add( + new ChainedProxyAdapter() { + @Override + public InetSocketAddress getChainedProxyAddress() { + return new InetSocketAddress("127.0.0.1", socksPort); + } + + @Override + public ChainedProxyType getChainedProxyType() { + return ChainedProxyType.SOCKS5; + } + + @Override + public String getUsername() { + return SOCKS_USERNAME; + } + + @Override + public String getPassword() { + return SOCKS_PASSWORD; + } + }); + } + + /** + * Tests that SOCKS5 proxy authentication works correctly. + * + *

This test verifies the authentication flow between LittleProxy and an upstream SOCKS5 proxy + * that requires username/password authentication (issue #57). + * + *

The test demonstrates that: 1. LittleProxy correctly negotiates authentication with the + * SOCKS5 proxy 2. The username and password are transmitted correctly 3. Authentication succeeds + */ + @Test + void testSocks5Authentication() throws Exception { + // Make a request through the proxy - this will trigger the authentication + // We expect this to fail because our test server closes the connection after auth + // but that's fine - we're testing authentication, not the full proxy functionality + + // Use a very short connect timeout so the test doesn't hang + org.apache.http.client.config.RequestConfig config = + org.apache.http.client.config.RequestConfig.custom() + .setConnectTimeout(2000) + .setSocketTimeout(2000) + .build(); + + try (CloseableHttpClient httpClient = + buildHttpClient(true, true, proxyServer.getListenAddress().getPort(), null, null)) { + HttpGet request = new HttpGet("http://127.0.0.1:" + webServerPort + "/"); + request.setConfig(config); + try { + HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", webServerPort), request); + EntityUtils.consumeQuietly(response.getEntity()); + } catch (Exception e) { + // Expected - our test SOCKS server closes connection after auth + log.info("Request failed as expected: {}", e.getMessage()); + } + } + + // Wait for the authentication to complete + boolean authCompleted = authCompleteLatch.await(5, TimeUnit.SECONDS); + assertThat(authCompleted).as("Authentication should complete").isTrue(); + + // Verify authentication was attempted + assertThat(authAttempts.get()).as("Authentication attempts").isGreaterThanOrEqualTo(1); + + // Verify authentication succeeded + assertThat(authSuccesses.get()).as("Authentication successes").isEqualTo(1); + + // Verify we received a CONNECT request after successful auth + assertThat(connectRequests.get()).as("CONNECT requests").isGreaterThanOrEqualTo(1); + + log.info("=== TEST PASSED ==="); + log.info("SOCKS5 authentication flow works correctly!"); + log.info("Authentication attempts: {}", authAttempts.get()); + log.info("Authentication successes: {}", authSuccesses.get()); + log.info("CONNECT requests: {}", connectRequests.get()); + } + + /** Custom SOCKS5 server initializer that supports username/password authentication. */ + private class Socks5AuthServerInitializer extends ChannelInitializer { + @Override + protected void initChannel(Channel ch) { + // Use SocksPortUnificationServerHandler to detect SOCKS protocol version automatically + ch.pipeline().addLast(new io.netty.handler.codec.socksx.SocksPortUnificationServerHandler()); + ch.pipeline().addLast(new Socks5AuthServerHandler()); + } + } + + /** + * SOCKS5 server handler that implements username/password authentication. This handler requires + * authentication and tracks the authentication flow. + */ + private class Socks5AuthServerHandler extends SimpleChannelInboundHandler { + + @Override + protected void channelRead0(ChannelHandlerContext ctx, SocksMessage socksRequest) { + log.info("SOCKS: Received {}", socksRequest.getClass().getSimpleName()); + + if (socksRequest instanceof io.netty.handler.codec.socksx.v5.Socks5InitialRequest request) { + log.info("SOCKS: Client supported auth methods: {}", request.authMethods()); + + // Verify client supports both NO_AUTH and PASSWORD + assertThat(request.authMethods()).contains(Socks5AuthMethod.NO_AUTH); + assertThat(request.authMethods()).contains(Socks5AuthMethod.PASSWORD); + + // Require password authentication + ctx.pipeline().addFirst(new Socks5PasswordAuthRequestDecoder()); + ctx.writeAndFlush(new DefaultSocks5InitialResponse(Socks5AuthMethod.PASSWORD)); + log.info("SOCKS: Sent auth method response - requiring PASSWORD"); + + } else if (socksRequest instanceof Socks5PasswordAuthRequest authRequest) { + String username = authRequest.username(); + String password = authRequest.password(); + + log.info("SOCKS: Received password auth - username:'{}'", username); + + authAttempts.incrementAndGet(); + + // Verify credentials are correct + assertThat(username).as("Username").isEqualTo(SOCKS_USERNAME); + assertThat(password).as("Password").isEqualTo(SOCKS_PASSWORD); + + // Accept credentials + ctx.pipeline().addFirst(new Socks5CommandRequestDecoder()); + ctx.writeAndFlush(new DefaultSocks5PasswordAuthResponse(Socks5PasswordAuthStatus.SUCCESS)); + authSuccesses.incrementAndGet(); + log.info("SOCKS: Authentication SUCCESS!"); + + } else if (socksRequest instanceof Socks5CommandRequest request) { + log.info("SOCKS: CONNECT request to '{}:{}'", request.dstAddr(), request.dstPort()); + + connectRequests.incrementAndGet(); + + // Signal that auth is complete + authCompleteLatch.countDown(); + + // Close connection (test mode) + ctx.close(); + log.info("SOCKS: Closed connection (test mode)"); + } + } + + @Override + public void channelReadComplete(ChannelHandlerContext ctx) { + ctx.flush(); + } + + @Override + public void exceptionCaught(ChannelHandlerContext ctx, Throwable cause) { + log.info("SOCKS: Exception - {}", cause.getMessage()); + ctx.close(); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyTest.java index 21685609..ff4db9cf 100644 --- a/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyTest.java @@ -1,8 +1,8 @@ package org.littleshoot.proxy; -public class Socks5ChainedProxyTest extends BaseChainedSocksProxyTest { - @Override - protected ChainedProxyType getSocksProxyType() { - return ChainedProxyType.SOCKS5; - } +public final class Socks5ChainedProxyTest extends BaseChainedSocksProxyTest { + @Override + protected ChainedProxyType getSocksProxyType() { + return ChainedProxyType.SOCKS5; + } } diff --git a/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyWithMissConfiguredSendProxyProtocolTest.java b/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyWithMissConfiguredSendProxyProtocolTest.java new file mode 100644 index 00000000..567ad1b0 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/Socks5ChainedProxyWithMissConfiguredSendProxyProtocolTest.java @@ -0,0 +1,25 @@ +package org.littleshoot.proxy; + +import org.junit.jupiter.api.Tag; + +@Tag("slow-test") +public final class Socks5ChainedProxyWithMissConfiguredSendProxyProtocolTest + extends BaseChainedSocksProxyTest { + @Override + protected ChainedProxyType getSocksProxyType() { + return ChainedProxyType.SOCKS5; + } + + @Override + protected void setUp() throws Exception { + initializeSocksServer(); + proxyServer = + bootstrapProxy() + .withName("Downstream") + .withPort(0) + .withChainProxyManager(chainedProxyManager()) + // misconfigured option + .withSendProxyProtocol(true) + .start(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/StopProxyTest.java b/src/test/java/org/littleshoot/proxy/StopProxyTest.java index 1a37081b..3142965f 100644 --- a/src/test/java/org/littleshoot/proxy/StopProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/StopProxyTest.java @@ -1,24 +1,20 @@ package org.littleshoot.proxy; -import org.junit.Test; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -public class StopProxyTest { - @Test - public void testStop() { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); +public final class StopProxyTest { + @Test + public void testStop() { + HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - proxyServer.stop(); - } + proxyServer.stop(); + } - @Test - public void testAbort() { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); + @Test + public void testAbort() { + HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - proxyServer.abort(); - } + proxyServer.abort(); + } } diff --git a/src/test/java/org/littleshoot/proxy/TestUtils.java b/src/test/java/org/littleshoot/proxy/TestUtils.java index f2332b50..2be13744 100644 --- a/src/test/java/org/littleshoot/proxy/TestUtils.java +++ b/src/test/java/org/littleshoot/proxy/TestUtils.java @@ -1,6 +1,24 @@ package org.littleshoot.proxy; +import static org.assertj.core.api.Assumptions.assumeThat; + import com.sun.management.UnixOperatingSystemMXBean; +import jakarta.servlet.http.HttpServletRequest; +import jakarta.servlet.http.HttpServletResponse; +import java.io.BufferedInputStream; +import java.io.IOException; +import java.io.InputStream; +import java.lang.management.ManagementFactory; +import java.lang.management.MemoryUsage; +import java.lang.management.OperatingSystemMXBean; +import java.net.InetSocketAddress; +import java.net.ServerSocket; +import java.security.KeyManagementException; +import java.security.KeyStoreException; +import java.security.NoSuchAlgorithmException; +import java.security.SecureRandom; +import java.util.Objects; +import javax.net.ssl.SSLContext; import org.apache.http.HttpHost; import org.apache.http.auth.AuthScope; import org.apache.http.auth.UsernamePasswordCredentials; @@ -12,315 +30,348 @@ import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.impl.client.HttpClientBuilder; import org.apache.http.ssl.SSLContextBuilder; -import org.eclipse.jetty.server.Connector; -import org.eclipse.jetty.server.Request; -import org.eclipse.jetty.server.Server; -import org.eclipse.jetty.server.ServerConnector; +import org.eclipse.jetty.server.*; import org.eclipse.jetty.server.handler.AbstractHandler; import org.eclipse.jetty.util.ssl.SslContextFactory; import org.littleshoot.proxy.extras.SelfSignedSslEngineSource; -import javax.net.ssl.SSLContext; -import javax.servlet.http.HttpServletRequest; -import javax.servlet.http.HttpServletResponse; -import java.io.BufferedInputStream; -import java.io.IOException; -import java.io.InputStream; -import java.lang.management.ManagementFactory; -import java.lang.management.MemoryUsage; -import java.lang.management.OperatingSystemMXBean; -import java.net.InetSocketAddress; -import java.net.ServerSocket; -import java.security.SecureRandom; -import java.util.Objects; - public class TestUtils { - public static final RequestConfig REQUEST_TIMEOUT_CONFIG = RequestConfig.custom().setConnectTimeout(5000).build(); - - private TestUtils() { - } - - /** - * Creates and starts an embedded web server on a JVM-assigned HTTP ports. - * Each response has a body that indicates how many bytes were received with - * a message like "Received x bytes\n". - * - * @return Instance of Server - */ - public static Server startWebServer() { - return startWebServer(false); + public static final RequestConfig REQUEST_TIMEOUT_CONFIG = + RequestConfig.custom().setConnectTimeout(5000).build(); + + private TestUtils() {} + + /** + * Creates and starts an embedded web server on a JVM-assigned HTTP ports. Each response has a + * body that indicates how many bytes were received with a message like "Received x bytes\n". + * + * @return Instance of Server + */ + public static Server startWebServer() { + return startWebServer(false, "target/littleproxy_keystore.jks"); + } + + /** + * Creates and starts an embedded web server on a JVM-assigned HTTP ports. Creates and starts + * embedded web server that is running on given port. Each response has a body that contains the + * specified contents. + * + * @return Instance of Server + */ + public static Server startWebServerWithResponse(byte[] content) { + return startWebServerWithResponse(false, content); + } + + /** + * Creates and starts an embedded web server on JVM-assigned HTTP and HTTPS ports. Each response + * has a body that indicates how many bytes were received with a message like "Received x + * bytes\n". + * + * @param enableHttps if true, an HTTPS connector will be added to the web server + * @param keyStorePath path to the keystore file for HTTPS + * @return Instance of Server + */ + public static Server startWebServer(boolean enableHttps, String keyStorePath) { + final Server httpServer = createWebServer(); + + if (enableHttps) { + // Add SSL connector + SelfSignedSslEngineSource contextSource = new SelfSignedSslEngineSource(keyStorePath); + ServerConnector connector = createServerConnector(contextSource, httpServer); + connector.setPort(0); + connector.setIdleTimeout(0); + httpServer.addConnector(connector); } - /** - * Creates and starts an embedded web server on a JVM-assigned HTTP ports. - * Creates and starts embedded web server that is running on given port. - * Each response has a body that contains the specified contents. - * - * @return Instance of Server - */ - public static Server startWebServerWithResponse(byte[] content) { - return startWebServerWithResponse(false, content); + try { + httpServer.start(); + } catch (Exception e) { + throw new RuntimeException("Error starting Jetty web server", e); } - /** - * Creates and starts an embedded web server on JVM-assigned HTTP and HTTPS ports. - * Each response has a body that indicates how many bytes were received with a message like - * "Received x bytes\n". - * - * @param enableHttps if true, an HTTPS connector will be added to the web server - * @return Instance of Server - */ - public static Server startWebServer(boolean enableHttps) { - final Server httpServer = new Server(0); - - httpServer.setHandler(new AbstractHandler() { - public void handle(String target, - Request baseRequest, - HttpServletRequest request, - HttpServletResponse response) throws IOException { - if (request.getRequestURI().contains("hang")) { - System.out.println("Hanging as requested"); - try { - Thread.sleep(90000); - } catch (InterruptedException ie) { - System.out.println("Stopped hanging due to interruption"); - } - } - - long numberOfBytesRead = 0; - try (InputStream in = new BufferedInputStream(request.getInputStream())) { - while (in.read() != -1) { - numberOfBytesRead += 1; - } - } - System.out.println("Done reading # of bytes: " - + numberOfBytesRead); - response.setStatus(HttpServletResponse.SC_OK); - baseRequest.setHandled(true); - byte[] content = ("Received " + numberOfBytesRead + " bytes\n").getBytes(); - response.addHeader("Content-Length", Integer.toString(content.length)); - response.getOutputStream().write(content); + return httpServer; + } + + private static ServerConnector createServerConnector( + SelfSignedSslEngineSource contextSource, Server httpServer) { + SSLContext sslContext = contextSource.getSslContext(); + + SslContextFactory.Server sslContextFactory = new SslContextFactory.Server(); + sslContextFactory.setSslContext(sslContext); + + SecureRequestCustomizer secureRequestCustomizer = new SecureRequestCustomizer(); + secureRequestCustomizer.setSniHostCheck(false); + + HttpConfiguration httpConfiguration = new HttpConfiguration(); + httpConfiguration.addCustomizer(secureRequestCustomizer); + + HttpConnectionFactory httpConnectionFactory = new HttpConnectionFactory(httpConfiguration); + return new ServerConnector(httpServer, sslContextFactory, httpConnectionFactory); + } + + private static Server createWebServer() { + final Server httpServer = new Server(0); + + httpServer.setHandler( + new AbstractHandler() { + @Override + public void handle( + String target, + Request baseRequest, + HttpServletRequest request, + HttpServletResponse response) + throws IOException { + if (request.getRequestURI().contains("hang")) { + System.out.println("Hanging as requested"); + try { + Thread.sleep(90000); + } catch (InterruptedException ie) { + System.out.println("Stopped hanging due to interruption"); + } } - }); - - if (enableHttps) { - // Add SSL connector - SslContextFactory sslContextFactory = new SslContextFactory.Server(); - - SelfSignedSslEngineSource contextSource = new SelfSignedSslEngineSource(); - SSLContext sslContext = contextSource.getSslContext(); - - sslContextFactory.setSslContext(sslContext); - ServerConnector connector = new ServerConnector(httpServer, sslContextFactory); - connector.setPort(0); - connector.setIdleTimeout(0); - httpServer.addConnector(connector); - } - - try { - httpServer.start(); - } catch (Exception e) { - throw new RuntimeException("Error starting Jetty web server", e); - } - - return httpServer; - } - /** - * Creates and starts an embedded web server on JVM-assigned HTTP and HTTPS ports. - * Each response has a body that contains the specified contents. - * - * @param enableHttps if true, an HTTPS connector will be added to the web server - * @param content The response the server will return - * @return Instance of Server - */ - public static Server startWebServerWithResponse(boolean enableHttps, final byte[] content) { - final Server httpServer = new Server(0); - httpServer.setHandler(new AbstractHandler() { - public void handle(String target, - Request baseRequest, - HttpServletRequest request, - HttpServletResponse response) throws IOException { - if (request.getRequestURI().contains("hang")) { - System.out.println("Hanging as requested"); - try { - Thread.sleep(90000); - } catch (InterruptedException ie) { - System.out.println("Stopped hanging due to interruption"); - } - } - - long numberOfBytesRead = 0; - try (InputStream in = new BufferedInputStream(request.getInputStream())) { - while (in.read() != -1) { - numberOfBytesRead += 1; - } - } - System.out.println("Done reading # of bytes: " - + numberOfBytesRead); - response.setStatus(HttpServletResponse.SC_OK); - baseRequest.setHandled(true); - - response.addHeader("Content-Length", Integer.toString(content.length)); - response.getOutputStream().write(content); + long numberOfBytesRead = 0; + try (InputStream in = new BufferedInputStream(request.getInputStream())) { + while (in.read() != -1) { + numberOfBytesRead += 1; + } } + System.out.println("Done reading # of bytes: " + numberOfBytesRead); + response.setStatus(HttpServletResponse.SC_OK); + baseRequest.setHandled(true); + byte[] content = ("Received " + numberOfBytesRead + " bytes\n").getBytes(); + response.addHeader("Content-Length", Integer.toString(content.length)); + response.getOutputStream().write(content); + } }); + return httpServer; + } + + /** + * Creates and starts an embedded web server on JVM-assigned HTTP and HTTPS ports. Each response + * has a body that contains the specified contents. + * + * @param enableHttps if true, an HTTPS connector will be added to the web server + * @param content The response the server will return + * @return Instance of Server + */ + public static Server startWebServerWithResponse(boolean enableHttps, final byte[] content) { + final Server httpServer = new Server(0); + httpServer.setHandler( + new AbstractHandler() { + @Override + public void handle( + String target, + Request baseRequest, + HttpServletRequest request, + HttpServletResponse response) + throws IOException { + if (request.getRequestURI().contains("hang")) { + System.out.println("Hanging as requested"); + try { + Thread.sleep(90000); + } catch (InterruptedException ie) { + System.out.println("Stopped hanging due to interruption"); + } + } - if (enableHttps) { - // Add SSL connector - SslContextFactory sslContextFactory = new SslContextFactory.Server(); + long numberOfBytesRead = 0; + try (InputStream in = new BufferedInputStream(request.getInputStream())) { + while (in.read() != -1) { + numberOfBytesRead += 1; + } + } + System.out.println("Done reading # of bytes: " + numberOfBytesRead); + response.setStatus(HttpServletResponse.SC_OK); + baseRequest.setHandled(true); - SelfSignedSslEngineSource contextSource = new SelfSignedSslEngineSource(); - SSLContext sslContext = contextSource.getSslContext(); + response.addHeader("Content-Length", Integer.toString(content.length)); + response.getOutputStream().write(content); + } + }); - sslContextFactory.setSslContext(sslContext); - ServerConnector connector = new ServerConnector(httpServer, sslContextFactory); - connector.setPort(0); - connector.setIdleTimeout(0); - httpServer.addConnector(connector); - } + if (enableHttps) { + // Add SSL connector + SslContextFactory.Server sslContextFactory = new SslContextFactory.Server(); - try { - httpServer.start(); - } catch (Exception e) { - throw new RuntimeException("Error starting Jetty web server", e); - } + SelfSignedSslEngineSource contextSource = + new SelfSignedSslEngineSource("target/littleproxy_keystore.jks"); + SSLContext sslContext = contextSource.getSslContext(); - return httpServer; + sslContextFactory.setSslContext(sslContext); + ServerConnector connector = new ServerConnector(httpServer, sslContextFactory); + connector.setPort(0); + connector.setIdleTimeout(0); + httpServer.addConnector(connector); } - /** - * Finds the port the specified server is listening for HTTP connections on. - * - * @param webServer started web server - * @return HTTP port, or -1 if no HTTP port was found - */ - public static int findLocalHttpPort(Server webServer) { - for (Connector connector : webServer.getConnectors()) { - if (!Objects.equals(connector.getDefaultConnectionFactory().getProtocol(), "SSL")) { - return ((ServerConnector) connector).getLocalPort(); - } - } - - return -1; + try { + httpServer.start(); + } catch (Exception e) { + throw new RuntimeException("Error starting Jetty web server", e); } - /** - * Finds the port the specified server is listening for HTTPS connections on. - * - * @param webServer started web server - * @return HTTP port, or -1 if no HTTPS port was found - */ - public static int findLocalHttpsPort(Server webServer) { - for (Connector connector : webServer.getConnectors()) { - if (Objects.equals(connector.getDefaultConnectionFactory().getProtocol(), "SSL")) { - return ((ServerConnector) connector).getLocalPort(); - } - } + return httpServer; + } + + /** + * Finds the port the specified server is listening for HTTP connections on. + * + * @param webServer started web server + * @return HTTP port, or -1 if no HTTP port was found + */ + public static int findLocalHttpPort(Server webServer) { + for (Connector connector : webServer.getConnectors()) { + if (!Objects.equals(connector.getDefaultConnectionFactory().getProtocol(), "SSL")) { + return ((ServerConnector) connector).getLocalPort(); + } + } - return -1; + return -1; + } + + /** + * Finds the port the specified server is listening for HTTPS connections on. + * + * @param webServer started web server + * @return HTTP port, or -1 if no HTTPS port was found + */ + public static int findLocalHttpsPort(Server webServer) { + for (Connector connector : webServer.getConnectors()) { + if (Objects.equals(connector.getDefaultConnectionFactory().getProtocol(), "SSL")) { + return ((ServerConnector) connector).getLocalPort(); + } } - /** - * Creates instance HttpClient that is configured to use proxy server. The - * proxy server should run on 127.0.0.1 and given port - * - * @param port - * the proxy port - * @return instance of HttpClient - */ - public static CloseableHttpClient createProxiedHttpClient(final int port) throws Exception { - return buildHttpClient(true, false, port, null, null); + return -1; + } + + /** + * Creates instance HttpClient that is configured to use proxy server. The proxy server should run + * on 127.0.0.1 and given port + * + * @param port the proxy port + * @return instance of HttpClient + */ + public static CloseableHttpClient createProxiedHttpClient(final int port) { + return buildHttpClient(true, false, port, null, null); + } + + public static int randomPort() { + final SecureRandom secureRandom = new SecureRandom(); + for (int i = 0; i < 20; i++) { + // The +1 on the random int is because + // Math.abs(Integer.MIN_VALUE) == Integer.MIN_VALUE -- caught + // by FindBugs. + final int randomPort = 1024 + (Math.abs(secureRandom.nextInt() + 1) % 60000); + try (ServerSocket sock = new ServerSocket()) { + sock.bind(new InetSocketAddress("127.0.0.1", randomPort)); + return sock.getLocalPort(); + } catch (final IOException ignored) { + } } - public static int randomPort() { - final SecureRandom secureRandom = new SecureRandom(); - for (int i = 0; i < 20; i++) { - // The +1 on the random int is because - // Math.abs(Integer.MIN_VALUE) == Integer.MIN_VALUE -- caught - // by FindBugs. - final int randomPort = 1024 + (Math.abs(secureRandom.nextInt() + 1) % 60000); - try (ServerSocket sock = new ServerSocket()) { - sock.bind(new InetSocketAddress("127.0.0.1", randomPort)); - return sock.getLocalPort(); - } catch (final IOException ignored) { - } - } - - // If we can't grab one of our securely chosen random ports, use - // whatever port the OS assigns. - try (ServerSocket sock = new ServerSocket()) { - sock.bind(null); - return sock.getLocalPort(); - } catch (final IOException e) { - return 1024 + (Math.abs(secureRandom.nextInt() + 1) % 60000); - } + // If we can't grab one of our securely chosen random ports, use + // whatever port the OS assigns. + try (ServerSocket sock = new ServerSocket()) { + sock.bind(null); + return sock.getLocalPort(); + } catch (final IOException e) { + return 1024 + (Math.abs(secureRandom.nextInt() + 1) % 60000); } - - public static long getOpenFileDescriptorsAndPrintMemoryUsage() { - // Below courtesy of: - // http://stackoverflow.com/questions/10999076/programmatically-print-the-heap-usage-that-is-typically-printed-on-jvm-exit-when - MemoryUsage mu = ManagementFactory.getMemoryMXBean() - .getHeapMemoryUsage(); - MemoryUsage muNH = ManagementFactory.getMemoryMXBean() - .getNonHeapMemoryUsage(); - System.out.println("Init :" + mu.getInit() + "\nMax :" + mu.getMax() - + "\nUsed :" + mu.getUsed() + "\nCommitted :" - + mu.getCommitted() + "\nInit NH :" + muNH.getInit() - + "\nMax NH :" + muNH.getMax() + "\nUsed NH:" + muNH.getUsed() - + "\nCommitted NH:" + muNH.getCommitted()); - - OperatingSystemMXBean osMxBean = ManagementFactory.getOperatingSystemMXBean(); - - if (osMxBean instanceof UnixOperatingSystemMXBean) { - UnixOperatingSystemMXBean unixOsMxBean = (UnixOperatingSystemMXBean) osMxBean; - return unixOsMxBean.getOpenFileDescriptorCount(); - } else { - throw new UnsupportedOperationException("Unable to determine number of open file handles on non-Unix system"); - } + } + + public static long getOpenFileDescriptorsAndPrintMemoryUsage() { + // Below courtesy of: + // http://stackoverflow.com/questions/10999076/programmatically-print-the-heap-usage-that-is-typically-printed-on-jvm-exit-when + MemoryUsage mu = ManagementFactory.getMemoryMXBean().getHeapMemoryUsage(); + MemoryUsage muNH = ManagementFactory.getMemoryMXBean().getNonHeapMemoryUsage(); + System.out.println( + "Init :" + + mu.getInit() + + "\nMax :" + + mu.getMax() + + "\nUsed :" + + mu.getUsed() + + "\nCommitted :" + + mu.getCommitted() + + "\nInit NH :" + + muNH.getInit() + + "\nMax NH :" + + muNH.getMax() + + "\nUsed NH:" + + muNH.getUsed() + + "\nCommitted NH:" + + muNH.getCommitted()); + + OperatingSystemMXBean osMxBean = ManagementFactory.getOperatingSystemMXBean(); + + if (osMxBean instanceof UnixOperatingSystemMXBean unixOsMxBean) { + return unixOsMxBean.getOpenFileDescriptorCount(); + } else { + throw new UnsupportedOperationException( + "Unable to determine number of open file handles on non-Unix system"); + } + } + + /** + * Determines if we are running on a Unix-like operating system that exposes a {@link + * com.sun.management.UnixOperatingSystemMXBean}. + * + * @return true if this is a Unix OS and the JVM exposes a {@link + * com.sun.management.UnixOperatingSystemMXBean}, otherwise false. + */ + public static boolean isUnixManagementCapable() { + OperatingSystemMXBean osMxBean = ManagementFactory.getOperatingSystemMXBean(); + + return (osMxBean instanceof UnixOperatingSystemMXBean); + } + + public static void requireUnix() { + assumeThat(isUnixManagementCapable()).as("Skipping test on non-Unix OS").isTrue(); + } + + public static void disableOnMac() { + assumeThat(System.getProperty("os.name")) + .as("Skipping test on Mac OS") + .doesNotContainIgnoringCase("mac"); + } + + /** + * Creates a DefaultHttpClient instance. + * + * @return instance of ClosableHttpClient + */ + public static CloseableHttpClient buildHttpClient( + boolean isProxied, boolean supportSsl, int proxyPort, String username, String password) { + + HttpClientBuilder builder = HttpClientBuilder.create().setSSLContext(createSslContext()); + + if (supportSsl) { + builder.setSSLHostnameVerifier(new NoopHostnameVerifier()); } - /** - * Determines if we are running on a Unix-like operating system that exposes a {@link com.sun.management.UnixOperatingSystemMXBean}. - * - * @return true if this is a Unix OS and the JVM exposes a {@link com.sun.management.UnixOperatingSystemMXBean}, otherwise false. - */ - public static boolean isUnixManagementCapable() { - OperatingSystemMXBean osMxBean = ManagementFactory.getOperatingSystemMXBean(); - - return (osMxBean instanceof UnixOperatingSystemMXBean); + if (isProxied) { + HttpHost proxy = new HttpHost("127.0.0.1", proxyPort); + builder.setProxy(proxy); + if (username != null && password != null) { + CredentialsProvider credentialsProvider = new BasicCredentialsProvider(); + credentialsProvider.setCredentials( + new AuthScope("127.0.0.1", proxyPort), + new UsernamePasswordCredentials(username, password)); + builder.setDefaultCredentialsProvider(credentialsProvider); + } } - /** - * Creates a DefaultHttpClient instance. - * - * @return instance of ClosableHttpClient - */ - public static CloseableHttpClient buildHttpClient(boolean isProxied, boolean supportSsl, int proxyPort, - String username, String password) throws Exception { - - HttpClientBuilder builder = HttpClientBuilder.create() - .setSSLContext(SSLContextBuilder.create() - .loadTrustMaterial(new TrustSelfSignedStrategy()) - .build()); - - if(supportSsl){ - builder.setSSLHostnameVerifier(new NoopHostnameVerifier()); - } - - if (isProxied) { - HttpHost proxy = new HttpHost("127.0.0.1", proxyPort); - builder.setProxy(proxy); - if (username != null && password != null) { - CredentialsProvider credentialsProvider = new BasicCredentialsProvider(); - credentialsProvider.setCredentials( - new AuthScope("127.0.0.1", proxyPort), - new UsernamePasswordCredentials(username, password)); - builder.setDefaultCredentialsProvider(credentialsProvider); - } - } + return builder.build(); + } - return builder.build(); + private static SSLContext createSslContext() { + try { + return SSLContextBuilder.create().loadTrustMaterial(new TrustSelfSignedStrategy()).build(); + } catch (NoSuchAlgorithmException | KeyManagementException | KeyStoreException e) { + throw new RuntimeException(e); } + } } diff --git a/src/test/java/org/littleshoot/proxy/ThrottledInputStream.java b/src/test/java/org/littleshoot/proxy/ThrottledInputStream.java index 1a437e2e..5cef8b1e 100644 --- a/src/test/java/org/littleshoot/proxy/ThrottledInputStream.java +++ b/src/test/java/org/littleshoot/proxy/ThrottledInputStream.java @@ -21,14 +21,12 @@ import java.io.InputStream; /** - * The ThrottleInputStream provides bandwidth throttling on a specified - * InputStream. It is implemented as a wrapper on top of another InputStream - * instance. - * The throttling works by examining the number of bytes read from the underlying - * InputStream from the beginning, and sleep()ing for a time interval if - * the byte-transfer is found exceed the specified tolerable maximum. - * (Thus, while the read-rate might exceed the maximum for a given short interval, - * the average tends towards the specified maximum, overall.) + * The ThrottleInputStream provides bandwidth throttling on a specified InputStream. It is + * implemented as a wrapper on top of another InputStream instance. The throttling works by + * examining the number of bytes read from the underlying InputStream from the beginning, and + * sleep()ing for a time interval if the byte-transfer is found exceed the specified tolerable + * maximum. (Thus, while the read-rate might exceed the maximum for a given short interval, the + * average tends towards the specified maximum, overall.) */ public class ThrottledInputStream extends InputStream { @@ -36,8 +34,8 @@ public class ThrottledInputStream extends InputStream { private final long maxBytesPerSec; private final long startTime = System.currentTimeMillis(); - private long bytesRead = 0; - private long totalSleepTime = 0; + private long bytesRead; + private long totalSleepTime; private static final long SLEEP_DURATION_MS = 50; @@ -46,7 +44,7 @@ public ThrottledInputStream(InputStream rawStream) { } public ThrottledInputStream(InputStream rawStream, long maxBytesPerSec) { - assert maxBytesPerSec > 0 : "Bandwidth " + maxBytesPerSec + " is invalid"; + assert maxBytesPerSec > 0 : "Bandwidth " + maxBytesPerSec + " is invalid"; this.rawStream = rawStream; this.maxBytesPerSec = maxBytesPerSec; } @@ -97,6 +95,7 @@ private void throttle() throws IOException { /** * Getter for the number of bytes read from this stream, since creation. + * * @return The number of bytes. */ public long getTotalBytesRead() { @@ -104,8 +103,9 @@ public long getTotalBytesRead() { } /** - * Getter for the read-rate from this stream, since creation. - * Calculated as bytesRead/elapsedTimeSinceStart. + * Getter for the read-rate from this stream, since creation. Calculated as + * bytesRead/elapsedTimeSinceStart. + * * @return Read rate, in bytes/sec. */ public long getBytesPerSec() { @@ -119,6 +119,7 @@ public long getBytesPerSec() { /** * Getter the total time spent in sleep. + * * @return Number of milliseconds spent in sleep. */ public long getTotalSleepTime() { @@ -128,11 +129,15 @@ public long getTotalSleepTime() { /** {@inheritDoc} */ @Override public String toString() { - return "ThrottledInputStream{" + - "bytesRead=" + bytesRead + - ", maxBytesPerSec=" + maxBytesPerSec + - ", bytesPerSec=" + getBytesPerSec() + - ", totalSleepTime=" + totalSleepTime + - '}'; + return "ThrottledInputStream{" + + "bytesRead=" + + bytesRead + + ", maxBytesPerSec=" + + maxBytesPerSec + + ", bytesPerSec=" + + getBytesPerSec() + + ", totalSleepTime=" + + totalSleepTime + + '}'; } -} \ No newline at end of file +} diff --git a/src/test/java/org/littleshoot/proxy/ThrottlingTest.java b/src/test/java/org/littleshoot/proxy/ThrottlingTest.java index 8233ccaf..8c4d04d1 100644 --- a/src/test/java/org/littleshoot/proxy/ThrottlingTest.java +++ b/src/test/java/org/littleshoot/proxy/ThrottlingTest.java @@ -1,344 +1,409 @@ package org.littleshoot.proxy; +import static org.apache.http.client.utils.HttpClientUtils.closeQuietly; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.createProxiedHttpClient; + +import java.util.Arrays; import org.apache.http.HttpHost; import org.apache.http.client.methods.HttpGet; import org.apache.http.client.methods.HttpPost; -import org.apache.http.client.utils.HttpClientUtils; import org.apache.http.entity.ByteArrayEntity; import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.util.EntityUtils; import org.eclipse.jetty.server.Server; -import org.junit.After; -import org.junit.Before; -import org.junit.FixMethodOrder; -import org.junit.Test; -import org.junit.runners.MethodSorters; +import org.junit.jupiter.api.*; +import org.junit.jupiter.api.parallel.Execution; +import org.junit.jupiter.api.parallel.ExecutionMode; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; - -import static org.hamcrest.Matchers.both; -import static org.hamcrest.Matchers.greaterThan; -import static org.hamcrest.Matchers.lessThan; -import static org.junit.Assert.assertEquals; -import static org.junit.Assert.assertThat; - -@FixMethodOrder(MethodSorters.JVM) -public class ThrottlingTest { - private static final long THROTTLED_READ_BYTES_PER_SECOND = 25000L; - private static final long THROTTLED_WRITE_BYTES_PER_SECOND = 25000L; - - // throttling is not guaranteed to be exact, so allow some variation in the amount of time the call takes. since we want - // these tests to take just a few seconds, allow significant variation. even with this large variation, if throttling - // is broken it should take much less time than expected. - private static final double ALLOWABLE_VARIATION = 0.30; - - private Server writeWebServer; - private Server readWebServer; - - private byte[] largeData; - - private int msToWriteThrottled; - private int msToReadThrottled; - - // time to allow for an unthrottled local request - private static final int UNTRHOTTLED_REQUEST_TIME_MS = 1000; - - private int writeWebServerPort; - private int readWebServerPort; - - @Before - public void setUp() { - // Set up some large data - largeData = new byte[100000]; - for (int i = 0; i < largeData.length; i++) { - largeData[i] = 1 % 256; +import org.littleshoot.proxy.test.EnableThreadDump; + +@Tag("slow-test") +@Timeout(25) +@Execution(ExecutionMode.SAME_THREAD) +@EnableThreadDump +public final class ThrottlingTest { + private static final int LARGE_DATA_SIZE = 200000; + private static final long THROTTLED_READ_BYTES_PER_SECOND = 25000L; + private static final long THROTTLED_WRITE_BYTES_PER_SECOND = 25000L; + + // throttling is not guaranteed to be exact, so allow some variation in the + // amount of time the call takes. since we want + // these tests to take just a few seconds, allow significant variation. even + // with this large variation, if throttling + // is broken it should take much less time than expected. + private static final double ALLOWABLE_VARIATION = 0.30; + private static final String DEFAULT_JKS_KEYSTORE_PATH = "target/littleproxy_keystore.jks"; + + private HttpProxyServer proxyServer; + private Server writeWebServer; + private Server readWebServer; + + private byte[] largeData; + + private int msToWriteThrottled; + private int msToReadThrottled; + + // time to allow for an unthrottled local request + private static final long UNTHROTTLED_REQUEST_TIME_MS = 1500; + + private int writeWebServerPort; + private int readWebServerPort; + + @BeforeEach + void setUp() { + // Set up some large data + largeData = new byte[LARGE_DATA_SIZE]; + Arrays.fill(largeData, (byte) (1 % 256)); + + msToWriteThrottled = largeData.length * 1000 / (int) THROTTLED_WRITE_BYTES_PER_SECOND; + msToReadThrottled = largeData.length * 1000 / (int) THROTTLED_READ_BYTES_PER_SECOND; + + writeWebServer = TestUtils.startWebServer(false, DEFAULT_JKS_KEYSTORE_PATH); + writeWebServerPort = TestUtils.findLocalHttpPort(writeWebServer); + + readWebServer = TestUtils.startWebServerWithResponse(false, largeData); + readWebServerPort = TestUtils.findLocalHttpPort(readWebServer); + } + + @AfterEach + void tearDown() throws Exception { + try { + if (proxyServer != null) { + proxyServer.abort(); + } + } finally { + try { + if (writeWebServer != null) { + writeWebServer.stop(); } - - msToWriteThrottled = largeData.length * 1000 / (int)THROTTLED_WRITE_BYTES_PER_SECOND; - msToReadThrottled = largeData.length * 1000 / (int)THROTTLED_READ_BYTES_PER_SECOND; - - writeWebServer = TestUtils.startWebServer(false); - writeWebServerPort = TestUtils.findLocalHttpPort(writeWebServer); - - readWebServer = TestUtils.startWebServerWithResponse(false, largeData); - readWebServerPort = TestUtils.findLocalHttpPort(readWebServer); - } - - @After - public void tearDown() throws Exception { - try { - if (writeWebServer != null) { - writeWebServer.stop(); - } - } finally { - if (readWebServer != null) { - readWebServer.stop(); - } + } finally { + if (readWebServer != null) { + readWebServer.stop(); } + } } - - @Test - public void aWarmUpTest() throws Exception { - // a "warm-up" test so the first test's results are not skewed due to classloading, etc. guaranteed to run - // first with the @FixMethodOrder(MethodSorters.NAME_ASCENDING) annotation on the class. - - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withThrottling(0, THROTTLED_WRITE_BYTES_PER_SECOND) - .start(); - - int proxyPort = proxyServer.getListenAddress().getPort(); - - HttpGet request = createHttpGet(); - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyPort); - - EntityUtils.consumeQuietly(httpClient.execute(new HttpHost("127.0.0.1", writeWebServerPort), request).getEntity()); - - EntityUtils.consumeQuietly(httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request).getEntity()); - } - - @Test - public void testThrottledWrite() throws Exception { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withThrottling(0, THROTTLED_WRITE_BYTES_PER_SECOND) - .start(); - - int proxyPort = proxyServer.getListenAddress().getPort(); - - final HttpPost request = createHttpPost(); - - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyPort); - - long start = System.currentTimeMillis(); - final org.apache.http.HttpResponse response = httpClient.execute( - new HttpHost("127.0.0.1", - writeWebServerPort), request); - long finish = System.currentTimeMillis(); - - assertEquals("Received " + largeData.length + " bytes\n", - EntityUtils.toString(response.getEntity())); - - assertThat("Expected throttled write to complete in approximately " + msToWriteThrottled + "ms" + " but took " + (finish - start) + "ms", - (double)(finish - start), both(greaterThan(msToWriteThrottled * (1 - ALLOWABLE_VARIATION))).and( - lessThan(msToWriteThrottled * (1 + ALLOWABLE_VARIATION)))); - - proxyServer.abort(); - } - - @Test - public void testUnthrottledWrite() throws Exception { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - - int proxyPort = proxyServer.getListenAddress().getPort(); - - final HttpPost request = createHttpPost(); - - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyPort); - - long start = System.currentTimeMillis(); - final org.apache.http.HttpResponse response = httpClient.execute( - new HttpHost("127.0.0.1", - writeWebServerPort), request); - long finish = System.currentTimeMillis(); - - assertEquals("Received " + largeData.length + " bytes\n", - EntityUtils.toString(response.getEntity())); - - assertThat("Unthrottled write took " + (finish - start) + "ms, but expected to complete in " + UNTRHOTTLED_REQUEST_TIME_MS + "ms", - finish - start, lessThan((long) UNTRHOTTLED_REQUEST_TIME_MS)); - - proxyServer.abort(); + } + + @Test + public void aWarmUpTest() throws Exception { + // a "warm-up" test so the first test's results are not skewed due to + // classloading, etc. guaranteed to run + // first with the @FixMethodOrder(MethodSorters.NAME_ASCENDING) annotation on + // the class. + + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withThrottling(0, THROTTLED_WRITE_BYTES_PER_SECOND) + .start(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + + HttpGet request = createHttpGet(); + try (CloseableHttpClient httpClient = createProxiedHttpClient(proxyPort)) { + EntityUtils.consumeQuietly( + httpClient.execute(new HttpHost("127.0.0.1", writeWebServerPort), request).getEntity()); + EntityUtils.consumeQuietly( + httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request).getEntity()); } + } - @Test - public void testThrottledRead() throws Exception { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withThrottling(THROTTLED_READ_BYTES_PER_SECOND, 0) - .start(); - - int proxyPort = proxyServer.getListenAddress().getPort(); + @Test + public void testThrottledWrite() throws Exception { + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withThrottling(0, THROTTLED_WRITE_BYTES_PER_SECOND) + .start(); - final HttpGet request = createHttpGet(); + int proxyPort = proxyServer.getListenAddress().getPort(); - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyPort); + final HttpPost request = createHttpPost(); - long start = System.currentTimeMillis(); - final org.apache.http.HttpResponse response = httpClient.execute( - new HttpHost("127.0.0.1", - readWebServerPort), request); - byte[] readContent = new byte[100000]; - - int bytesRead = 0; - while (bytesRead < largeData.length) { - int read = response.getEntity().getContent().read(readContent, bytesRead, largeData.length - bytesRead); - bytesRead += read; - } + try (CloseableHttpClient httpClient = createProxiedHttpClient(proxyPort)) { - long finish = System.currentTimeMillis(); + long start = System.currentTimeMillis(); + org.apache.http.HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", writeWebServerPort), request); + long finish = System.currentTimeMillis(); - assertThat("Expected throttled read to complete in approximately " + msToReadThrottled + "ms" + " but took " + (finish - start) + "ms", - (double)(finish - start), both(greaterThan(msToReadThrottled * (1 - ALLOWABLE_VARIATION))) - .and(lessThan(msToReadThrottled * (1 + ALLOWABLE_VARIATION)))); + assertThat(EntityUtils.toString(response.getEntity())) + .isEqualTo("Received " + largeData.length + " bytes\n"); - proxyServer.abort(); + assertThat((double) (finish - start)) + .as( + "Expected throttled write to complete in approximately %s ms but took %s ms", + msToWriteThrottled, finish - start) + .isGreaterThan(msToWriteThrottled * (1 - ALLOWABLE_VARIATION)) + .isLessThan(msToWriteThrottled * (1 + ALLOWABLE_VARIATION)); } + } - @Test - public void testUnthrottledRead() throws Exception { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); - - int proxyPort = proxyServer.getListenAddress().getPort(); + @Test + public void testUnthrottledWrite() throws Exception { + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - final HttpGet request = createHttpGet(); + int proxyPort = proxyServer.getListenAddress().getPort(); - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyPort); + final HttpPost request = createHttpPost(); - long start = System.currentTimeMillis(); - final org.apache.http.HttpResponse response = httpClient.execute( - new HttpHost("127.0.0.1", - readWebServerPort), request); + try (CloseableHttpClient httpClient = createProxiedHttpClient(proxyPort)) { - byte[] readContent = new byte[100000]; - int bytesRead = 0; - while (bytesRead < largeData.length) { - int read = response.getEntity().getContent().read(readContent, bytesRead, largeData.length - bytesRead); - bytesRead += read; - } + long start = System.currentTimeMillis(); + final org.apache.http.HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", writeWebServerPort), request); + long finish = System.currentTimeMillis(); - long finish = System.currentTimeMillis(); + assertThat(EntityUtils.toString(response.getEntity())) + .isEqualTo("Received " + largeData.length + " bytes\n"); - assertThat("Unthrottled read took " + (finish - start) + "ms, but expected to complete in " + UNTRHOTTLED_REQUEST_TIME_MS + "ms", - finish - start, lessThan((long)UNTRHOTTLED_REQUEST_TIME_MS)); - - proxyServer.abort(); + assertThat(finish - start) + .as( + "Unthrottled write took %s ms, but expected to complete in %s ms", + finish - start, UNTHROTTLED_REQUEST_TIME_MS) + .isLessThan(UNTHROTTLED_REQUEST_TIME_MS); } - - @Test - public void testChangeThrottling() throws Exception { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withThrottling(THROTTLED_READ_BYTES_PER_SECOND, 0) - .start(); - - int proxyPort = proxyServer.getListenAddress().getPort(); - - final HttpGet request = createHttpGet(); - - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyPort); - - long firstStart = System.currentTimeMillis(); - org.apache.http.HttpResponse response = httpClient.execute( - new HttpHost("127.0.0.1", - readWebServerPort), request); - byte[] readContent = new byte[100000]; - - int bytesRead = 0; - while (bytesRead < largeData.length) { - int read = response.getEntity().getContent().read(readContent, bytesRead, largeData.length - bytesRead); - bytesRead += read; - } - - long firstFinish = System.currentTimeMillis(); - - HttpClientUtils.closeQuietly(response); - - proxyServer.setThrottle(THROTTLED_READ_BYTES_PER_SECOND * 2, 0); - - long secondStart = System.currentTimeMillis(); - response = httpClient.execute( - new HttpHost("127.0.0.1", - readWebServerPort), request); - readContent = new byte[100000]; - - bytesRead = 0; - while (bytesRead < largeData.length) { - int read = response.getEntity().getContent().read(readContent, bytesRead, largeData.length - bytesRead); - bytesRead += read; - } - - long secondFinish = System.currentTimeMillis(); - - HttpClientUtils.closeQuietly(response); - - assertThat("Expected second read to take approximately half as long as first throttled read. First read took " + (firstFinish - firstStart) + "ms" + " but second read took " + (secondFinish - secondStart) + "ms", - (double)(firstFinish - firstStart) / 2, both(greaterThan((secondFinish - secondStart) * (1 - ALLOWABLE_VARIATION))) - .and(lessThan((secondFinish - secondStart) * (1 + ALLOWABLE_VARIATION)))); - - proxyServer.abort(); + } + + @Test + public void testThrottledRead() throws Exception { + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withThrottling(THROTTLED_READ_BYTES_PER_SECOND, 0) + .start(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + + final HttpGet request = createHttpGet(); + + try (CloseableHttpClient httpClient = createProxiedHttpClient(proxyPort)) { + + long start = System.currentTimeMillis(); + final org.apache.http.HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request); + byte[] readContent = new byte[LARGE_DATA_SIZE]; + + int bytesRead = 0; + while (bytesRead < largeData.length) { + int read = + response + .getEntity() + .getContent() + .read(readContent, bytesRead, largeData.length - bytesRead); + bytesRead += read; + } + + long finish = System.currentTimeMillis(); + + assertThat(bytesRead) + .as("Expected to read %s bytes but actually read %s bytes", LARGE_DATA_SIZE, bytesRead) + .isEqualTo(LARGE_DATA_SIZE); + + assertThat((double) (finish - start)) + .as( + "Expected throttled read to complete in approximately %s ms but took %s ms", + msToReadThrottled, finish - start) + .isGreaterThan(msToReadThrottled * (1 - ALLOWABLE_VARIATION)) + .isLessThan(msToReadThrottled * (1 + ALLOWABLE_VARIATION)); } + } - @Test - public void testDisableThrottling() throws Exception { - HttpProxyServer proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withThrottling(THROTTLED_READ_BYTES_PER_SECOND, 0) - .start(); - - int proxyPort = proxyServer.getListenAddress().getPort(); - - final HttpGet request = createHttpGet(); + @Test + public void testUnthrottledRead() throws Exception { + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyPort); + int proxyPort = proxyServer.getListenAddress().getPort(); - long firstStart = System.currentTimeMillis(); - org.apache.http.HttpResponse response = httpClient.execute( - new HttpHost("127.0.0.1", - readWebServerPort), request); - byte[] readContent = new byte[100000]; + final HttpGet request = createHttpGet(); - int bytesRead = 0; - while (bytesRead < largeData.length) { - int read = response.getEntity().getContent().read(readContent, bytesRead, largeData.length - bytesRead); - bytesRead += read; - } - - long firstFinish = System.currentTimeMillis(); - - HttpClientUtils.closeQuietly(response); - - proxyServer.setThrottle(0, 0); - - long secondStart = System.currentTimeMillis(); - response = httpClient.execute( - new HttpHost("127.0.0.1", - readWebServerPort), request); - readContent = new byte[100000]; + try (CloseableHttpClient httpClient = createProxiedHttpClient(proxyPort)) { - bytesRead = 0; - while (bytesRead < largeData.length) { - int read = response.getEntity().getContent().read(readContent, bytesRead, largeData.length - bytesRead); - bytesRead += read; - } - - long secondFinish = System.currentTimeMillis(); + long start = System.currentTimeMillis(); + final org.apache.http.HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request); - HttpClientUtils.closeQuietly(response); + byte[] readContent = new byte[LARGE_DATA_SIZE]; + int bytesRead = 0; + while (bytesRead < largeData.length) { + int read = + response + .getEntity() + .getContent() + .read(readContent, bytesRead, largeData.length - bytesRead); + bytesRead += read; + } - assertThat("Expected second read to complete within " + UNTRHOTTLED_REQUEST_TIME_MS + "ms, without throttling. First read took " - + (firstFinish - firstStart) + "ms" + ". Second read took " + (secondFinish - secondStart) + "ms", - secondFinish - secondStart, lessThan((long) UNTRHOTTLED_REQUEST_TIME_MS)); + long finish = System.currentTimeMillis(); - proxyServer.abort(); + assertThat(bytesRead) + .as("Expected to read %s bytes but actually read %s bytes", LARGE_DATA_SIZE, bytesRead) + .isEqualTo(LARGE_DATA_SIZE); + assertThat(finish - start) + .as( + "Unthrottled read took %s ms, but expected to complete in %s ms", + finish - start, UNTHROTTLED_REQUEST_TIME_MS) + .isLessThan(UNTHROTTLED_REQUEST_TIME_MS); } - - private HttpGet createHttpGet() { - final HttpGet request = new HttpGet("/"); - request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); - return request; + } + + @Test + @Timeout(30) + public void testChangeThrottling() throws Exception { + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withThrottling(THROTTLED_READ_BYTES_PER_SECOND, 0) + .start(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + + final HttpGet request = createHttpGet(); + + try (CloseableHttpClient httpClient = createProxiedHttpClient(proxyPort)) { + + long firstStart = System.currentTimeMillis(); + org.apache.http.HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request); + byte[] readContent = new byte[LARGE_DATA_SIZE]; + + int bytesRead = 0; + while (bytesRead < largeData.length) { + int read = + response + .getEntity() + .getContent() + .read(readContent, bytesRead, largeData.length - bytesRead); + bytesRead += read; + } + + long firstFinish = System.currentTimeMillis(); + + assertThat(bytesRead) + .as("Expected to read %s bytes but actually read %s bytes", LARGE_DATA_SIZE, bytesRead) + .isEqualTo(LARGE_DATA_SIZE); + + closeQuietly(response); + + proxyServer.setThrottle(THROTTLED_READ_BYTES_PER_SECOND * 2, 0); + Thread.sleep(1000); // necessary for the traffic shaping to reset + + long secondStart = System.currentTimeMillis(); + response = httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request); + readContent = new byte[LARGE_DATA_SIZE]; + + bytesRead = 0; + while (bytesRead < largeData.length) { + int read = + response + .getEntity() + .getContent() + .read(readContent, bytesRead, largeData.length - bytesRead); + bytesRead += read; + } + + long secondFinish = System.currentTimeMillis(); + + assertThat(bytesRead) + .as("Expected to read %s bytes but actually read %s bytes", LARGE_DATA_SIZE, bytesRead) + .isEqualTo(LARGE_DATA_SIZE); + + closeQuietly(response); + + assertThat((double) (firstFinish - firstStart) / 2) + .as( + "Expected second read to take approximately half as long as first throttled read. First read took %s ms but second read took %s ms", + firstFinish - firstStart, secondFinish - secondStart) + .isGreaterThan((secondFinish - secondStart) * (1 - ALLOWABLE_VARIATION)) + .isLessThan((secondFinish - secondStart) * (1 + ALLOWABLE_VARIATION)); } - - private HttpPost createHttpPost() { - final HttpPost request = new HttpPost("/"); - request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); - final ByteArrayEntity entity = new ByteArrayEntity(largeData); - entity.setChunked(true); - request.setEntity(entity); - return request; + } + + @Test + public void testDisableThrottling() throws Exception { + proxyServer = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withThrottling(THROTTLED_READ_BYTES_PER_SECOND, 0) + .start(); + + int proxyPort = proxyServer.getListenAddress().getPort(); + + final HttpGet request = createHttpGet(); + + try (CloseableHttpClient httpClient = createProxiedHttpClient(proxyPort)) { + + long firstStart = System.currentTimeMillis(); + org.apache.http.HttpResponse response = + httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request); + byte[] readContent = new byte[LARGE_DATA_SIZE]; + + int bytesRead = 0; + while (bytesRead < largeData.length) { + int read = + response + .getEntity() + .getContent() + .read(readContent, bytesRead, largeData.length - bytesRead); + bytesRead += read; + } + + long firstFinish = System.currentTimeMillis(); + + assertThat(bytesRead) + .as("Expected to read %s bytes but actually read %s bytes", LARGE_DATA_SIZE, bytesRead) + .isEqualTo(LARGE_DATA_SIZE); + + closeQuietly(response); + + proxyServer.setThrottle(0, 0); + Thread.sleep(1000); // necessary for the traffic shaping to reset + + long secondStart = System.currentTimeMillis(); + response = httpClient.execute(new HttpHost("127.0.0.1", readWebServerPort), request); + readContent = new byte[LARGE_DATA_SIZE]; + + bytesRead = 0; + while (bytesRead < largeData.length) { + int read = + response + .getEntity() + .getContent() + .read(readContent, bytesRead, largeData.length - bytesRead); + bytesRead += read; + } + + long secondFinish = System.currentTimeMillis(); + + assertThat(bytesRead) + .as("Expected to read %s bytes but actually read %s bytes", LARGE_DATA_SIZE, bytesRead) + .isEqualTo(LARGE_DATA_SIZE); + + closeQuietly(response); + + assertThat(secondFinish - secondStart) + .as( + "Expected second read to complete within %s ms, without throttling. First read took %s ms" + + ". Second read took %s ms.", + UNTHROTTLED_REQUEST_TIME_MS, firstFinish - firstStart, secondFinish - secondStart) + .isLessThan(UNTHROTTLED_REQUEST_TIME_MS); } + } + + private HttpGet createHttpGet() { + final HttpGet request = new HttpGet("/"); + request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); + return request; + } + + private HttpPost createHttpPost() { + final HttpPost request = new HttpPost("/"); + request.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); + final ByteArrayEntity entity = new ByteArrayEntity(largeData); + entity.setChunked(true); + request.setEntity(entity); + return request; + } } diff --git a/src/test/java/org/littleshoot/proxy/TimeoutTest.java b/src/test/java/org/littleshoot/proxy/TimeoutTest.java index 8e3e69b0..7379decb 100644 --- a/src/test/java/org/littleshoot/proxy/TimeoutTest.java +++ b/src/test/java/org/littleshoot/proxy/TimeoutTest.java @@ -1,133 +1,137 @@ package org.littleshoot.proxy; +import static com.github.tomakehurst.wiremock.client.WireMock.*; +import static com.github.tomakehurst.wiremock.core.WireMockConfiguration.options; +import static java.util.concurrent.TimeUnit.MILLISECONDS; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.createProxiedHttpClient; + +import com.github.tomakehurst.wiremock.WireMockServer; +import java.io.IOException; +import java.net.Socket; +import java.util.concurrent.TimeUnit; import org.apache.http.HttpResponse; import org.apache.http.client.methods.HttpGet; import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.util.EntityUtils; -import org.junit.After; -import org.junit.Before; -import org.junit.Test; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.test.EnableThreadDump; import org.littleshoot.proxy.test.SocketClientUtil; -import org.mockserver.integration.ClientAndServer; -import org.mockserver.matchers.Times; -import org.mockserver.model.Delay; - -import java.io.IOException; -import java.net.Socket; -import java.util.concurrent.TimeUnit; - -import static org.hamcrest.Matchers.lessThan; -import static org.junit.Assert.assertEquals; -import static org.junit.Assert.assertFalse; -import static org.junit.Assert.assertThat; -import static org.mockserver.model.HttpRequest.request; -import static org.mockserver.model.HttpResponse.response; - -public class TimeoutTest { - - private static final String UNUSED_URI_FOR_BAD_GATEWAY = "http://1.2.3.6:53540"; - - private ClientAndServer mockServer; - private int mockServerPort; - private HttpProxyServer proxyServer; - - @Before - public void setUp() { - mockServer = new ClientAndServer(0); - mockServerPort = mockServer.getLocalPort(); - } - - @After - public void tearDown() { - try { - if (mockServer != null) { - mockServer.stop(); - } - } finally { - if (proxyServer != null) { - proxyServer.abort(); - } - } +@EnableThreadDump +public final class TimeoutTest { + + private static final String UNUSED_URI_FOR_BAD_GATEWAY = "http://1.2.3.6:53540"; + + private WireMockServer mockServer; + private int mockServerPort; + + private HttpProxyServer proxyServer; + + @BeforeEach + void setUp() { + mockServer = new WireMockServer(options().dynamicPort()); + mockServer.start(); + mockServerPort = mockServer.port(); + } + + @AfterEach + void tearDown() { + try { + if (mockServer != null) { + mockServer.stop(); + } + } finally { + if (proxyServer != null) { + proxyServer.abort(); + } } - - @Test - public void testIdleConnectionTimeout() throws Exception { - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withIdleConnectionTimeout(1) - .start(); - - mockServer.when(request() - .withMethod("GET") - .withPath("/idleconnection"), - Times.exactly(1)) - .respond(response() - .withStatusCode(200) - .withDelay(new Delay(TimeUnit.SECONDS, 5)) - ); - - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyServer.getListenAddress().getPort()); - - long start = System.nanoTime(); - HttpGet get = new HttpGet("http://127.0.0.1:" + mockServerPort + "/idleconnection"); - long stop = System.nanoTime(); - - HttpResponse response = httpClient.execute(get); - EntityUtils.consumeQuietly(response.getEntity()); - - assertEquals("Expected to receive an HTTP 504 (Gateway Timeout) response after proxy did not receive a response within 1 second", 504, response.getStatusLine().getStatusCode()); - assertThat("Expected idle connection timeout to happen after approximately 1 second", - TimeUnit.MILLISECONDS.convert(stop - start, TimeUnit.NANOSECONDS), lessThan(2000L)); - } - - @Test - public void testConnectionTimeout() throws Exception { - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withConnectTimeout(1000) - .start(); - - CloseableHttpClient httpClient = TestUtils.createProxiedHttpClient(proxyServer.getListenAddress().getPort()); - - HttpGet get = new HttpGet(UNUSED_URI_FOR_BAD_GATEWAY); - - long start = System.nanoTime(); - HttpResponse response = httpClient.execute(get); - long stop = System.nanoTime(); - - EntityUtils.consumeQuietly(response.getEntity()); - - assertEquals("Expected to receive an HTTP 502 (Bad Gateway) response after proxy could not connect within 1 second", 502, response.getStatusLine().getStatusCode()); - assertThat("Expected connection timeout to happen after approximately 1 second", - TimeUnit.MILLISECONDS.convert(stop - start, TimeUnit.NANOSECONDS), lessThan(2000L)); + } + + @Test + public void testIdleConnectionTimeout() throws Exception { + proxyServer = + DefaultHttpProxyServer.bootstrap().withPort(0).withIdleConnectionTimeout(1).start(); + + mockServer.stubFor( + get(urlEqualTo("/idleconnection")) + .willReturn(aResponse().withStatus(200).withFixedDelay(5000))); + + try (CloseableHttpClient httpClient = + createProxiedHttpClient(proxyServer.getListenAddress().getPort())) { + long start = System.nanoTime(); + HttpGet get = new HttpGet("http://127.0.0.1:" + mockServerPort + "/idleconnection"); + long stop = System.nanoTime(); + + HttpResponse response = httpClient.execute(get); + EntityUtils.consumeQuietly(response.getEntity()); + + assertThat(response.getStatusLine().getStatusCode()) + .as( + "Expected to receive an HTTP 504 (Gateway Timeout) response after proxy did not receive a response within 1 second") + .isEqualTo(504); + assertThat(MILLISECONDS.convert(stop - start, TimeUnit.NANOSECONDS)) + .as("Expected idle connection timeout to happen after approximately 1 second") + .isLessThan(2000L); } + } - /** - * Verifies that when the client times out sending the initial request, the proxy still returns a Gateway Timeout. - */ - @Test - public void testClientIdleBeforeRequestReceived() throws IOException, InterruptedException { - this.proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withIdleConnectionTimeout(1) - .start(); - - // connect to the proxy and begin to transmit the request, but don't send the trailing \r\n that indicates the client has completely transmitted the request - String successfulGet = "GET http://localhost:" + mockServerPort + "/success HTTP/1.1"; + @Test + public void testConnectionTimeout() throws Exception { + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).withConnectTimeout(1000).start(); - // using the SocketClientUtil since we want the client to fail to send the entire GET request, which is not possible with - // Apache HTTP client and most other HTTP clients - Socket socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + try (CloseableHttpClient httpClient = + createProxiedHttpClient(proxyServer.getListenAddress().getPort())) { - SocketClientUtil.writeStringToSocket(successfulGet, socket); + HttpGet get = new HttpGet(UNUSED_URI_FOR_BAD_GATEWAY); - // wait a bit to allow the proxy server to respond - Thread.sleep(1500); + long start = System.nanoTime(); + HttpResponse response = httpClient.execute(get); + long stop = System.nanoTime(); - assertFalse("Client to proxy connection should be closed", SocketClientUtil.isSocketReadyToRead(socket)); + EntityUtils.consumeQuietly(response.getEntity()); - socket.close(); + assertThat(response.getStatusLine().getStatusCode()) + .as( + "Expected to receive an HTTP 502 (Bad Gateway) response after proxy could not connect within 1 second") + .isEqualTo(502); + assertThat(MILLISECONDS.convert(stop - start, TimeUnit.NANOSECONDS)) + .as("Expected connection timeout to happen after approximately 1 second") + .isLessThan(2000L); } + } + + /** + * Verifies that when the client times out sending the initial request, the proxy still returns a + * Gateway Timeout. + */ + @Test + public void testClientIdleBeforeRequestReceived() throws IOException, InterruptedException { + proxyServer = + DefaultHttpProxyServer.bootstrap().withPort(0).withIdleConnectionTimeout(1).start(); + + // connect to the proxy and begin to transmit the request, but don't send the + // trailing \r\n that indicates the client has completely transmitted the + // request + String successfulGet = "GET http://localhost:" + mockServerPort + "/success HTTP/1.1"; + + // using the SocketClientUtil since we want the client to fail to send the + // entire GET request, which is not possible with + // Apache HTTP client and most other HTTP clients + Socket socket = SocketClientUtil.getSocketToProxyServer(proxyServer); + + SocketClientUtil.writeStringToSocket(successfulGet, socket); + + // wait a bit to allow the proxy server to respond + Thread.sleep(1500); + + assertThat(SocketClientUtil.isSocketReadyToRead(socket)) + .as("Client to proxy connection should be closed") + .isFalse(); + + socket.close(); + } } diff --git a/src/test/java/org/littleshoot/proxy/UnencryptedTCPChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/UnencryptedTCPChainedProxyTest.java index 7641444a..65835e3f 100644 --- a/src/test/java/org/littleshoot/proxy/UnencryptedTCPChainedProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/UnencryptedTCPChainedProxyTest.java @@ -1,26 +1,8 @@ package org.littleshoot.proxy; -import static org.littleshoot.proxy.TransportProtocol.TCP; - -public class UnencryptedTCPChainedProxyTest extends BaseChainedProxyTest { - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(TCP); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.TCP; - } - - @Override - public boolean requiresEncryption() { - return false; - } - }; - } +public final class UnencryptedTCPChainedProxyTest extends BaseChainedProxyTest { + @Override + protected HttpProxyServerBootstrap upstreamProxy() { + return super.upstreamProxy(); + } } diff --git a/src/test/java/org/littleshoot/proxy/UnencryptedUDTChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/UnencryptedUDTChainedProxyTest.java deleted file mode 100644 index 33b574ae..00000000 --- a/src/test/java/org/littleshoot/proxy/UnencryptedUDTChainedProxyTest.java +++ /dev/null @@ -1,26 +0,0 @@ -package org.littleshoot.proxy; - -import static org.littleshoot.proxy.TransportProtocol.UDT; - -public class UnencryptedUDTChainedProxyTest extends BaseChainedProxyTest { - @Override - protected HttpProxyServerBootstrap upstreamProxy() { - return super.upstreamProxy() - .withTransportProtocol(UDT); - } - - @Override - protected ChainedProxy newChainedProxy() { - return new BaseChainedProxy() { - @Override - public TransportProtocol getTransportProtocol() { - return TransportProtocol.UDT; - } - - @Override - public boolean requiresEncryption() { - return false; - } - }; - } -} diff --git a/src/test/java/org/littleshoot/proxy/UsernamePasswordAuthenticatingProxyTest.java b/src/test/java/org/littleshoot/proxy/UsernamePasswordAuthenticatingProxyTest.java index f3c704ae..d7ec9a2b 100644 --- a/src/test/java/org/littleshoot/proxy/UsernamePasswordAuthenticatingProxyTest.java +++ b/src/test/java/org/littleshoot/proxy/UsernamePasswordAuthenticatingProxyTest.java @@ -1,40 +1,35 @@ package org.littleshoot.proxy; -/** - * Tests a single proxy that requires username/password authentication. - */ +/** Tests a single proxy that requires username/password authentication. */ public class UsernamePasswordAuthenticatingProxyTest extends BaseProxyTest - implements ProxyAuthenticator { - @Override - protected void setUp() { - this.proxyServer = bootstrapProxy() - .withPort(0) - .withProxyAuthenticator(this) - .start(); - } + implements ProxyAuthenticator { + @Override + protected void setUp() { + proxyServer = bootstrapProxy().withPort(0).withProxyAuthenticator(this).start(); + } - @Override - protected String getUsername() { - return "user1"; - } + @Override + protected String getUsername() { + return "user1"; + } - @Override - protected String getPassword() { - return "user2"; - } + @Override + protected String getPassword() { + return "user2"; + } - @Override - public boolean authenticate(String userName, String password) { - return getUsername().equals(userName) && getPassword().equals(password); - } + @Override + public boolean authenticate(String userName, String password) { + return getUsername().equals(userName) && getPassword().equals(password); + } - @Override - protected boolean isAuthenticating() { - return true; - } + @Override + protected boolean isAuthenticating() { + return true; + } - @Override - public String getRealm() { - return null; - } + @Override + public String getRealm() { + return null; + } } diff --git a/src/test/java/org/littleshoot/proxy/VariableSpeedClientServerTest.java b/src/test/java/org/littleshoot/proxy/VariableSpeedClientServerTest.java index 320031ee..6cba5a5b 100644 --- a/src/test/java/org/littleshoot/proxy/VariableSpeedClientServerTest.java +++ b/src/test/java/org/littleshoot/proxy/VariableSpeedClientServerTest.java @@ -1,157 +1,166 @@ package org.littleshoot.proxy; +import static java.nio.charset.StandardCharsets.UTF_8; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.TestUtils.createProxiedHttpClient; + +import java.io.*; +import java.net.ServerSocket; +import java.net.Socket; +import java.util.Arrays; import org.apache.http.HttpEntity; import org.apache.http.HttpResponse; import org.apache.http.client.methods.HttpPost; import org.apache.http.entity.InputStreamEntity; import org.apache.http.impl.client.CloseableHttpClient; import org.apache.http.util.EntityUtils; -import org.junit.Ignore; -import org.junit.Test; +import org.junit.jupiter.api.Disabled; +import org.junit.jupiter.api.Test; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; - -import java.io.BufferedReader; -import java.io.InputStream; -import java.io.InputStreamReader; -import java.io.OutputStream; -import java.net.ServerSocket; -import java.net.Socket; -import java.nio.charset.Charset; -import java.util.Arrays; - -import static org.junit.Assert.assertEquals; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; /** * Tests cases where either the client or the server is slower than the other. - * - * Ignored because this doesn't quite trigger OOME for some reason. It also - * takes too long to include in normal tests. + * + *

Ignored because this doesn't quite trigger OOME for some reason. It also takes too long to + * include in normal tests. */ -@Ignore -public class VariableSpeedClientServerTest { - - private static final int PORT = TestUtils.randomPort(); - private static final int PROXY_PORT = TestUtils.randomPort(); - private static final int CONTENT_LENGTH = 1000000000; - - @Test - public void testServerFaster() throws Exception { - doTest(PORT, PROXY_PORT, false); - } - - @Test - public void testServerSlower() throws Exception { - doTest(PORT, PROXY_PORT, true); +@Disabled +public final class VariableSpeedClientServerTest { + private static final Logger log = LoggerFactory.getLogger(VariableSpeedClientServerTest.class); + + private static final int PORT = TestUtils.randomPort(); + private static final int PROXY_PORT = TestUtils.randomPort(); + private static final int CONTENT_LENGTH = 1000000000; + + @Test + public void testServerFaster() throws Exception { + doTest(PORT, PROXY_PORT, false); + } + + @Test + public void testServerSlower() throws Exception { + doTest(PORT, PROXY_PORT, true); + } + + private void doTest(int port, int proxyPort, boolean slowServer) throws Exception { + startServer(port, slowServer); + Thread.yield(); + DefaultHttpProxyServer.bootstrap().withPort(proxyPort).start(); + Thread.yield(); + Thread.sleep(400); + try (CloseableHttpClient client = createProxiedHttpClient(proxyPort)) { + + log.info("------------------ Memory Usage At Beginning ------------------"); + TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); + + final HttpPost post = createHttpPost("http://127.0.0.1:" + port + "/"); + final HttpResponse response = client.execute(post); + + final HttpEntity entity = response.getEntity(); + final long cl = entity.getContentLength(); + assertThat(cl).isEqualTo(CONTENT_LENGTH); + + int bytesRead = 0; + try (InputStream content = + slowServer + ? new ThrottledInputStream(entity.getContent(), 10 * 1000) + : entity.getContent()) { + final byte[] input = new byte[100000]; + int read = content.read(input); + + while (read != -1) { + bytesRead += read; + read = content.read(input); + } + } + assertThat(bytesRead).isEqualTo(CONTENT_LENGTH); + // final String body = IOUtils.toString(entity.getContent()); + EntityUtils.consume(entity); + log.info("------------------ Memory Usage At Beginning ------------------"); + TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); } - - private void doTest(int port, int proxyPort, boolean slowServer) - throws Exception { - startServer(port, slowServer); - Thread.yield(); - DefaultHttpProxyServer.bootstrap().withPort(proxyPort).start(); - Thread.yield(); - Thread.sleep(400); - final CloseableHttpClient client = TestUtils.createProxiedHttpClient(proxyPort); - - System.out - .println("------------------ Memory Usage At Beginning ------------------"); - TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); - - final String endpoint = "http://127.0.0.1:" + port + "/"; - final HttpPost post = new HttpPost(endpoint); - post.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); - post.setEntity(new InputStreamEntity(new InputStream() { - private int remaining = CONTENT_LENGTH; - - @Override - public int read() { + } + + private static HttpPost createHttpPost(String endpoint) { + final HttpPost post = new HttpPost(endpoint); + post.setConfig(TestUtils.REQUEST_TIMEOUT_CONFIG); + post.setEntity( + new InputStreamEntity( + new InputStream() { + private int remaining = CONTENT_LENGTH; + + @Override + public int read() { if (remaining > 0) { - remaining -= 1; - return 77; + remaining -= 1; + return 77; } else { - return 0; + return 0; } - } + } - @Override - public int available() { + @Override + public int available() { return remaining; - } - }, CONTENT_LENGTH)); - final HttpResponse response = client.execute(post); - - final HttpEntity entity = response.getEntity(); - final long cl = entity.getContentLength(); - assertEquals(CONTENT_LENGTH, cl); - - int bytesRead = 0; - try (InputStream content = slowServer ? new ThrottledInputStream(entity.getContent(), 10 * 1000) : entity.getContent()) { - final byte[] input = new byte[100000]; - int read = content.read(input); - - while (read != -1) { - bytesRead += read; - read = content.read(input); - } - } - assertEquals(CONTENT_LENGTH, bytesRead); - // final String body = IOUtils.toString(entity.getContent()); - EntityUtils.consume(entity); - System.out - .println("------------------ Memory Usage At Beginning ------------------"); - TestUtils.getOpenFileDescriptorsAndPrintMemoryUsage(); - } - - private void startServer(final int port, final boolean slowReader) { - final Thread t = new Thread(() -> { - try { + } + }, + CONTENT_LENGTH)); + return post; + } + + private void startServer(final int port, final boolean slowReader) { + final Thread t = + new Thread( + () -> { + try { startServerOnThread(port, slowReader); - } catch (Exception e) { - e.printStackTrace(); - } - }, "Test-Server-Thread"); - t.setDaemon(true); - t.start(); - } - - private void startServerOnThread(int port, boolean slowReader) - throws Exception { - try (ServerSocket server = new ServerSocket(port)) { - server.setSoTimeout(100000); - final Socket sock = server.accept(); - InputStream is = sock.getInputStream(); - if (slowReader) { - is = new ThrottledInputStream(is, 10 * 1000); - } - BufferedReader br = new BufferedReader(new InputStreamReader(is)); - while (br.read() != 0) { - } - final OutputStream os = sock.getOutputStream(); - final String responseHeaders = - "HTTP/1.1 200 OK\r\n" + - "Date: Sun, 20 Jan 2013 00:16:23 GMT\r\n" + - "Expires: -1\r\n" + - "Cache-Control: private, max-age=0\r\n" + - "Content-Type: text/html; charset=ISO-8859-1\r\n" + - "Server: gws\r\n" + - "Content-Length: " + CONTENT_LENGTH + "\r\n\r\n"; // 10 - // gigs - // or - // so. - - os.write(responseHeaders.getBytes(Charset.forName("UTF-8"))); - - int bufferSize = 100000; - final byte[] bytes = new byte[bufferSize]; - Arrays.fill(bytes, (byte) 77); - int remainingBytes = CONTENT_LENGTH; - - while (remainingBytes > 0) { - int numberOfBytesToWrite = Math.min(remainingBytes, bufferSize); - os.write(bytes, 0, numberOfBytesToWrite); - remainingBytes -= numberOfBytesToWrite; - } - os.close(); + } catch (IOException e) { + log.error( + "Failed to start server on port {} (slowReader: {})", port, slowReader, e); + } + }, + "Test-Server-Thread"); + t.setDaemon(true); + t.start(); + } + + private void startServerOnThread(int port, boolean slowReader) throws IOException { + try (ServerSocket server = new ServerSocket(port)) { + server.setSoTimeout(100000); + final Socket sock = server.accept(); + InputStream is = sock.getInputStream(); + if (slowReader) { + is = new ThrottledInputStream(is, 10 * 1000); + } + BufferedReader br = new BufferedReader(new InputStreamReader(is)); + while (br.read() != 0) {} + try (final OutputStream os = sock.getOutputStream()) { + final String responseHeaders = + "HTTP/1.1 200 OK\r\n" + + "Date: Sun, 20 Jan 2013 00:16:23 GMT\r\n" + + "Expires: -1\r\n" + + "Cache-Control: private, max-age=0\r\n" + + "Content-Type: text/html; charset=ISO-8859-1\r\n" + + "Server: gws\r\n" + + "Content-Length: " + + CONTENT_LENGTH + + "\r\n\r\n"; // ~10 gigs + + os.write(responseHeaders.getBytes(UTF_8)); + + int bufferSize = 100000; + final byte[] bytes = new byte[bufferSize]; + Arrays.fill(bytes, (byte) 77); + int remainingBytes = CONTENT_LENGTH; + + while (remainingBytes > 0) { + int numberOfBytesToWrite = Math.min(remainingBytes, bufferSize); + os.write(bytes, 0, numberOfBytesToWrite); + remainingBytes -= numberOfBytesToWrite; } + } } + } } diff --git a/src/test/java/org/littleshoot/proxy/extras/ActivityLoggerTest.java b/src/test/java/org/littleshoot/proxy/extras/ActivityLoggerTest.java new file mode 100644 index 00000000..de24d00b --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/extras/ActivityLoggerTest.java @@ -0,0 +1,205 @@ +package org.littleshoot.proxy.extras; + +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +import io.netty.handler.codec.http.*; +import java.net.InetAddress; +import java.net.InetSocketAddress; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.FlowContext; +import org.littleshoot.proxy.FullFlowContext; +import org.mockito.Mock; +import org.mockito.MockitoAnnotations; + +class ActivityLoggerTest { + + @Mock private FlowContext flowContext; + @Mock private FullFlowContext fullFlowContext; + @Mock private HttpRequest request; + @Mock private HttpResponse response; + @Mock private HttpHeaders requestHeaders; + @Mock private HttpHeaders responseHeaders; + + @BeforeEach + void setUp() { + MockitoAnnotations.openMocks(this); + when(request.headers()).thenReturn(requestHeaders); + when(response.headers()).thenReturn(responseHeaders); + } + + @Test + void testClfFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.CLF); + setupMocks(); + + tracker.requestReceivedFromClient(flowContext, request); + tracker.responseSentToClient(flowContext, response); + + System.out.println("CLF Log: " + tracker.lastLogMessage); + // Expecting: 127.0.0.1 - - [Date] "GET /test HTTP/1.1" 200 100 + assertTrue(tracker.lastLogMessage.contains("127.0.0.1 - - [")); + assertTrue(tracker.lastLogMessage.contains("] \"GET /test HTTP/1.1\" 200 100")); + } + + @Test + void testJsonFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.JSON); + setupMocks(); + + tracker.requestReceivedFromClient(flowContext, request); + tracker.responseSentToClient(flowContext, response); + + System.out.println("JSON Log: " + tracker.lastLogMessage); + assertTrue(tracker.lastLogMessage.startsWith("{")); + assertTrue(tracker.lastLogMessage.contains("\"client_ip\":\"127.0.0.1\"")); + assertTrue(tracker.lastLogMessage.contains("\"method\":\"GET\"")); + assertTrue(tracker.lastLogMessage.contains("\"uri\":\"/test\"")); + assertTrue(tracker.lastLogMessage.contains("\"status\":200")); + assertTrue(tracker.lastLogMessage.contains("\"bytes\":100")); + } + + @Test + void testElfFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.ELF); + setupMocks(); + when(requestHeaders.get("Referer")).thenReturn("http://referrer.com"); + when(requestHeaders.get("User-Agent")).thenReturn("Mozilla/5.0"); + + tracker.requestReceivedFromClient(flowContext, request); + tracker.responseSentToClient(flowContext, response); + + System.out.println("ELF Log: " + tracker.lastLogMessage); + // host ident authuser [date] "request" status bytes "referer" "user-agent" + // 127.0.0.1 - - [Date] "GET /test HTTP/1.1" 200 100 "http://referrer.com" + // "Mozilla/5.0" + assertTrue(tracker.lastLogMessage.startsWith("127.0.0.1 - - [")); + assertTrue( + tracker.lastLogMessage.contains( + "] \"GET /test HTTP/1.1\" 200 100 \"http://referrer.com\" \"Mozilla/5.0\"")); + } + + @Test + void testW3cFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.W3C); + setupMocks(); + when(requestHeaders.get("User-Agent")).thenReturn("Mozilla/5.0"); + + tracker.requestReceivedFromClient(flowContext, request); + tracker.responseSentToClient(flowContext, response); + + System.out.println("W3C Log: " + tracker.lastLogMessage); + // date time c-ip cs-method cs-uri-stem sc-status sc-bytes cs(User-Agent) + // YYYY-MM-DD HH:MM:SS 127.0.0.1 GET /test 200 100 "Mozilla/5.0" + assertTrue(tracker.lastLogMessage.contains(" 127.0.0.1 GET /test 200 100 \"Mozilla/5.0\"")); + } + + @Test + void testLtsvFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.LTSV); + setupMocksWithDelay(); + + tracker.requestReceivedFromClient(flowContext, request); + // Simulate delay + try { + Thread.sleep(10); + } catch (InterruptedException ignored) { + } + tracker.responseSentToClient(flowContext, response); + + System.out.println("LTSV Log: " + tracker.lastLogMessage); + // time:... host:127.0.0.1 method:GET uri:/test status:200 size:100 duration:>=0 + // ua:Mozilla/5.0 + assertTrue( + tracker.lastLogMessage.contains( + "host:127.0.0.1\tmethod:GET\turi:/test\tstatus:200\tsize:100\tduration:")); + } + + @Test + void testCsvFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.CSV); + setupMocksWithDelay(); + + tracker.requestReceivedFromClient(flowContext, request); + tracker.responseSentToClient(flowContext, response); + + System.out.println("CSV Log: " + tracker.lastLogMessage); + // "timestamp","127.0.0.1","GET","/test",200,100,duration,"Mozilla/5.0" + assertTrue(tracker.lastLogMessage.contains("\",\"127.0.0.1\",\"GET\",\"/test\",200,100,")); + assertTrue(tracker.lastLogMessage.endsWith(",\"Mozilla/5.0\"")); + } + + @Test + void testHaproxyFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.HAPROXY); + setupMocksWithDelay(); + + tracker.requestReceivedFromClient(flowContext, request); + tracker.responseSentToClient(flowContext, response); + + System.out.println("HAProxy Log: " + tracker.lastLogMessage); + // 127.0.0.1 [date] "GET /test HTTP/1.1" 200 100 duration + assertTrue(tracker.lastLogMessage.startsWith("127.0.0.1 [")); + assertTrue(tracker.lastLogMessage.contains("] \"GET /test HTTP/1.1\" 200 100 ")); + } + + @Test + void testSquidFormat() { + TestableActivityLogger tracker = new TestableActivityLogger(LogFormat.SQUID); + setupMocks(); + + tracker.requestReceivedFromClient(flowContext, request); + tracker.responseSentToClient(flowContext, response); + + System.out.println("Squid Log: " + tracker.lastLogMessage); + // time elapsed remotehost code/status bytes method URL rfc931 + // peerstatus/peerhost type + // Check that elapsed time is present (we can't check exact value easily but + // check structure) + // 1234567890.123 0 127.0.0.1 ... + // We now expect something >= 0, not necessarily hardcoded 0. + // Regex: timestamp space duration space ip ... + assertTrue( + tracker.lastLogMessage.matches( + ".*\\d+ \\d+ 127\\.0\\.0\\.1 TCP_MISS/200 100 GET /test - DIRECT/- -.*")); + } + + private static class TestableActivityLogger extends ActivityLogger { + String lastLogMessage; + + public TestableActivityLogger(LogFormat logFormat) { + super(logFormat); + } + + @Override + protected void log(String message) { + this.lastLogMessage = message; + } + } + + private void setupMocks() { + setupMocksCommon(); + } + + private void setupMocksWithDelay() { + setupMocksCommon(); + when(requestHeaders.get("User-Agent")).thenReturn("Mozilla/5.0"); + } + + private void setupMocksCommon() { + InetSocketAddress clientAddr = mock(InetSocketAddress.class); + InetAddress inetAddr = mock(InetAddress.class); + when(flowContext.getClientAddress()).thenReturn(clientAddr); + when(clientAddr.getAddress()).thenReturn(inetAddr); + when(inetAddr.getHostAddress()).thenReturn("127.0.0.1"); + + when(request.method()).thenReturn(HttpMethod.GET); + when(request.uri()).thenReturn("/test"); + when(request.protocolVersion()).thenReturn(HttpVersion.HTTP_1_1); + + when(response.status()).thenReturn(HttpResponseStatus.OK); + when(responseHeaders.get("Content-Length")).thenReturn("100"); + } +} diff --git a/src/test/java/org/littleshoot/proxy/extras/SelfSignedMitmManagerTest.java b/src/test/java/org/littleshoot/proxy/extras/SelfSignedMitmManagerTest.java index 20b5b7df..3044c71e 100644 --- a/src/test/java/org/littleshoot/proxy/extras/SelfSignedMitmManagerTest.java +++ b/src/test/java/org/littleshoot/proxy/extras/SelfSignedMitmManagerTest.java @@ -1,45 +1,44 @@ package org.littleshoot.proxy.extras; -import io.netty.handler.codec.http.HttpRequest; -import org.junit.Test; +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.*; +import io.netty.handler.codec.http.HttpRequest; import javax.net.ssl.SSLEngine; import javax.net.ssl.SSLSession; +import org.junit.jupiter.api.Test; -import static org.junit.Assert.assertEquals; -import static org.mockito.Mockito.*; - -public class SelfSignedMitmManagerTest { +public final class SelfSignedMitmManagerTest { - @Test - public void testServerSslEnginePeerAndPort() { - String peer = "localhost"; - int port = 8090; - SelfSignedSslEngineSource source = mock(SelfSignedSslEngineSource.class); - SelfSignedMitmManager manager = new SelfSignedMitmManager(source); - SSLEngine engine = mock(SSLEngine.class); - when(source.newSslEngine(peer, port)).thenReturn(engine); - assertEquals(engine, manager.serverSslEngine(peer, port)); - } + @Test + public void testServerSslEnginePeerAndPort() { + String peer = "localhost"; + int port = 8090; + SelfSignedSslEngineSource source = mock(); + SelfSignedMitmManager manager = new SelfSignedMitmManager(source); + SSLEngine engine = mock(); + when(source.newSslEngine(peer, port)).thenReturn(engine); + assertThat(manager.serverSslEngine(peer, port)).isEqualTo(engine); + } - @Test - public void testServerSslEngine() { - SelfSignedSslEngineSource source = mock(SelfSignedSslEngineSource.class); - SelfSignedMitmManager manager = new SelfSignedMitmManager(source); - SSLEngine engine = mock(SSLEngine.class); - when(source.newSslEngine()).thenReturn(engine); - assertEquals(engine, manager.serverSslEngine()); - } + @Test + public void testServerSslEngine() { + SelfSignedSslEngineSource source = mock(); + SelfSignedMitmManager manager = new SelfSignedMitmManager(source); + SSLEngine engine = mock(); + when(source.newSslEngine()).thenReturn(engine); + assertThat(manager.serverSslEngine()).isEqualTo(engine); + } - @Test - public void testClientSslEngineFor() { - HttpRequest request = mock(HttpRequest.class); - SSLSession session = mock(SSLSession.class); - SelfSignedSslEngineSource source = mock(SelfSignedSslEngineSource.class); - SelfSignedMitmManager manager = new SelfSignedMitmManager(source); - SSLEngine engine = mock(SSLEngine.class); - when(source.newSslEngine()).thenReturn(engine); - assertEquals(engine, manager.clientSslEngineFor(request, session)); - verifyZeroInteractions(request, session); - } + @Test + public void testClientSslEngineFor() { + HttpRequest request = mock(); + SSLSession session = mock(); + SelfSignedSslEngineSource source = mock(); + SelfSignedMitmManager manager = new SelfSignedMitmManager(source); + SSLEngine engine = mock(); + when(source.newSslEngine()).thenReturn(engine); + assertThat(manager.clientSslEngineFor(request, session)).isEqualTo(engine); + verifyNoMoreInteractions(request, session); + } } diff --git a/src/test/java/org/littleshoot/proxy/extras/SelfSignedSslEngineSourceTest.java b/src/test/java/org/littleshoot/proxy/extras/SelfSignedSslEngineSourceTest.java new file mode 100644 index 00000000..c5ffb035 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/extras/SelfSignedSslEngineSourceTest.java @@ -0,0 +1,406 @@ +package org.littleshoot.proxy.extras; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.io.File; +import java.io.FileInputStream; +import java.io.IOException; +import java.lang.reflect.Field; +import java.nio.ByteBuffer; +import java.security.KeyStore; +import java.security.KeyStoreException; +import java.util.Collections; +import java.util.List; +import java.util.concurrent.TimeUnit; +import javax.net.ssl.SSLContext; +import javax.net.ssl.SSLEngine; +import javax.net.ssl.SSLEngineResult; +import org.junit.jupiter.api.*; +import org.junit.jupiter.api.io.TempDir; + +@TestMethodOrder(MethodOrderer.OrderAnnotation.class) +final class SelfSignedSslEngineSourceTest { + + private static final String DEFAULT_KEYSTORE = "littleproxy_keystore.jks"; + private static final int BUFFER_CAPACITY = 32768; + + @TempDir File tempDir; + + @BeforeEach + void cleanDefaultKeystore() { + deleteIfExistsOrFail(new File(DEFAULT_KEYSTORE)); + deleteIfExistsOrFail(new File("littleproxy_cert")); + } + + private static boolean isKeytoolAvailable() { + ProcessBuilder pb = new ProcessBuilder("keytool", "-help"); + pb.redirectErrorStream(true); + try { + Process p = pb.start(); + boolean finished = p.waitFor(5, TimeUnit.SECONDS); + return finished && p.exitValue() == 0; + } catch (IOException | InterruptedException e) { + return false; + } + } + + private static void deleteIfExistsOrFail(File file) { + if (file.exists()) { + assertThat(file.delete()) + .as("Failed to delete stale test artifact: " + file.getAbsolutePath()) + .isTrue(); + } + } + + @BeforeAll + static void setUp() { + Assumptions.assumeTrue(isKeytoolAvailable(), "keytool is not installed, test ignored"); + } + + @AfterAll + static void cleanUp() { + deleteIfExistsOrFail(new File(DEFAULT_KEYSTORE)); + deleteIfExistsOrFail(new File("littleproxy_cert")); + } + + @Test + @Order(1) + void defaultConstructor() { + // The default constructor creates "littleproxy_keystore.jks" in the working directory + // Clean up any existing keystore first to ensure the test is deterministic + File keystore = new File(DEFAULT_KEYSTORE); + if (keystore.exists()) { + // Attempt to delete, but don't fail test if delete fails (file may be locked) + // The assertion on doesNotExist() below will catch if cleanup failed + keystore.delete(); + } + + // Create instance using default constructor (generates the keystore via keytool) + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(); + + // Verify the keystore file was created and is non-empty + assertThat(keystore).as("Keystore file should exist after default constructor").exists(); + assertThat(keystore.length()).as("Keystore file should not be empty").isGreaterThan(0); + + // Verify the instance is functional + assertThat(source.getSslContext()).isNotNull(); + assertThat(source.getSslContext().getProtocol()).isEqualTo("TLS"); + + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + assertThat(engine.getUseClientMode()).isFalse(); + } + + @Test + void constructorWithKeyStorePathReusesExistingKeystore() { + // Test that existing keystore is reused when file already exists + String keystorePath = new File(tempDir, "test_keystore_existing.jks").getAbsolutePath(); + File keystoreFile = new File(keystorePath); + + // First instantiation creates the keystore + SelfSignedSslEngineSource first = new SelfSignedSslEngineSource(keystorePath); + long originalSize = keystoreFile.length(); + assertThat(originalSize).as("Original keystore should have content").isGreaterThan(0); + + // Second instantiation should reuse existing keystore without modification + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(keystorePath); + + // Verify source is functional and keystore wasn't regenerated + assertThat(source.getSslContext()).isNotNull(); + assertThat(keystoreFile.length()) + .as("Keystore size should not change when reusing") + .isEqualTo(originalSize); + + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + } + + @Test + @Order(2) + void constructorWithTrustAllServersGeneratesDefaultKeystore() { + // Test trustAllServers=true with default keystore path + // Delete any existing default keystore to test generation + File defaultKeystore = new File(DEFAULT_KEYSTORE); + if (defaultKeystore.exists()) { + // Attempt to delete, but don't fail test if delete fails (file may be locked) + defaultKeystore.delete(); + } + + // Create source with trustAllServers enabled + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(true); + + // Verify source and SSLContext are initialized + assertThat(source.getSslContext()).isNotNull(); + assertThat(source.getSslContext().getProtocol()).isEqualTo("TLS"); + // Verify default keystore was created + assertThat(defaultKeystore).as("Default keystore should be created").exists(); + assertThat(defaultKeystore.length()) + .as("Default keystore should not be empty") + .isGreaterThan(0); + + // Verify SSLEngine is functional + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + assertThat(engine.getUseClientMode()).isFalse(); + } + + @Test + @Order(3) + void constructorWithTrustAllServersReusesExistingKeystore() { + // Test that existing default keystore is reused + File defaultKeystore = new File(DEFAULT_KEYSTORE); + // Ensure keystore exists before testing reuse + if (!defaultKeystore.exists()) { + new SelfSignedSslEngineSource(true); + } + long originalSize = defaultKeystore.length(); + assertThat(originalSize).as("Original keystore should have content").isGreaterThan(0); + + // Create another instance - should reuse existing keystore + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(true); + + // Verify keystore wasn't regenerated + assertThat(source.getSslContext()).isNotNull(); + assertThat(defaultKeystore.length()) + .as("Keystore size should not change when reusing") + .isEqualTo(originalSize); + + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + } + + @Test + void constructorWithFullParametersGeneratesNewKeystore() throws Exception { + // Test constructor with custom alias and password parameters + String keystorePath = new File(tempDir, "test_full.jks").getAbsolutePath(); + String alias = "test_alias"; + String password = "test_password"; + File keystoreFile = new File(keystorePath); + + SelfSignedSslEngineSource source = + new SelfSignedSslEngineSource(keystorePath, true, false, alias, password); + + assertThat(source.getSslContext()).isNotNull(); + assertThat(source.getSslContext().getProtocol()).isEqualTo("TLS"); + assertThat(keystoreFile).as("Keystore should be created").exists(); + assertThat(keystoreFile.length()).as("Keystore should not be empty").isGreaterThan(0); + + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + assertThat(engine.getUseClientMode()).isFalse(); + // this is an indirect check that password is used + assertThatThrownBy( + () -> { + KeyStore keyStore = KeyStore.getInstance("JKS"); + try (FileInputStream fis = new FileInputStream(keystoreFile)) { + keyStore.load(fis, "Be Your Own Lantern".toCharArray()); + } + }) + .as("Keystore should not be loadable with default password") + .isInstanceOf(IOException.class) + .hasMessageStartingWith("keystore password was incorrect"); + + KeyStore keyStore = KeyStore.getInstance("JKS"); + try (FileInputStream fis = new FileInputStream(keystoreFile)) { + keyStore.load(fis, password.toCharArray()); + } + assertThat(keyStore.containsAlias(alias)) + .as(() -> "Keystore should contain the custom alias, but received: " + aliasesOf(keyStore)) + .isTrue(); + } + + @Test + void constructorWithKeyStorePathGeneratesNewKeystore() { + // Test that a new keystore is generated when the specified file doesn't exist + String keystorePath = new File(tempDir, "test_keystore.jks").getAbsolutePath(); + File keystoreFile = new File(keystorePath); + // Verify keystore doesn't exist before instantiation + assertThat(keystoreFile).as("Keystore should not exist before test").doesNotExist(); + + // Create source - this should generate a new keystore + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(keystorePath); + + // Verify source is properly initialized + assertThat(source.getSslContext()).isNotNull(); + assertThat(source.getSslContext().getProtocol()).isEqualTo("TLS"); + // Verify keystore was created with content + assertThat(keystoreFile).as("Keystore should be created").exists(); + assertThat(keystoreFile.length()).as("Keystore should not be empty").isGreaterThan(0); + + // Verify SSLEngine is created in server mode + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + assertThat(engine.getUseClientMode()).isFalse(); + } + + @Test + void newSslEngineWithPeerInfoGeneratesNewKeystore() { + // Test newSslEngine(peerHost, peerPort) with custom peer information + String keystorePath = new File(tempDir, "peer_test.jks").getAbsolutePath(); + File keystoreFile = new File(keystorePath); + + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(keystorePath); + + SSLEngine engine = source.newSslEngine("example.com", -1); + + assertThat(engine).isNotNull(); + assertThat(engine.getPeerHost()).isEqualTo("example.com"); + assertThat(engine.getPeerPort()).isEqualTo(-1); + assertThat(keystoreFile).as("Keystore should be created").exists(); + assertThat(keystoreFile.length()).as("Keystore should not be empty").isGreaterThan(0); + } + + @Test + void getSslContextGeneratesNewKeystore() { + // Test getSslContext() returns valid SSLContext and generates keystore + String keystorePath = new File(tempDir, "context_test.jks").getAbsolutePath(); + File keystoreFile = new File(keystorePath); + + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(keystorePath); + + SSLContext context = source.getSslContext(); + + assertThat(context).isNotNull(); + assertThat(context.getProtocol()).isEqualTo("TLS"); + assertThat(keystoreFile).as("Keystore should be created").exists(); + assertThat(keystoreFile.length()).as("Keystore should not be empty").isGreaterThan(0); + } + + @Test + void trustAllServersOptionGeneratesNewKeystore() { + // Test trustAllServers=true option with sendCerts=true + String keystorePath = new File(tempDir, "trust_test.jks").getAbsolutePath(); + File keystoreFile = new File(keystorePath); + + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(keystorePath, true, true); + + assertThat(source.getSslContext()).isNotNull(); + assertThat(source.getSslContext().getProtocol()).isEqualTo("TLS"); + assertThat(keystoreFile).as("Keystore should be created").exists(); + assertThat(keystoreFile.length()).as("Keystore should not be empty").isGreaterThan(0); + + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + } + + @Test + void sendCertsOptionGeneratesNewKeystore() { + // Test sendCerts=false option (trustAllServers=false) + String keystorePath = new File(tempDir, "certs_test.jks").getAbsolutePath(); + File keystoreFile = new File(keystorePath); + + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(keystorePath, false, false); + + assertThat(getBooleanField(source, "trustAllServers")) + .as("trustAllServers should be initialized from constructor") + .isFalse(); + assertThat(getBooleanField(source, "sendCerts")) + .as("sendCerts should be initialized from constructor") + .isFalse(); + + assertThat(source.getSslContext()).isNotNull(); + assertThat(source.getSslContext().getProtocol()).isEqualTo("TLS"); + assertThat(keystoreFile).as("Keystore should be created").exists(); + assertThat(keystoreFile.length()).as("Keystore should not be empty").isGreaterThan(0); + + SSLEngine engine = source.newSslEngine(); + assertThat(engine).isNotNull(); + } + + @Test + void tlsHandshakeGeneratesNewKeystore() throws Exception { + // Test that generated keystore enables successful TLS handshake + String keystorePath = new File(tempDir, "handshake_test.jks").getAbsolutePath(); + File keystoreFile = new File(keystorePath); + + SelfSignedSslEngineSource source = new SelfSignedSslEngineSource(keystorePath); + + assertThat(keystoreFile).as("Keystore should be created").exists(); + assertThat(keystoreFile.length()).as("Keystore should not be empty").isGreaterThan(0); + + // Create server-side SSLEngine + SSLEngine serverEngine = source.newSslEngine(); + serverEngine.setUseClientMode(false); + + // Create client-side SSLEngine using same SSLContext + SSLContext clientContext = source.getSslContext(); + SSLEngine clientEngine = clientContext.createSSLEngine("localhost", -1); + clientEngine.setUseClientMode(true); + + // Perform TLS handshake in-memory + performHandshake(clientEngine, serverEngine); + + // Verify handshake completed successfully + assertThat(clientEngine.getSession()).isNotNull(); + assertThat(serverEngine.getSession()).isNotNull(); + } + + private List aliasesOf(KeyStore keyStore) { + try { + return Collections.list(keyStore.aliases()); + } catch (KeyStoreException e) { + throw new RuntimeException(e); + } + } + + // Performs an in-memory TLS handshake between two SSLEngines using ByteBuffers + // This validates that the keystore produces valid certificates without needing a real network + private void performHandshake(SSLEngine client, SSLEngine server) throws Exception { + ByteBuffer clientOut = ByteBuffer.allocate(BUFFER_CAPACITY); + ByteBuffer serverOut = ByteBuffer.allocate(BUFFER_CAPACITY); + clientOut.flip(); + serverOut.flip(); + + int maxIterations = 100; + while (client.getHandshakeStatus() != SSLEngineResult.HandshakeStatus.NOT_HANDSHAKING + || server.getHandshakeStatus() != SSLEngineResult.HandshakeStatus.NOT_HANDSHAKING) { + maxIterations--; + assertThat(maxIterations).isPositive(); + + if (client.getHandshakeStatus() == SSLEngineResult.HandshakeStatus.NEED_WRAP) { + clientOut.clear(); + SSLEngineResult cr = client.wrap(ByteBuffer.allocate(0), clientOut); + clientOut.flip(); + assertThat(cr.getStatus()).isEqualTo(SSLEngineResult.Status.OK); + } + + if (server.getHandshakeStatus() == SSLEngineResult.HandshakeStatus.NEED_WRAP) { + serverOut.clear(); + SSLEngineResult sr = server.wrap(ByteBuffer.allocate(0), serverOut); + serverOut.flip(); + assertThat(sr.getStatus()).isEqualTo(SSLEngineResult.Status.OK); + } + + if (client.getHandshakeStatus() == SSLEngineResult.HandshakeStatus.NEED_UNWRAP) { + if (serverOut.hasRemaining()) { + ByteBuffer temp = ByteBuffer.allocate(BUFFER_CAPACITY); + temp.put(serverOut); + temp.flip(); + client.unwrap(temp, ByteBuffer.allocate(BUFFER_CAPACITY)); + } + } + + if (server.getHandshakeStatus() == SSLEngineResult.HandshakeStatus.NEED_UNWRAP) { + if (clientOut.hasRemaining()) { + ByteBuffer temp = ByteBuffer.allocate(BUFFER_CAPACITY); + temp.put(clientOut); + temp.flip(); + server.unwrap(temp, ByteBuffer.allocate(BUFFER_CAPACITY)); + } + } + + Thread.sleep(1); + } + } + + private static boolean getBooleanField(SelfSignedSslEngineSource source, String fieldName) { + try { + Field field = SelfSignedSslEngineSource.class.getDeclaredField(fieldName); + field.setAccessible(true); + return field.getBoolean(source); + } catch (NoSuchFieldException | IllegalAccessException e) { + throw new AssertionError("Unable to read field '" + fieldName + "'", e); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/extras/TestMitmManager.java b/src/test/java/org/littleshoot/proxy/extras/TestMitmManager.java new file mode 100644 index 00000000..ca4f576e --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/extras/TestMitmManager.java @@ -0,0 +1,7 @@ +package org.littleshoot.proxy.extras; + +public class TestMitmManager extends SelfSignedMitmManager { + public TestMitmManager() { + super("target/littleproxy_keystore.jks", true, true); + } +} diff --git a/src/test/java/org/littleshoot/proxy/extras/TrustingTrustManagerTest.java b/src/test/java/org/littleshoot/proxy/extras/TrustingTrustManagerTest.java new file mode 100644 index 00000000..1634758f --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/extras/TrustingTrustManagerTest.java @@ -0,0 +1,111 @@ +package org.littleshoot.proxy.extras; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatCode; + +import java.security.cert.X509Certificate; +import org.junit.jupiter.api.Test; + +class TrustingTrustManagerTest { + + private final TrustingTrustManager trustManager = new TrustingTrustManager(); + + @Test + void testCheckClientTrusted() { + X509Certificate[] certs = new X509Certificate[0]; + + // Should not throw any exception for any client certificate + assertThatCode(() -> trustManager.checkClientTrusted(certs, "RSA")).doesNotThrowAnyException(); + } + + @Test + void testCheckServerTrusted() { + X509Certificate[] certs = new X509Certificate[0]; + + // Should not throw any exception for any server certificate + assertThatCode(() -> trustManager.checkServerTrusted(certs, "RSA")).doesNotThrowAnyException(); + } + + @Test + void testCheckClientTrustedWithNullCerts() { + // Should not throw even with null certificates + assertThatCode(() -> trustManager.checkClientTrusted(null, "RSA")).doesNotThrowAnyException(); + } + + @Test + void testCheckServerTrustedWithNullCerts() { + // Should not throw even with null certificates + assertThatCode(() -> trustManager.checkServerTrusted(null, "RSA")).doesNotThrowAnyException(); + } + + @Test + void testCheckClientTrustedWithNullAuthType() { + // Should not throw even with null auth type + assertThatCode(() -> trustManager.checkClientTrusted(new X509Certificate[0], null)) + .doesNotThrowAnyException(); + } + + @Test + void testCheckServerTrustedWithNullAuthType() { + // Should not throw even with null auth type + assertThatCode(() -> trustManager.checkServerTrusted(new X509Certificate[0], null)) + .doesNotThrowAnyException(); + } + + @Test + void testCheckClientTrustedWithBothNull() { + // Should not throw even when both parameters are null + assertThatCode(() -> trustManager.checkClientTrusted(null, null)).doesNotThrowAnyException(); + } + + @Test + void testCheckServerTrustedWithBothNull() { + // Should not throw even when both parameters are null + assertThatCode(() -> trustManager.checkServerTrusted(null, null)).doesNotThrowAnyException(); + } + + @Test + void testGetAcceptedIssuers() { + // Should return null (accepts all issuers) + X509Certificate[] issuers = trustManager.getAcceptedIssuers(); + assertThat(issuers).isNull(); + } + + @Test + void testTrustsAllClients() { + assertThatCode(() -> trustManager.checkClientTrusted(new X509Certificate[0], "RSA")) + .doesNotThrowAnyException(); + } + + @Test + void testTrustsAllServers() { + assertThatCode(() -> trustManager.checkServerTrusted(new X509Certificate[0], "RSA")) + .doesNotThrowAnyException(); + } + + @Test + void testTrustsWithVariousAuthTypes() { + // Test with various authentication types + X509Certificate[] certs = new X509Certificate[0]; + + assertThatCode(() -> trustManager.checkClientTrusted(certs, "RSA")).doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkClientTrusted(certs, "DSA")).doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkClientTrusted(certs, "EC")).doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkClientTrusted(certs, "DiffieHellman")) + .doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkClientTrusted(certs, "")).doesNotThrowAnyException(); + + assertThatCode(() -> trustManager.checkServerTrusted(certs, "RSA")).doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkServerTrusted(certs, "DSA")).doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkServerTrusted(certs, "EC")).doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkServerTrusted(certs, "DiffieHellman")) + .doesNotThrowAnyException(); + assertThatCode(() -> trustManager.checkServerTrusted(certs, "")).doesNotThrowAnyException(); + } + + @Test + void testTrustManagerInstanceIsNotNull() { + // Basic sanity check that the instance exists + assertThat(trustManager).isNotNull(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/haproxy/BaseProxyProtocolTest.java b/src/test/java/org/littleshoot/proxy/haproxy/BaseProxyProtocolTest.java index 7fcde2c3..3a703d2c 100644 --- a/src/test/java/org/littleshoot/proxy/haproxy/BaseProxyProtocolTest.java +++ b/src/test/java/org/littleshoot/proxy/haproxy/BaseProxyProtocolTest.java @@ -14,123 +14,224 @@ import io.netty.handler.codec.haproxy.HAProxyMessageDecoder; import io.netty.handler.codec.http.HttpRequestDecoder; import io.netty.handler.codec.http.HttpRequestEncoder; +import io.netty.handler.codec.http.HttpResponseDecoder; +import io.netty.handler.ssl.SslContext; +import io.netty.handler.ssl.SslContextBuilder; +import io.netty.handler.ssl.util.InsecureTrustManagerFactory; +import io.netty.handler.ssl.util.SelfSignedCertificate; import io.netty.handler.timeout.ReadTimeoutHandler; -import org.junit.After; +import java.net.InetSocketAddress; +import java.security.cert.CertificateException; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import javax.net.ssl.SSLEngine; +import javax.net.ssl.SSLException; +import org.junit.jupiter.api.AfterEach; import org.littleshoot.proxy.HttpProxyServer; +import org.littleshoot.proxy.HttpProxyServerBootstrap; +import org.littleshoot.proxy.SslEngineSource; import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import java.net.InetSocketAddress; - /** - * Base for running Proxy protocol tests. - * Proxy Protocol tests need special client and servers that are - * capable of emitting and consuming proxy protocol headers. + * Base for running Proxy protocol tests. Proxy Protocol tests need special client and servers that + * are capable of emitting and consuming proxy protocol headers. + * + *

Subclasses may override {@link #useTlsInbound()} to enable TLS between the client and the + * proxy. When TLS is enabled, the client sends the PROXY protocol header as cleartext + * before the TLS handshake, matching real-world deployments (e.g. AWS NLB → proxy). + * + *

Subclasses may override {@link #sendProxyHeaderBeforeTls()} to control the ordering of the + * PROXY header relative to the TLS handshake. Returning {@code false} causes the PROXY header to be + * sent inside the TLS tunnel (after the handshake), which is useful for negative testing. */ public abstract class BaseProxyProtocolTest { - private EventLoopGroup childGroup; - private EventLoopGroup parentGroup; - private EventLoopGroup clientWorkGroup; - private ProxyProtocolServerHandler proxyProtocolServerHandler; - private HttpProxyServer proxyServer; - private int proxyPort; - private boolean acceptProxy = true; - private boolean sendProxy = true; - int serverPort; - static final String SOURCE_ADDRESS = "192.168.0.153"; - static final String DESTINATION_ADDRESS = "192.168.0.154"; - static final String SOURCE_PORT = "123"; - static final String DESTINATION_PORT = "456"; - - - public void setup(boolean acceptProxy, boolean sendProxy) throws Exception { - this.acceptProxy = acceptProxy; - this.sendProxy = sendProxy; - startProxyServer(); - startServer(); - startClient(); + protected CountDownLatch clientTlsHandshakeDone; + private CountDownLatch serverHandlerReady; + private EventLoopGroup childGroup; + private EventLoopGroup parentGroup; + private EventLoopGroup clientWorkGroup; + private ProxyProtocolServerHandler proxyProtocolServerHandler; + + private HttpProxyServer proxyServer; + private int proxyPort; + private boolean acceptProxy = true; + private boolean sendProxy = true; + int serverPort; + static final String SOURCE_ADDRESS = "192.168.0.153"; + static final String DESTINATION_ADDRESS = "192.168.0.154"; + static final String SOURCE_PORT = "123"; + static final String DESTINATION_PORT = "456"; + private volatile ProxyProtocolClientHandler clientHandler; + + protected boolean useTlsInbound() { + return false; + } + + /** Controls whether the PROXY protocol header is sent before or after the TLS handshake. */ + protected boolean sendProxyHeaderBeforeTls() { + return true; + } + + boolean isClientTlsHandshakeSuccess() { + return clientHandler != null && clientHandler.isTlsHandshakeSuccess(); + } + + Throwable getClientTlsHandshakeFailureCause() { + return clientHandler != null ? clientHandler.getTlsHandshakeFailureCause() : null; + } + + protected final void setup(boolean acceptProxy, boolean sendProxy) throws Exception { + this.acceptProxy = acceptProxy; + this.sendProxy = sendProxy; + this.serverHandlerReady = new CountDownLatch(1); + this.clientTlsHandshakeDone = new CountDownLatch(1); + startProxyServer(); + startServer(); + startClient(); + } + + final void startServer() { + parentGroup = new NioEventLoopGroup(); + childGroup = new NioEventLoopGroup(); + ServerBootstrap b = new ServerBootstrap(); + b.group(parentGroup, childGroup) + .channelFactory(NioServerSocketChannel::new) + .childHandler( + new ChannelInitializer() { + @Override + public void initChannel(SocketChannel ch) { + proxyProtocolServerHandler = new ProxyProtocolServerHandler(); + ch.pipeline() + .addLast(new HAProxyMessageDecoder()) + .addLast(new HttpRequestDecoder()) + .addLast(proxyProtocolServerHandler); + serverHandlerReady.countDown(); + } + }) + .option(ChannelOption.SO_BACKLOG, 128) + .childOption(ChannelOption.SO_KEEPALIVE, true); + + ChannelFuture f = b.bind(0).awaitUninterruptibly(); + Throwable cause = f.cause(); + if (cause != null) { + throw new RuntimeException(cause); } - - void startServer() { - parentGroup = new NioEventLoopGroup(); - childGroup = new NioEventLoopGroup(); - ServerBootstrap b = new ServerBootstrap(); - b.group(parentGroup, childGroup) - .channelFactory(NioServerSocketChannel::new) - .childHandler(new ChannelInitializer() { - @Override - public void initChannel(SocketChannel ch) { - proxyProtocolServerHandler = new ProxyProtocolServerHandler(); - ch.pipeline().addLast(new HAProxyMessageDecoder()).addLast(new HttpRequestDecoder()).addLast(proxyProtocolServerHandler); - } - }).option(ChannelOption.SO_BACKLOG, 128) - .childOption(ChannelOption.SO_KEEPALIVE, true); - - ChannelFuture f = b.bind(0) - .awaitUninterruptibly(); - Throwable cause = f.cause(); - if (cause != null) { - throw new RuntimeException(cause); - } - serverPort = ((InetSocketAddress) f.channel().localAddress()).getPort(); - Runtime.getRuntime().addShutdownHook(new Thread(new Runnable() { - public void run() { - stopServer(); + serverPort = ((InetSocketAddress) f.channel().localAddress()).getPort(); + Runtime.getRuntime().addShutdownHook(new Thread(this::stopServer, "stopServerHook")); + } + + final void startClient() throws Exception { + String host = "localhost"; + clientWorkGroup = new NioEventLoopGroup(); + Bootstrap b = new Bootstrap(); + b.group(clientWorkGroup); + b.channel(NioSocketChannel.class); + b.option(ChannelOption.SO_KEEPALIVE, true); + b.handler( + new ChannelInitializer() { + @Override + public void initChannel(SocketChannel ch) throws SSLException { + ch.pipeline().addLast(new ReadTimeoutHandler(1)); + if (!useTlsInbound() && acceptProxy) { + ch.pipeline().addLast(new ProxyProtocolTestEncoder()); } - }, "stopServerHook")); - } - void startClient() throws Exception { - String host = "localhost"; - clientWorkGroup = new NioEventLoopGroup(); - Bootstrap b = new Bootstrap(); - b.group(clientWorkGroup); - b.channel(NioSocketChannel.class); - b.option(ChannelOption.SO_KEEPALIVE, true); - b.handler(new ChannelInitializer() { - @Override - public void initChannel(SocketChannel ch) { - ch.pipeline().addLast(new ReadTimeoutHandler(1)); - if (acceptProxy) { - ch.pipeline().addLast(new ProxyProtocolTestEncoder()); - } - ch.pipeline().addLast(new HttpRequestEncoder()).addLast(new ProxyProtocolClientHandler(serverPort, getProxyProtocolHeader())); + // Build client SslContext for TLS cases. + SslContext clientSslCtx = null; + if (useTlsInbound()) { + clientSslCtx = + SslContextBuilder.forClient() + .trustManager(InsecureTrustManagerFactory.INSTANCE) + .build(); } - }); - ChannelFuture f = b.connect(host, proxyPort).sync(); - f.channel().closeFuture().sync(); - } - - HAProxyMessage getRelayedHaProxyMessage() { - return proxyProtocolServerHandler.getHaProxyMessage(); - } - private void stopServer() { - childGroup.shutdownGracefully(); - parentGroup.shutdownGracefully(); - } + clientHandler = + new ProxyProtocolClientHandler( + serverPort, + getProxyProtocolHeader(), + clientSslCtx, + proxyPort, + clientTlsHandshakeDone, + sendProxyHeaderBeforeTls(), + acceptProxy); + + ch.pipeline() + .addLast(new HttpResponseDecoder()) + .addLast(new HttpRequestEncoder()) + .addLast(clientHandler); + } + }); + b.connect(host, proxyPort).sync(); + } - private void stopProxyServer() { - proxyServer.abort(); + HAProxyMessage getRelayedHaProxyMessage() throws InterruptedException { + if (!serverHandlerReady.await(5, TimeUnit.SECONDS)) { + return null; } + return proxyProtocolServerHandler.awaitHaProxyMessage(3, TimeUnit.SECONDS); + } + + private void stopServer() { + childGroup.shutdownGracefully(); + parentGroup.shutdownGracefully(); + } + + private void stopProxyServer() { + proxyServer.abort(); + } + + private void startProxyServer() throws CertificateException, SSLException { + HttpProxyServerBootstrap builder = + DefaultHttpProxyServer.bootstrap() + .withPort(0) + .withAcceptProxyProtocol(acceptProxy) + .withSendProxyProtocol(sendProxy); + + if (useTlsInbound()) { + SelfSignedCertificate ssc = new SelfSignedCertificate(); + SslContext sslCtx = SslContextBuilder.forServer(ssc.certificate(), ssc.privateKey()).build(); + + builder.withSslEngineSource( + new SslEngineSource() { + @Override + public SSLEngine newSslEngine() { + return sslCtx.newEngine(io.netty.buffer.ByteBufAllocator.DEFAULT); + } - private void startProxyServer() { - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .withAcceptProxyProtocol(acceptProxy) - .withSendProxyProtocol(sendProxy) - .start(); - proxyPort = proxyServer.getListenAddress().getPort(); - - } + @Override + public SSLEngine newSslEngine(String peerHost, int peerPort) { + return sslCtx.newEngine(io.netty.buffer.ByteBufAllocator.DEFAULT, peerHost, peerPort); + } + }); - private ProxyProtocolHeader getProxyProtocolHeader() { - return new ProxyProtocolHeader(SOURCE_ADDRESS, DESTINATION_ADDRESS, SOURCE_PORT, DESTINATION_PORT); + builder.withAuthenticateSslClients(false); } - @After - public void tearDown() { - stopServer(); - stopProxyServer(); + customizeProxyServer(builder); + + proxyServer = builder.start(); + proxyPort = proxyServer.getListenAddress().getPort(); + } + + /** + * Hook for subclasses to further configure the proxy bootstrap (e.g. attach an {@link + * org.littleshoot.proxy.ActivityTracker} or {@link org.littleshoot.proxy.ChainedProxyManager}) + * before it is started. Default is a no-op. + */ + protected void customizeProxyServer(HttpProxyServerBootstrap builder) {} + + private ProxyProtocolHeader getProxyProtocolHeader() { + return new ProxyProtocolHeader( + SOURCE_ADDRESS, DESTINATION_ADDRESS, SOURCE_PORT, DESTINATION_PORT); + } + + @AfterEach + final void tearDown() { + stopServer(); + stopProxyServer(); + if (clientWorkGroup != null) { + clientWorkGroup.shutdownGracefully(); } + } } diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolClientHandler.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolClientHandler.java index 288439bf..f743aa54 100644 --- a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolClientHandler.java +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolClientHandler.java @@ -1,37 +1,173 @@ package org.littleshoot.proxy.haproxy; -import io.netty.channel.ChannelHandlerContext; -import io.netty.channel.ChannelInboundHandlerAdapter; +import io.netty.buffer.ByteBuf; +import io.netty.channel.*; import io.netty.handler.codec.http.DefaultHttpRequest; import io.netty.handler.codec.http.HttpMethod; import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; import io.netty.handler.codec.http.HttpVersion; - +import io.netty.handler.ssl.SslContext; +import io.netty.handler.ssl.SslHandler; +import io.netty.util.ReferenceCountUtil; +import io.netty.util.concurrent.Future; +import io.netty.util.concurrent.GenericFutureListener; +import java.nio.charset.StandardCharsets; +import java.util.concurrent.CountDownLatch; public class ProxyProtocolClientHandler extends ChannelInboundHandlerAdapter { - private static final String HOST = "http://localhost"; - private int serverPort; - private ProxyProtocolHeader proxyProtocolHeader; + private static final String HOST = "http://localhost"; + private final int serverPort; + private final ProxyProtocolHeader proxyProtocolHeader; + private final SslContext clientSslContext; + private final int proxyPort; + private final CountDownLatch tlsHandshakeDone; + private final boolean sendProxyBeforeTls; + private final boolean sendInboundProxyHeader; + private volatile boolean tlsHandshakeSuccess; + private volatile Throwable tlsHandshakeFailureCause; + + ProxyProtocolClientHandler( + int serverPort, + ProxyProtocolHeader proxyProtocolHeader, + SslContext clientSslContext, + int proxyPort, + CountDownLatch tlsHandshakeDone, + boolean sendProxyBeforeTls, + boolean sendInboundProxyHeader) { + this.serverPort = serverPort; + this.proxyProtocolHeader = proxyProtocolHeader; + this.clientSslContext = clientSslContext; + this.proxyPort = proxyPort; + this.tlsHandshakeDone = tlsHandshakeDone; + this.sendProxyBeforeTls = sendProxyBeforeTls; + this.sendInboundProxyHeader = sendInboundProxyHeader; + } - ProxyProtocolClientHandler(int serverPort, ProxyProtocolHeader proxyProtocolHeader) { - this.serverPort = serverPort; - this.proxyProtocolHeader = proxyProtocolHeader; + @Override + public void channelActive(ChannelHandlerContext ctx) { + if (clientSslContext != null) { + if (sendProxyBeforeTls) { + // Correct order: write PROXY header as raw cleartext bytes, then + // add SslHandler and send CONNECT once TLS handshake completes. + sendProxyThenTls(ctx); + } else { + // Wrong order: TLS first, then PROXY header inside the encrypted tunnel. + sendTlsThenProxy(ctx); + } + } else { + // Non-TLS mode: send PROXY header + CONNECT directly. + if (sendInboundProxyHeader) { + ctx.write(getHAProxyHeader()); + } + ctx.writeAndFlush(getConnectRequest()); } + } + + /** Correct order: cleartext PROXY header → TLS handshake → HTTP CONNECT. */ + private void sendProxyThenTls(ChannelHandlerContext ctx) { + ByteBuf buf = ctx.alloc().buffer(); + buf.writeBytes(getHAProxyHeader().getBytes(StandardCharsets.US_ASCII)); + + // Write through the pipeline head to bypass HttpRequestEncoder, + // which only handles HttpRequest/HttpContent, not raw ByteBuf. + ctx.pipeline() + .firstContext() + .writeAndFlush(buf) + .addListener( + (ChannelFutureListener) + future -> { + if (!future.isSuccess()) { + signalTlsDone(false, future.cause()); + ctx.close(); + return; + } + addSslAndConnect(ctx); + }); + } + + /** + * Wrong order (for negative testing): TLS handshake first → PROXY header sent inside the + * encrypted tunnel → proxy cannot decode it. + */ + private void sendTlsThenProxy(ChannelHandlerContext ctx) { + SslHandler sslHandler = clientSslContext.newHandler(ctx.alloc(), "localhost", proxyPort); + ctx.pipeline().addFirst("ssl", sslHandler); + + sslHandler + .handshakeFuture() + .addListener( + (GenericFutureListener>) + hsFuture -> { + signalTlsDone(hsFuture.isSuccess(), hsFuture.cause()); + if (hsFuture.isSuccess()) { + // PROXY header is now encrypted — proxy can't decode it. + ByteBuf buf = ctx.alloc().buffer(); + buf.writeBytes(getHAProxyHeader().getBytes(StandardCharsets.US_ASCII)); + ctx.writeAndFlush(buf); + ctx.writeAndFlush(getConnectRequest()); + } else { + ctx.close(); + } + }); + } - @Override - public void channelActive(ChannelHandlerContext ctx) { - ctx.write(getHAProxyHeader()); - ctx.writeAndFlush(getConnectRequest()); + /** Adds the SslHandler and sends CONNECT after the handshake succeeds. */ + private void addSslAndConnect(ChannelHandlerContext ctx) { + ChannelPipeline pipeline = ctx.pipeline(); + SslHandler sslHandler = clientSslContext.newHandler(ctx.alloc(), "localhost", proxyPort); + pipeline.addFirst("ssl", sslHandler); + + sslHandler + .handshakeFuture() + .addListener( + (GenericFutureListener>) + hsFuture -> { + signalTlsDone(hsFuture.isSuccess(), hsFuture.cause()); + if (hsFuture.isSuccess()) { + ctx.writeAndFlush(getConnectRequest()); + } else { + ctx.close(); + } + }); + } + + @Override + public void channelRead(ChannelHandlerContext ctx, Object msg) { + if (msg instanceof HttpResponse) { + ReferenceCountUtil.release(msg); + ctx.close(); } + } - private HttpRequest getConnectRequest() { - return new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.CONNECT, HOST + ":" + serverPort); + private void signalTlsDone(boolean success, Throwable cause) { + tlsHandshakeSuccess = success; + tlsHandshakeFailureCause = cause; + if (tlsHandshakeDone != null) { + tlsHandshakeDone.countDown(); } + } + boolean isTlsHandshakeSuccess() { + return tlsHandshakeSuccess; + } - private String getHAProxyHeader() { - return String.format("PROXY TCP4 %s %s %s %s\r\n", proxyProtocolHeader.getSourceAddress(), proxyProtocolHeader.getDestinationAddress(), - proxyProtocolHeader.getSourcePort(), proxyProtocolHeader.getDestinationPort()); - } + Throwable getTlsHandshakeFailureCause() { + return tlsHandshakeFailureCause; + } + + private HttpRequest getConnectRequest() { + return new DefaultHttpRequest( + HttpVersion.HTTP_1_1, HttpMethod.CONNECT, HOST + ":" + serverPort); + } + + private String getHAProxyHeader() { + return String.format( + "PROXY TCP4 %s %s %s %s\r\n", + proxyProtocolHeader.getSourceAddress(), + proxyProtocolHeader.getDestinationAddress(), + proxyProtocolHeader.getSourcePort(), + proxyProtocolHeader.getDestinationPort()); + } } diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolHeader.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolHeader.java index d989139b..02ac8fe6 100644 --- a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolHeader.java +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolHeader.java @@ -2,32 +2,32 @@ class ProxyProtocolHeader { - private String sourceAddress; - private String destinationAddress; - private String sourcePort; - private String destinationPort; - - ProxyProtocolHeader(String sourceAddress, String destinationAddress, String sourcePort, String destinationPort) { - this.sourceAddress = sourceAddress; - this.destinationAddress = destinationAddress; - this.sourcePort = sourcePort; - this.destinationPort = destinationPort; - } - - String getSourceAddress() { - return sourceAddress; - } - - String getDestinationAddress() { - return destinationAddress; - } - - String getSourcePort() { - return sourcePort; - } - - String getDestinationPort() { - return destinationPort; - } - + private final String sourceAddress; + private final String destinationAddress; + private final String sourcePort; + private final String destinationPort; + + ProxyProtocolHeader( + String sourceAddress, String destinationAddress, String sourcePort, String destinationPort) { + this.sourceAddress = sourceAddress; + this.destinationAddress = destinationAddress; + this.sourcePort = sourcePort; + this.destinationPort = destinationPort; + } + + String getSourceAddress() { + return sourceAddress; + } + + String getDestinationAddress() { + return destinationAddress; + } + + String getSourcePort() { + return sourcePort; + } + + String getDestinationPort() { + return destinationPort; + } } diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolHttpConnectChainedProxyTest.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolHttpConnectChainedProxyTest.java new file mode 100644 index 00000000..9c7309f9 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolHttpConnectChainedProxyTest.java @@ -0,0 +1,233 @@ +package org.littleshoot.proxy.haproxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import io.netty.bootstrap.Bootstrap; +import io.netty.bootstrap.ServerBootstrap; +import io.netty.buffer.ByteBuf; +import io.netty.buffer.Unpooled; +import io.netty.channel.Channel; +import io.netty.channel.ChannelFutureListener; +import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.ChannelInboundHandlerAdapter; +import io.netty.channel.ChannelInitializer; +import io.netty.channel.EventLoopGroup; +import io.netty.channel.nio.NioEventLoopGroup; +import io.netty.channel.socket.SocketChannel; +import io.netty.channel.socket.nio.NioServerSocketChannel; +import io.netty.channel.socket.nio.NioSocketChannel; +import io.netty.util.ReferenceCountUtil; +import java.io.OutputStream; +import java.net.InetSocketAddress; +import java.net.Socket; +import java.nio.charset.StandardCharsets; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ChainedProxyAdapter; +import org.littleshoot.proxy.HttpProxyServer; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; + +/** + * With {@code sendProxyProtocol=true} and an HTTP CONNECT chained proxy, the PROXY header must be + * tunnelled to the final server (first bytes after the CONNECT handshake), not sent to the + * intermediate proxy ahead of the CONNECT. + * + *

Topology: raw client → downstream LittleProxy → intermediate HTTP CONNECT proxy → final + * server. + */ +@Tag("slow-test") +public final class ProxyProtocolHttpConnectChainedProxyTest { + + private EventLoopGroup bossGroup; + private EventLoopGroup workerGroup; + private HttpProxyServer downstreamProxy; + private Socket clientSocket; + + private final AtomicReference intermediateFirstBytes = new AtomicReference<>(); + private final AtomicReference finalServerFirstBytes = new AtomicReference<>(); + private final CountDownLatch intermediateReceivedConnect = new CountDownLatch(1); + private final CountDownLatch finalServerReceivedData = new CountDownLatch(1); + + @AfterEach + void tearDown() throws Exception { + if (clientSocket != null) { + clientSocket.close(); + } + if (downstreamProxy != null) { + downstreamProxy.abort(); + } + if (bossGroup != null) { + bossGroup.shutdownGracefully(); + } + if (workerGroup != null) { + workerGroup.shutdownGracefully(); + } + } + + @Test + void proxyProtocolHeaderIsTunnelledToFinalServerNotIntermediateProxy() throws Exception { + bossGroup = new NioEventLoopGroup(1); + workerGroup = new NioEventLoopGroup(); + + int finalServerPort = startFinalServer(); + int intermediatePort = startIntermediateConnectProxy(finalServerPort); + int downstreamPort = startDownstreamProxy(intermediatePort); + + sendConnectThroughDownstream(downstreamPort, finalServerPort); + + assertThat(intermediateReceivedConnect.await(5, TimeUnit.SECONDS)) + .as("intermediate proxy should receive the CONNECT") + .isTrue(); + assertThat(finalServerReceivedData.await(5, TimeUnit.SECONDS)) + .as("final server should receive tunnelled bytes") + .isTrue(); + + assertThat(intermediateFirstBytes.get()) + .as("intermediate HTTP proxy must see a plain CONNECT as its first bytes") + .startsWith("CONNECT"); + assertThat(intermediateFirstBytes.get()) + .as("intermediate HTTP proxy must NOT receive a PROXY header") + .doesNotContain("PROXY TCP"); + + assertThat(finalServerFirstBytes.get()) + .as("final server must receive the PROXY protocol header as the first tunnelled bytes") + .startsWith("PROXY TCP"); + } + + /** Captures the first bytes the ultimate destination receives. */ + private int startFinalServer() { + ServerBootstrap b = new ServerBootstrap(); + b.group(bossGroup, workerGroup) + .channel(NioServerSocketChannel.class) + .childHandler( + new ChannelInitializer() { + @Override + protected void initChannel(SocketChannel ch) { + ch.pipeline() + .addLast( + new ChannelInboundHandlerAdapter() { + @Override + public void channelRead(ChannelHandlerContext ctx, Object msg) { + ByteBuf buf = (ByteBuf) msg; + finalServerFirstBytes.compareAndSet( + null, buf.toString(StandardCharsets.US_ASCII)); + buf.release(); + finalServerReceivedData.countDown(); + } + }); + } + }); + Channel ch = b.bind(0).syncUninterruptibly().channel(); + return ((InetSocketAddress) ch.localAddress()).getPort(); + } + + /** + * Minimal HTTP CONNECT proxy: records its first bytes, answers 200, then relays to the final + * server. + */ + private int startIntermediateConnectProxy(int finalServerPort) { + ServerBootstrap b = new ServerBootstrap(); + b.group(bossGroup, workerGroup) + .channel(NioServerSocketChannel.class) + .childHandler( + new ChannelInitializer() { + @Override + protected void initChannel(SocketChannel ch) { + ch.pipeline() + .addLast( + new ChannelInboundHandlerAdapter() { + private volatile Channel outbound; + private boolean connectHandled; + + @Override + public void channelRead(ChannelHandlerContext ctx, Object msg) { + ByteBuf buf = (ByteBuf) msg; + if (!connectHandled) { + connectHandled = true; + intermediateFirstBytes.compareAndSet( + null, buf.toString(StandardCharsets.US_ASCII)); + buf.release(); + intermediateReceivedConnect.countDown(); + connectToFinalServerThenRespond(ctx, finalServerPort); + } else if (outbound != null) { + outbound.writeAndFlush(msg); + } else { + buf.release(); + } + } + + private void connectToFinalServerThenRespond( + ChannelHandlerContext ctx, int port) { + Bootstrap cb = new Bootstrap(); + cb.group(workerGroup) + .channel(NioSocketChannel.class) + .handler( + new ChannelInboundHandlerAdapter() { + @Override + public void channelRead(ChannelHandlerContext c, Object m) { + ReferenceCountUtil.release(m); + } + }); + cb.connect("localhost", port) + .addListener( + (ChannelFutureListener) + f -> { + if (f.isSuccess()) { + outbound = f.channel(); + ctx.writeAndFlush( + Unpooled.copiedBuffer( + "HTTP/1.1 200 Connection Established\r\n\r\n", + StandardCharsets.US_ASCII)); + } else { + ctx.close(); + } + }); + } + }); + } + }); + Channel ch = b.bind(0).syncUninterruptibly().channel(); + return ((InetSocketAddress) ch.localAddress()).getPort(); + } + + private int startDownstreamProxy(int intermediatePort) { + downstreamProxy = + DefaultHttpProxyServer.bootstrap() + .withName("Downstream") + .withPort(0) + .withSendProxyProtocol(true) + .withChainProxyManager( + (httpRequest, chainedProxies, clientDetails) -> + chainedProxies.add( + new ChainedProxyAdapter() { + @Override + public InetSocketAddress getChainedProxyAddress() { + return new InetSocketAddress("localhost", intermediatePort); + } + })) + .start(); + return downstreamProxy.getListenAddress().getPort(); + } + + private void sendConnectThroughDownstream(int downstreamPort, int finalServerPort) + throws Exception { + clientSocket = new Socket("localhost", downstreamPort); + clientSocket.setSoTimeout(5000); + OutputStream out = clientSocket.getOutputStream(); + String connect = + "CONNECT localhost:" + + finalServerPort + + " HTTP/1.1\r\nHost: localhost:" + + finalServerPort + + "\r\n\r\n"; + out.write(connect.getBytes(StandardCharsets.US_ASCII)); + out.flush(); + // Discard the CONNECT response; the socket stays open until tearDown so the tunnel lives while + // the PROXY header propagates. + clientSocket.getInputStream().read(new byte[1024]); + } +} diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolOrderTest.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolOrderTest.java new file mode 100644 index 00000000..b46ac4de --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolOrderTest.java @@ -0,0 +1,29 @@ +package org.littleshoot.proxy.haproxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import io.netty.handler.codec.haproxy.HAProxyMessage; +import org.junit.jupiter.api.Test; + +/** + * Verifies that LittleProxy's inbound pipeline decodes the PROXY protocol header performing the TLS + * handshake. + */ +public final class ProxyProtocolOrderTest extends BaseProxyProtocolTest { + + @Override + protected boolean useTlsInbound() { + return true; + } + + @Test + void proxyProtocolIsDecodedBeforeTlsOnInbound() throws Exception { + setup(true, true); + + HAProxyMessage relayed = getRelayedHaProxyMessage(); + + assertThat(relayed).as("PROXY protocol message should be decoded even with TLS").isNotNull(); + assertThat(relayed.sourceAddress()).isEqualTo(SOURCE_ADDRESS); + assertThat(relayed.destinationAddress()).isEqualTo(DESTINATION_ADDRESS); + } +} diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolServerHandler.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolServerHandler.java index eec2e72d..b5ad7924 100644 --- a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolServerHandler.java +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolServerHandler.java @@ -3,19 +3,28 @@ import io.netty.channel.ChannelHandlerContext; import io.netty.channel.ChannelInboundHandlerAdapter; import io.netty.handler.codec.haproxy.HAProxyMessage; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; public class ProxyProtocolServerHandler extends ChannelInboundHandlerAdapter { - private HAProxyMessage haProxyMessage; + private final CountDownLatch messageLatch = new CountDownLatch(1); + private volatile HAProxyMessage haProxyMessage; - @Override - public void channelRead(ChannelHandlerContext ctx, Object msg) { - if ( msg instanceof HAProxyMessage){ - this.haProxyMessage = (HAProxyMessage) msg; - } + @Override + public void channelRead(ChannelHandlerContext ctx, Object msg) { + if (msg instanceof HAProxyMessage) { + haProxyMessage = (HAProxyMessage) msg; + messageLatch.countDown(); } + } - HAProxyMessage getHaProxyMessage() { - return haProxyMessage; - } + /** + * Waits up to the given timeout for an HAProxyMessage to arrive. Returns the message, or null if + * none arrived within the timeout. + */ + HAProxyMessage awaitHaProxyMessage(long timeout, TimeUnit unit) throws InterruptedException { + messageLatch.await(timeout, unit); + return haProxyMessage; + } } diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTest.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTest.java index a0c6baee..658a4ae7 100644 --- a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTest.java +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTest.java @@ -1,45 +1,137 @@ package org.littleshoot.proxy.haproxy; +import static java.lang.String.valueOf; +import static org.assertj.core.api.Assertions.assertThat; + import io.netty.handler.codec.haproxy.HAProxyMessage; -import org.junit.Assert; -import org.junit.Test; - -public class ProxyProtocolTest extends BaseProxyProtocolTest { - - - private static final String LOCALHOST = "127.0.0.1"; - private static final boolean ACCEPT_PROXY = true; - private static final boolean SEND_PROXY = true; - private static final boolean DO_NOT_ACCEPT_PROXY = false; - private static final boolean DO_NOT_SEND_PROXY = false; - - @Test - public void canRelayProxyProtocolHeader() throws Exception { - setup(ACCEPT_PROXY, SEND_PROXY); - HAProxyMessage haProxyMessage = getRelayedHaProxyMessage(); - Assert.assertNotNull(haProxyMessage); - Assert.assertEquals(SOURCE_ADDRESS, haProxyMessage.sourceAddress()); - Assert.assertEquals(DESTINATION_ADDRESS, haProxyMessage.destinationAddress()); - Assert.assertEquals(SOURCE_PORT, String.valueOf(haProxyMessage.sourcePort())); - Assert.assertEquals(DESTINATION_PORT, String.valueOf(haProxyMessage.destinationPort())); - } +import java.net.InetSocketAddress; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import javax.net.ssl.SSLSession; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ActivityTrackerAdapter; +import org.littleshoot.proxy.ChainedProxyAdapter; +import org.littleshoot.proxy.FlowContext; +import org.littleshoot.proxy.HttpProxyServerBootstrap; - @Test - public void canSendProxyProtocolHeader() throws Exception { - setup(DO_NOT_ACCEPT_PROXY, SEND_PROXY); - HAProxyMessage haProxyMessage = getRelayedHaProxyMessage(); - Assert.assertNotNull(haProxyMessage); - Assert.assertEquals(LOCALHOST, haProxyMessage.sourceAddress()); - Assert.assertEquals(LOCALHOST, haProxyMessage.destinationAddress()); - Assert.assertEquals(String.valueOf(serverPort), String.valueOf(haProxyMessage.destinationPort())); - } +public final class ProxyProtocolTest extends BaseProxyProtocolTest { + + private static final String LOCALHOST = "127.0.0.1"; + private static final boolean ACCEPT_PROXY = true; + private static final boolean SEND_PROXY = true; + private static final boolean DO_NOT_ACCEPT_PROXY = false; + private static final boolean DO_NOT_SEND_PROXY = false; + + private final AtomicInteger clientConnectedCount = new AtomicInteger(); + private final AtomicReference trackerClientAddress = new AtomicReference<>(); + private final AtomicReference chainedProxyClientAddress = + new AtomicReference<>(); + private final CountDownLatch clientDisconnectedLatch = new CountDownLatch(1); - @Test - public void canAcceptProxyProtocolHeader() throws Exception { - setup(ACCEPT_PROXY, DO_NOT_SEND_PROXY); - HAProxyMessage haProxyMessage = getRelayedHaProxyMessage(); - Assert.assertNull(haProxyMessage); + /** Set by a test before {@code setup(...)} to also capture the chained-routing address. */ + private volatile boolean captureChainedProxyClientDetails; + + @Override + protected void customizeProxyServer(HttpProxyServerBootstrap builder) { + // Observational tracker; harmless for the header-relay tests that ignore these captures. + builder.plusActivityTracker( + new ActivityTrackerAdapter() { + @Override + public void clientConnected(FlowContext flowContext) { + clientConnectedCount.incrementAndGet(); + trackerClientAddress.set(flowContext.getClientAddress()); + } + + @Override + public void clientDisconnected(FlowContext flowContext, SSLSession sslSession) { + clientDisconnectedLatch.countDown(); + } + }); + + if (captureChainedProxyClientDetails) { + // Capture the routing address, then fall back to direct so the flow completes. + builder.withChainProxyManager( + (httpRequest, chainedProxies, clientDetails) -> { + chainedProxyClientAddress.set(clientDetails.getClientAddress()); + chainedProxies.add(ChainedProxyAdapter.FALLBACK_TO_DIRECT_CONNECTION); + }); } + } + + @Test + public void canRelayProxyProtocolHeader() throws Exception { + setup(ACCEPT_PROXY, SEND_PROXY); + HAProxyMessage haProxyMessage = getRelayedHaProxyMessage(); + assertThat(haProxyMessage).isNotNull(); + assertThat(haProxyMessage.sourceAddress()).isEqualTo(SOURCE_ADDRESS); + assertThat(haProxyMessage.destinationAddress()).isEqualTo(DESTINATION_ADDRESS); + assertThat(valueOf(haProxyMessage.sourcePort())).isEqualTo(SOURCE_PORT); + assertThat(valueOf(haProxyMessage.destinationPort())).isEqualTo(DESTINATION_PORT); + } + + @Test + public void canSendProxyProtocolHeader() throws Exception { + setup(DO_NOT_ACCEPT_PROXY, SEND_PROXY); + HAProxyMessage haProxyMessage = getRelayedHaProxyMessage(); + assertThat(haProxyMessage).isNotNull(); + assertThat(haProxyMessage.sourceAddress()).isEqualTo(LOCALHOST); + assertThat(haProxyMessage.destinationAddress()).isEqualTo(LOCALHOST); + assertThat(valueOf(haProxyMessage.destinationPort())).isEqualTo(valueOf(serverPort)); + } + + @Test + public void canAcceptProxyProtocolHeader() throws Exception { + setup(ACCEPT_PROXY, DO_NOT_SEND_PROXY); + HAProxyMessage haProxyMessage = getRelayedHaProxyMessage(); + assertThat(haProxyMessage).isNull(); + } + + /** The PROXY header source address is surfaced to both the tracker and chained routing. */ + @Test + public void surfacesRealClientAddressFromProxyProtocolHeader() throws Exception { + captureChainedProxyClientDetails = true; + setup(ACCEPT_PROXY, SEND_PROXY); + + assertThat(clientDisconnectedLatch.await(5, TimeUnit.SECONDS)) + .as("client should connect and then disconnect within the timeout") + .isTrue(); + + assertThat(trackerClientAddress.get()) + .as("ActivityTracker should see the PROXY header source address") + .isNotNull(); + assertThat(trackerClientAddress.get().getHostString()).isEqualTo(SOURCE_ADDRESS); + assertThat(trackerClientAddress.get().getPort()).isEqualTo(Integer.parseInt(SOURCE_PORT)); + + assertThat(chainedProxyClientAddress.get()) + .as("ChainedProxyManager (ClientDetails) should see the PROXY header source address") + .isNotNull(); + assertThat(chainedProxyClientAddress.get().getHostString()).isEqualTo(SOURCE_ADDRESS); + assertThat(chainedProxyClientAddress.get().getPort()).isEqualTo(Integer.parseInt(SOURCE_PORT)); + } + + /** + * With no PROXY header, {@code clientConnected} still fires once, from the first request, with + * the TCP peer address. + */ + @Test + public void clientConnectedFiresOnceWithTcpPeerWhenNoProxyHeader() throws Exception { + setup(DO_NOT_ACCEPT_PROXY, DO_NOT_SEND_PROXY); + + assertThat(clientDisconnectedLatch.await(5, TimeUnit.SECONDS)) + .as("client should connect and then disconnect within the timeout") + .isTrue(); + assertThat(clientConnectedCount.get()) + .as("clientConnected must fire exactly once across the connection lifecycle") + .isEqualTo(1); + assertThat(trackerClientAddress.get()) + .as("clientConnected should carry the TCP peer address when no PROXY header is present") + .isNotNull(); + assertThat(trackerClientAddress.get().getAddress().isLoopbackAddress()) + .as("client connected over loopback, so the reported address should be a loopback address") + .isTrue(); + } } diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTestEncoder.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTestEncoder.java index cc8bbe89..1f9e03f5 100644 --- a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTestEncoder.java +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolTestEncoder.java @@ -7,15 +7,14 @@ public class ProxyProtocolTestEncoder extends MessageToByteEncoder { - @Override - protected void encode(ChannelHandlerContext ctx, String msg, ByteBuf out) { - out.writeBytes(msg.getBytes()); - } - - @Override - public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise promise) throws Exception { - super.write(ctx, msg, promise); - } - + @Override + protected void encode(ChannelHandlerContext ctx, String msg, ByteBuf out) { + out.writeBytes(msg.getBytes()); + } + @Override + public void write(ChannelHandlerContext ctx, Object msg, ChannelPromise promise) + throws Exception { + super.write(ctx, msg, promise); + } } diff --git a/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolWrongOrderTest.java b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolWrongOrderTest.java new file mode 100644 index 00000000..bb64cc9b --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/haproxy/ProxyProtocolWrongOrderTest.java @@ -0,0 +1,49 @@ +package org.littleshoot.proxy.haproxy; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.util.concurrent.TimeUnit; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; + +/** + * Negative test: when the PROXY protocol header is sent after the TLS handshake (i.e. encrypted + * inside the tunnel), the proxy must NOT successfully decode it. + */ +public final class ProxyProtocolWrongOrderTest extends BaseProxyProtocolTest { + + @Override + protected boolean useTlsInbound() { + return true; + } + + @Override + protected boolean sendProxyHeaderBeforeTls() { + return false; + } + + @Test + @Tag("slow-test") + void proxyProtocolInsideTlsTunnelIsNotDecoded() throws Exception { + setup(true, true); + + boolean tlsCompleted = clientTlsHandshakeDone.await(5, TimeUnit.SECONDS); + assertThat(tlsCompleted) + .as("TLS handshake should complete or fail within the timeout") + .isTrue(); + if (isClientTlsHandshakeSuccess()) { + assertThat(getRelayedHaProxyMessage()) + .as("PROXY header sent inside TLS must not be decoded") + .isNull(); + } else { + assertThat(getClientTlsHandshakeFailureCause()) + .as("TLS failure should expose a cause") + .isNotNull(); + assertThat(isClientTlsHandshakeSuccess()) + .as( + "TLS should fail when PROXY header is not sent before TLS. " + "Cause: %s", + getClientTlsHandshakeFailureCause()) + .isFalse(); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionBackpressureTest.java b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionBackpressureTest.java new file mode 100644 index 00000000..3596b532 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionBackpressureTest.java @@ -0,0 +1,277 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.*; + +import io.netty.channel.Channel; +import io.netty.channel.ChannelConfig; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.traffic.GlobalTrafficShapingHandler; +import java.lang.reflect.Field; +import java.util.concurrent.ConcurrentMap; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ChainedProxyManager; +import org.littleshoot.proxy.HttpFiltersSource; + +class ClientToProxyConnectionBackpressureTest { + + private DefaultHttpProxyServer mockProxyServer; + private GlobalTrafficShapingHandler mockTrafficHandler; + private ClientToProxyConnection clientConn; + private EmbeddedChannel clientChannel; + + private ProxyToServerConnection mockServerInMap; + private ProxyToServerConnection mockCurrentServer; + + @BeforeEach + void setUp() throws Exception { + mockProxyServer = mock(); + mockTrafficHandler = mock(); + when(mockProxyServer.getChainProxyManager()).thenReturn(mock(ChainedProxyManager.class)); + when(mockProxyServer.getFiltersSource()).thenReturn(mock(HttpFiltersSource.class)); + when(mockProxyServer.getMaxInitialLineLength()).thenReturn(8192); + when(mockProxyServer.getMaxHeaderSize()).thenReturn(16384); + when(mockProxyServer.getMaxChunkSize()).thenReturn(16384); + when(mockProxyServer.getIdleConnectionTimeout()).thenReturn(0); + when(mockProxyServer.isAcceptProxyProtocol()).thenReturn(false); + when(mockProxyServer.getProxyAlias()).thenReturn("test"); + when(mockProxyServer.isAllowRequestsToOriginServer()).thenReturn(true); + when(mockProxyServer.getActivityTrackers()).thenReturn(java.util.Collections.emptyList()); + + clientChannel = new EmbeddedChannel(); + clientConn = + new ClientToProxyConnection( + mockProxyServer, null, false, clientChannel.pipeline(), mockTrafficHandler); + + // Create mock server connections + mockServerInMap = createServerConnectionMock(); + mockCurrentServer = createServerConnectionMock(); + + // Populate serverConnectionsByHostAndPort (private field, set via reflection since + // there is no public add method — connections are added internally during request flow) + @SuppressWarnings("unchecked") + ConcurrentMap serverMap = + (ConcurrentMap) + field(ClientToProxyConnection.class, "serverConnectionsByHostAndPort").get(clientConn); + serverMap.put("example.com:80", mockServerInMap); + + // Set currentServerConnection (private field, no public setter available) + field(ClientToProxyConnection.class, "currentServerConnection") + .set(clientConn, mockCurrentServer); + } + + /** + * Creates a mock ProxyToServerConnection with a proper channel so stopReading/resumeReading work. + */ + private static ProxyToServerConnection createServerConnectionMock() throws Exception { + ProxyToServerConnection conn = mock(); + Channel ch = mock(); + ChannelConfig cfg = mock(); + when(ch.config()).thenReturn(cfg); + field(ProxyConnection.class, "channel").set(conn, ch); + return conn; + } + + // ----------------------------------------------------------------------- + // becameSaturated + // ----------------------------------------------------------------------- + + @Test + @DisplayName( + "becameSaturated should stop reading on currentServerConnection when client is saturated") + void becameSaturatedShouldStopReadingOnCurrentServerConnection() throws Exception { + mockClientChannelNotWritable(); + clientConn.becameSaturated(); + verify(mockCurrentServer).stopReading(); + } + + @Test + @DisplayName( + "becameSaturated should stop reading on mapped server connections when client is saturated") + void becameSaturatedShouldStopReadingOnAllServerConnections() throws Exception { + mockClientChannelNotWritable(); + clientConn.becameSaturated(); + verify(mockServerInMap).stopReading(); + } + + @Test + @DisplayName("becameSaturated should NOT stop reading when client is not saturated") + void becameSaturatedShouldNotStopReadingWhenNotSaturated() { + // clientChannel is writable by default → isSaturated() returns false + clientConn.becameSaturated(); + verify(mockCurrentServer, never()).stopReading(); + verify(mockServerInMap, never()).stopReading(); + } + + @Test + @DisplayName("becameSaturated should handle null currentServerConnection (pooled connection)") + void becameSaturatedShouldHandleNullCurrentServerConnection() throws Exception { + field(ClientToProxyConnection.class, "currentServerConnection").set(clientConn, null); + mockClientChannelNotWritable(); + clientConn.becameSaturated(); + // Should not throw NPE; the mapped connection should still be stopped + verify(mockServerInMap).stopReading(); + } + + // ----------------------------------------------------------------------- + // becameWritable + // ----------------------------------------------------------------------- + + @Test + @DisplayName("becameWritable should resume reading on currentServerConnection") + void becameWritableShouldResumeReadingOnCurrentServerConnection() throws Exception { + clientConn.becameWritable(); + verify(mockCurrentServer).resumeReading(); + } + + @Test + @DisplayName("becameWritable should resume reading on mapped server connections") + void becameWritableShouldResumeReadingOnMappedConnections() throws Exception { + clientConn.becameWritable(); + verify(mockServerInMap).resumeReading(); + } + + @Test + @DisplayName("becameWritable should NOT resume reading when client is still saturated") + void becameWritableShouldNotResumeReadingWhenSaturated() throws Exception { + mockClientChannelNotWritable(); + clientConn.becameWritable(); + verify(mockCurrentServer, never()).resumeReading(); + verify(mockServerInMap, never()).resumeReading(); + } + + @Test + @DisplayName("becameWritable should handle null currentServerConnection") + void becameWritableShouldHandleNullCurrentServerConnection() throws Exception { + field(ClientToProxyConnection.class, "currentServerConnection").set(clientConn, null); + clientConn.becameWritable(); + verify(mockServerInMap).resumeReading(); + } + + // ----------------------------------------------------------------------- + // serverBecameSaturated + // ----------------------------------------------------------------------- + + @Test + @DisplayName("serverBecameSaturated should stop client reading when server is saturated") + void serverBecameSaturatedShouldStopClientWhenServerSaturated() { + clientChannel.config().setAutoRead(true); + when(mockCurrentServer.isSaturated()).thenReturn(true); + + clientConn.serverBecameSaturated(mockCurrentServer); + + assertThat(clientChannel.config().isAutoRead()).isFalse(); + } + + @Test + @DisplayName("serverBecameSaturated should not stop client reading when server is not saturated") + void serverBecameSaturatedShouldNotStopClientWhenServerNotSaturated() { + clientChannel.config().setAutoRead(true); + when(mockCurrentServer.isSaturated()).thenReturn(false); + + clientConn.serverBecameSaturated(mockCurrentServer); + + assertThat(clientChannel.config().isAutoRead()).isTrue(); + } + + // ----------------------------------------------------------------------- + // serverBecameWriteable + // ----------------------------------------------------------------------- + + @Test + @DisplayName("serverBecameWriteable should resume client reading when no servers are saturated") + void serverBecameWriteableShouldResumeWhenNoServerSaturated() { + clientChannel.config().setAutoRead(false); + when(mockServerInMap.isSaturated()).thenReturn(false); + when(mockCurrentServer.isSaturated()).thenReturn(false); + + clientConn.serverBecameWriteable(mockCurrentServer); + + assertThat(clientChannel.config().isAutoRead()).isTrue(); + } + + @Test + @DisplayName("serverBecameWriteable should not resume when a mapped server is still saturated") + void serverBecameWriteableShouldNotResumeWhenMappedSaturated() { + clientChannel.config().setAutoRead(false); + when(mockServerInMap.isSaturated()).thenReturn(true); + when(mockCurrentServer.isSaturated()).thenReturn(false); + + clientConn.serverBecameWriteable(mockCurrentServer); + + assertThat(clientChannel.config().isAutoRead()).isFalse(); + } + + @Test + @DisplayName( + "serverBecameWriteable should check currentServerConnection when " + + "it is a different connection than the one that became writeable") + void serverBecameWriteableShouldCheckCurrentServerConnection() throws Exception { + clientChannel.config().setAutoRead(false); + // The connection that became writeable is NOT currentServerConnection + ProxyToServerConnection differentServer = createServerConnectionMock(); + when(mockServerInMap.isSaturated()).thenReturn(false); + when(mockCurrentServer.isSaturated()).thenReturn(true); + + clientConn.serverBecameWriteable(differentServer); + + // Should NOT resume because currentServerConnection is still saturated + assertThat(clientChannel.config().isAutoRead()).isFalse(); + } + + @Test + @DisplayName("serverBecameWriteable should resume when currentServerConnection is the source") + void serverBecameWriteableShouldResumeWhenCurrentIsSource() { + clientChannel.config().setAutoRead(false); + when(mockServerInMap.isSaturated()).thenReturn(false); + when(mockCurrentServer.isSaturated()).thenReturn(false); + + clientConn.serverBecameWriteable(mockCurrentServer); + + assertThat(clientChannel.config().isAutoRead()).isTrue(); + } + + @Test + @DisplayName("serverBecameWriteable should handle null currentServerConnection") + void serverBecameWriteableShouldHandleNullCurrentServerConnection() throws Exception { + clientChannel.config().setAutoRead(false); + field(ClientToProxyConnection.class, "currentServerConnection").set(clientConn, null); + when(mockServerInMap.isSaturated()).thenReturn(false); + + clientConn.serverBecameWriteable(mockCurrentServer); + + assertThat(clientChannel.config().isAutoRead()).isTrue(); + } + + // ----------------------------------------------------------------------- + // Helpers + // ----------------------------------------------------------------------- + + /** + * Replaces the client channel with a non-writable mock to make isSaturated() return true. The + * EmbeddedChannel always reports isWritable() = true, so a mock is required here. + */ + private void mockClientChannelNotWritable() throws Exception { + Channel ch = mock(); + ChannelConfig cfg = mock(); + when(ch.isWritable()).thenReturn(false); + when(ch.config()).thenReturn(cfg); + field(ProxyConnection.class, "channel").set(clientConn, ch); + } + + private static Field field(Class clazz, String name) throws NoSuchFieldException { + Class current = clazz; + while (current != null) { + try { + Field f = current.getDeclaredField(name); + f.setAccessible(true); + return f; + } catch (NoSuchFieldException e) { + current = current.getSuperclass(); + } + } + throw new NoSuchFieldException(name + " in " + clazz.getName()); + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionShortCircuitTest.java b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionShortCircuitTest.java new file mode 100644 index 00000000..0966ff09 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionShortCircuitTest.java @@ -0,0 +1,161 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.*; + +import io.netty.channel.Channel; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.codec.http.DefaultHttpRequest; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpVersion; +import io.netty.handler.traffic.GlobalTrafficShapingHandler; +import java.lang.reflect.Field; +import java.net.InetSocketAddress; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.HttpFiltersSource; + +class ClientToProxyConnectionShortCircuitTest { + + private DefaultHttpProxyServer mockProxyServer; + private ClientToProxyConnection clientConn; + private EmbeddedChannel clientChannel; + private GlobalTrafficShapingHandler mockTrafficHandler; + + @BeforeEach + void setUp() throws Exception { + mockProxyServer = mock(); + mockTrafficHandler = mock(); + + when(mockProxyServer.getChainProxyManager()).thenReturn(null); + when(mockProxyServer.getServerResolver()) + .thenReturn(mock(org.littleshoot.proxy.HostResolver.class)); + when(mockProxyServer.getServerResolver().resolve(anyString(), anyInt())) + .thenReturn(new InetSocketAddress("127.0.0.1", 8080)); + when(mockProxyServer.getFiltersSource()).thenReturn(mock(HttpFiltersSource.class)); + when(mockProxyServer.getMaxInitialLineLength()).thenReturn(8192); + when(mockProxyServer.getMaxHeaderSize()).thenReturn(16384); + when(mockProxyServer.getMaxChunkSize()).thenReturn(16384); + when(mockProxyServer.getIdleConnectionTimeout()).thenReturn(0); + when(mockProxyServer.isAcceptProxyProtocol()).thenReturn(false); + when(mockProxyServer.getProxyAlias()).thenReturn("test"); + when(mockProxyServer.isAllowRequestsToOriginServer()).thenReturn(true); + when(mockProxyServer.getActivityTrackers()).thenReturn(java.util.Collections.emptyList()); + + clientChannel = new EmbeddedChannel(); + clientConn = + new ClientToProxyConnection( + mockProxyServer, null, false, clientChannel.pipeline(), mockTrafficHandler); + } + + // ----------------------------------------------------------------------- + // setCurrentClientConnectionForRequest(null) — covers the short-circuit path + // ----------------------------------------------------------------------- + + @Test + @DisplayName("setCurrentClientConnectionForRequest(null) should be callable") + void setCurrentClientConnectionForRequestShouldAcceptNull() throws Exception { + // Create a real ProxyToServerConnection to test the setter on + ProxyToServerConnection conn = createRealPooledConnection(); + + // This is the exact call made at ClientToProxyConnection line 393 when + // proxyToServerRequest short-circuits with a shared pool + conn.setCurrentClientConnectionForRequest(null); + + // Verify the field was set to null + Field field = + ProxyToServerConnection.class.getDeclaredField("currentClientConnectionForRequest"); + field.setAccessible(true); + assertThat(field.get(conn)).isNull(); + } + + @Test + @DisplayName("releaseToPool after setCurrentClientConnectionForRequest(null) should not throw") + void releaseToPoolAfterNullShouldNotThrow() throws Exception { + ProxyToServerConnection conn = createRealPooledConnection(); + + // Same sequence as ClientToProxyConnection lines 392-395 + conn.setCurrentClientConnectionForRequest(null); + conn.releaseToPool(); + + // Verify the connection was released back (currentHttpRequest cleared) + Field reqField = ProxyToServerConnection.class.getDeclaredField("currentHttpRequest"); + reqField.setAccessible(true); + assertThat(reqField.get(conn)).isNull(); + + Field clientField = + ProxyToServerConnection.class.getDeclaredField("currentClientConnectionForRequest"); + clientField.setAccessible(true); + assertThat(clientField.get(conn)).isNull(); + } + + // ----------------------------------------------------------------------- + // getClientAddress null-safety + // ----------------------------------------------------------------------- + + @Test + @DisplayName("getClientAddress should return null when channel is null") + void getClientAddressShouldBeNullWhenChannelNull() throws Exception { + setField(clientConn, "channel", null); + assertThat(clientConn.getClientAddress()).isNull(); + } + + @Test + @DisplayName("getClientAddress should return null for non-InetSocketAddress remote address") + void getClientAddressShouldBeNullForNonInetSocketAddress() { + // EmbeddedSocketAddress is NOT an InetSocketAddress + assertThat(clientConn.getClientAddress()).isNull(); + } + + @Test + @DisplayName("getClientAddress should return InetSocketAddress when present") + void getClientAddressShouldReturnInetSocketAddress() throws Exception { + Channel ch = mock(); + when(ch.remoteAddress()).thenReturn(new InetSocketAddress("192.168.1.1", 12345)); + setField(clientConn, "channel", ch); + + InetSocketAddress addr = clientConn.getClientAddress(); + assertThat(addr).isNotNull(); + assertThat(addr.getHostString()).isEqualTo("192.168.1.1"); + assertThat(addr.getPort()).isEqualTo(12345); + } + + // ----------------------------------------------------------------------- + // Helpers + // ----------------------------------------------------------------------- + + /** Creates a minimal ProxyToServerConnection to test setter calls on. */ + private ProxyToServerConnection createRealPooledConnection() throws Exception { + ClientToProxyConnection mockClient = mock(); + when(mockClient.flowContext()).thenReturn(mock()); + when(mockClient.flowContextForServerConnection(any(ProxyToServerConnection.class))) + .thenReturn(mock()); + return ProxyToServerConnection.createForPool( + mockProxyServer, + mock(ServerConnectionPool.class), + mockClient, + "example.com:80", + null, + mock(), + new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"), + mockTrafficHandler); + } + + private static void setField(Object obj, String name, Object value) throws Exception { + Field f = findField(obj.getClass(), name); + f.setAccessible(true); + f.set(obj, value); + } + + private static Field findField(Class clazz, String name) throws NoSuchFieldException { + try { + return clazz.getDeclaredField(name); + } catch (NoSuchFieldException e) { + if (clazz.getSuperclass() != null) { + return findField(clazz.getSuperclass(), name); + } + throw e; + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionTest.java b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionTest.java new file mode 100644 index 00000000..8f00634a --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyConnectionTest.java @@ -0,0 +1,107 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import io.netty.channel.ChannelPipeline; +import io.netty.channel.embedded.EmbeddedChannel; +import java.util.List; +import javax.net.ssl.SSLContext; +import javax.net.ssl.SSLEngine; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ActivityTracker; +import org.littleshoot.proxy.HttpFiltersSource; + +class ClientToProxyConnectionTest { + + private final DefaultHttpProxyServer proxyServer = mock(); + + private ClientToProxyConnection createConnection(ActivityTracker... trackers) { + when(proxyServer.getMaxInitialLineLength()).thenReturn(4096); + when(proxyServer.getMaxHeaderSize()).thenReturn(8192); + when(proxyServer.getMaxChunkSize()).thenReturn(8192); + when(proxyServer.getIdleConnectionTimeout()).thenReturn(60); + when(proxyServer.getFiltersSource()).thenReturn(mock(HttpFiltersSource.class)); + when(proxyServer.getActivityTrackers()).thenReturn(List.of(trackers)); + + return new ClientToProxyConnection(proxyServer, null, false, mock(ChannelPipeline.class), null); + } + + @Test + @DisplayName("disconnected should notify all trackers even if one throws") + void disconnectedShouldNotifyAllTrackersEvenIfOneThrows() { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")) + .when(throwingTracker) + .clientDisconnected(any(), any()); + ActivityTracker normalTracker = mock(); + ClientToProxyConnection connection = createConnection(throwingTracker, normalTracker); + + connection.disconnected(); + + verify(throwingTracker).clientDisconnected(connection.flowContext(), null); + verify(normalTracker).clientDisconnected(connection.flowContext(), null); + } + + @Test + @DisplayName("disconnected should notify trackers when no exception occurs") + void disconnectedShouldNotifyTrackersWhenNoException() { + ActivityTracker normalTracker = mock(); + ClientToProxyConnection connection = createConnection(normalTracker); + + connection.disconnected(); + + verify(normalTracker).clientDisconnected(connection.flowContext(), null); + } + + @Test + @DisplayName("encrypt requires a client certificate when authenticateClients is true") + void encryptRequiresClientCertificateWhenAuthenticateClientsIsTrue() throws Exception { + ClientToProxyConnection connection = createConnection(); + SSLEngine engine = newServerEngine(); + + connection.encrypt(newRealPipeline(), engine, true); + + assertThat(engine.getNeedClientAuth()).isTrue(); + } + + @Test + @DisplayName("encrypt leaves a plain engine unauthenticated when authenticateClients is false") + void encryptLeavesPlainEngineUnauthenticatedWhenAuthenticateClientsIsFalse() throws Exception { + ClientToProxyConnection connection = createConnection(); + SSLEngine engine = newServerEngine(); // default engine: neither need nor want + + connection.encrypt(newRealPipeline(), engine, false); + + assertThat(engine.getNeedClientAuth()).isFalse(); + assertThat(engine.getWantClientAuth()).isFalse(); + } + + @Test + @DisplayName( + "encrypt preserves setWantClientAuth from the SslEngineSource when authenticateClients is false") + void encryptPreservesWantClientAuthWhenAuthenticateClientsIsFalse() throws Exception { + ClientToProxyConnection connection = createConnection(); + SSLEngine engine = newServerEngine(); + engine.setWantClientAuth(true); // e.g. an engine built with Netty ClientAuth.OPTIONAL + + connection.encrypt(newRealPipeline(), engine, false); + + // Before the fix, encrypt() called setNeedClientAuth(false) here, which cleared this flag. + assertThat(engine.getWantClientAuth()).isTrue(); + assertThat(engine.getNeedClientAuth()).isFalse(); + } + + private static SSLEngine newServerEngine() throws Exception { + return SSLContext.getDefault().createSSLEngine(); + } + + private static ChannelPipeline newRealPipeline() { + return new EmbeddedChannel().pipeline(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ClientToProxyTimeoutBugTest.java b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyTimeoutBugTest.java new file mode 100644 index 00000000..8a41cb73 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ClientToProxyTimeoutBugTest.java @@ -0,0 +1,322 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.*; + +import io.netty.channel.Channel; +import io.netty.channel.ChannelConfig; +import io.netty.channel.ChannelFuture; +import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.ChannelPipeline; +import io.netty.channel.EventLoop; +import java.lang.reflect.Field; +import java.lang.reflect.Method; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +/** + * Unit test for issue #439: ClientToProxyConnection.timedOut() bug. + * + *

The bug: When a server connection has been created but has never read any data (lastReadTime + * == 0), the idle timeout check incorrectly prevents the client connection from being closed. + * + *

Original buggy code in ClientToProxyConnection.timedOut(): + * + *

+ * if (currentServerConnection == null || lastReadTime <= currentServerConnection.lastReadTime) {
+ *     super.timedOut();
+ * }
+ * 
+ * + *

When currentServerConnection.lastReadTime == 0 and client.lastReadTime > 0: - Condition: + * lastReadTime <= 0 evaluates to FALSE - So super.timedOut() is NOT called - BUG: The idle client + * connection should be closed! + * + *

The fix adds additional checks to avoid closing when a request is pending: + * + *

+ * boolean requestHasBeenWritten = false;
+ * if (currentServerConnection != null) {
+ *     HttpRequest initialRequest = currentServerConnection.getInitialRequest();
+ *     requestHasBeenWritten = initialRequest != null;
+ * }
+ *
+ * if (currentServerConnection == null
+ *     || (currentServerConnection.lastReadTime == 0
+ *         && !requestHasBeenWritten
+ *         && currentRequest == null)
+ *     || lastReadTime <= currentServerConnection.lastReadTime) {
+ *     super.timedOut();
+ * }
+ * 
+ * + *

This test MUST FAIL with the current buggy code and MUST PASS when the fix is applied. + * + * @see Issue #439 + */ +public class ClientToProxyTimeoutBugTest { + + private ChannelHandlerContext ctx; + private Channel channel; + private ChannelPipeline pipeline; + private ChannelConfig channelConfig; + private EventLoop eventLoop; + private DefaultHttpProxyServer proxyServer; + + private ClientToProxyConnection clientToProxyConnection; + + /** Sets up the test by creating a ClientToProxyConnection with mocked dependencies. */ + @BeforeEach + public void setUp() throws Exception { + // Create mocks manually + ctx = mock(ChannelHandlerContext.class); + channel = mock(Channel.class); + pipeline = mock(ChannelPipeline.class); + channelConfig = mock(ChannelConfig.class); + eventLoop = mock(EventLoop.class); + proxyServer = mock(DefaultHttpProxyServer.class); + + // Mock channel and its components + when(channel.pipeline()).thenReturn(pipeline); + when(channel.config()).thenReturn(channelConfig); + when(ctx.channel()).thenReturn(channel); + when(ctx.pipeline()).thenReturn(pipeline); + when(channel.eventLoop()).thenReturn(eventLoop); + + // Mock proxyServer to return valid configuration values for pipeline initialization + when(proxyServer.getMaxInitialLineLength()).thenReturn(4096); + when(proxyServer.getMaxHeaderSize()).thenReturn(8192); + when(proxyServer.getMaxChunkSize()).thenReturn(8192); + when(proxyServer.getIdleConnectionTimeout()).thenReturn(60); + when(proxyServer.isAcceptProxyProtocol()).thenReturn(false); + when(proxyServer.isTransparent()).thenReturn(false); + when(proxyServer.getProxyAlias()).thenReturn("test-proxy"); + when(proxyServer.getFiltersSource()) + .thenReturn(mock(org.littleshoot.proxy.HttpFiltersSource.class)); + + // Create ClientToProxyConnection - constructor is package-private + clientToProxyConnection = + new ClientToProxyConnection( + proxyServer, + null, // no SSL + false, // don't authenticate clients + pipeline, + null // no traffic shaping handler + ); + + // Inject the mocked context using reflection + Field ctxField = ProxyConnection.class.getDeclaredField("ctx"); + ctxField.setAccessible(true); + ctxField.set(clientToProxyConnection, ctx); + + // Inject the mocked channel using reflection + Field channelField = ProxyConnection.class.getDeclaredField("channel"); + channelField.setAccessible(true); + channelField.set(clientToProxyConnection, channel); + } + + /** + * This test directly verifies the bug in ClientToProxyConnection.timedOut(). + * + *

Scenario: - Client has read data (lastReadTime > 0) - Server connection exists but has never + * read any data (lastReadTime == 0) - No request has been written to the server + * (requestHasBeenWritten = false) - No pending request from client (currentRequest = null) - Idle + * timeout fires on the client channel + * + *

Expected behavior (with fix): - The condition: server.lastReadTime == 0 && + * !requestHasBeenWritten && currentRequest == null - This evaluates to: true && true && true = + * TRUE - So super.timedOut() SHOULD be called to close the idle client connection + * + *

Actual behavior (with bug): - The condition: lastReadTime <= 0 evaluates to FALSE - So + * super.timedOut() is NOT called - The idle client connection is NOT closed (BUG!) + * + *

This test MUST FAIL with the current buggy code and MUST PASS when the fix is applied. + */ + @Test + public void testTimedOut_WhenServerNeverRead_ShouldCallSuperTimedOut() throws Exception { + // Set up the scenario that triggers the bug: + // 1. Client has read data (lastReadTime > 0) + // 2. Server connection exists but has never read (lastReadTime == 0) + + // Set client.lastReadTime > 0 using reflection + long clientLastReadTime = System.currentTimeMillis(); + Field lastReadTimeField = ProxyConnection.class.getDeclaredField("lastReadTime"); + lastReadTimeField.setAccessible(true); + lastReadTimeField.set(clientToProxyConnection, clientLastReadTime); + + // Create a mock ProxyToServerConnection with lastReadTime = 0 + ProxyToServerConnection mockServerConnection = mock(ProxyToServerConnection.class); + + // Use reflection to set lastReadTime = 0 on the mock server connection + Field serverLastReadTimeField = ProxyConnection.class.getDeclaredField("lastReadTime"); + serverLastReadTimeField.setAccessible(true); + serverLastReadTimeField.set(mockServerConnection, 0L); + + // Set currentServerConnection using reflection + Field currentServerConnectionField = + ClientToProxyConnection.class.getDeclaredField("currentServerConnection"); + currentServerConnectionField.setAccessible(true); + currentServerConnectionField.set(clientToProxyConnection, mockServerConnection); + + // Use Mockito spy to verify that super.timedOut() (which calls disconnect()) is called + ClientToProxyConnection spyConnection = spy(clientToProxyConnection); + + // Stub disconnect to avoid actual channel operations - return a mock future + ChannelFuture mockFuture = mock(ChannelFuture.class); + doReturn(null).when(spyConnection).disconnect(); + + // Call timedOut() on the spy + Method timedOutMethod = ClientToProxyConnection.class.getDeclaredMethod("timedOut"); + timedOutMethod.setAccessible(true); + timedOutMethod.invoke(spyConnection); + + // VERIFICATION: + // With the FIX: disconnect() SHOULD be called because: + // - server.lastReadTime == 0 && !requestHasBeenWritten && currentRequest == null + // - evaluates to: true && true && true = TRUE + // + // With the BUG: disconnect() is NOT called because: + // - lastReadTime <= 0 evaluates to FALSE + // - so the whole condition is FALSE + + // This assertion will FAIL with the buggy code because disconnect() is NOT called + // when server.lastReadTime == 0 and no request has been sent + // With the fix, the condition evaluates to true so disconnect() IS called + try { + verify(spyConnection, times(1)).disconnect(); + } catch (AssertionError e) { + throw new AssertionError( + "BUG DETECTED: super.timedOut() was NOT called when it SHOULD have been!\n" + + "When server.lastReadTime == 0 and client.lastReadTime > 0,\n" + + "and no request has been sent to the server,\n" + + "the idle timeout condition incorrectly evaluates to false.\n" + + "The fix should check: server.lastReadTime == 0 && !requestHasBeenWritten && currentRequest == null\n" + + "See issue #439", + e); + } + } + + /** + * Additional test to verify the behavior when currentServerConnection is null. This should always + * call super.timedOut() - this works correctly even with the bug. + * + *

When currentServerConnection is null, the first part of the OR condition is true: + * "currentServerConnection == null" evaluates to TRUE So disconnect() is called regardless of the + * fix. + */ + @Test + public void testTimedOut_WhenNoServerConnection_ShouldCallSuperTimedOut() throws Exception { + // Ensure currentServerConnection is null + Field currentServerConnectionField = + ClientToProxyConnection.class.getDeclaredField("currentServerConnection"); + currentServerConnectionField.setAccessible(true); + currentServerConnectionField.set(clientToProxyConnection, null); + + // Set client.lastReadTime > 0 + long clientLastReadTime = System.currentTimeMillis(); + Field lastReadTimeField = ProxyConnection.class.getDeclaredField("lastReadTime"); + lastReadTimeField.setAccessible(true); + lastReadTimeField.set(clientToProxyConnection, clientLastReadTime); + + // Use Mockito spy to verify disconnect() is called + ClientToProxyConnection spyConnection = spy(clientToProxyConnection); + doReturn(null).when(spyConnection).disconnect(); + + // Call timedOut() using reflection + Method timedOutMethod = ClientToProxyConnection.class.getDeclaredMethod("timedOut"); + timedOutMethod.setAccessible(true); + timedOutMethod.invoke(spyConnection); + + // When currentServerConnection is null, the first part of the OR condition is true + // so disconnect() should be called - this works correctly even with the bug + verify(spyConnection, times(1)).disconnect(); + } + + /** + * Test to verify the behavior when server.lastReadTime > client.lastReadTime. + * + *

In this case, the server has been more active (sent data more recently). The condition + * "lastReadTime <= currentServerConnection.lastReadTime" evaluates to TRUE (e.g., 1000 <= 2000), + * so disconnect() SHOULD be called. + * + *

This is correct behavior and works both with and without the fix. + */ + @Test + public void testTimedOut_WhenServerMoreActive_ShouldCallSuperTimedOut() throws Exception { + // Set client.lastReadTime = 1000 (less recent) + Field lastReadTimeField = ProxyConnection.class.getDeclaredField("lastReadTime"); + lastReadTimeField.setAccessible(true); + lastReadTimeField.set(clientToProxyConnection, 1000L); + + // Create mock server connection with lastReadTime = 2000 (more recent) + ProxyToServerConnection mockServerConnection = mock(ProxyToServerConnection.class); + Field serverLastReadTimeField = ProxyConnection.class.getDeclaredField("lastReadTime"); + serverLastReadTimeField.setAccessible(true); + serverLastReadTimeField.set(mockServerConnection, 2000L); + + // Set currentServerConnection + Field currentServerConnectionField = + ClientToProxyConnection.class.getDeclaredField("currentServerConnection"); + currentServerConnectionField.setAccessible(true); + currentServerConnectionField.set(clientToProxyConnection, mockServerConnection); + + // Use Mockito spy to verify disconnect() is called + ClientToProxyConnection spyConnection = spy(clientToProxyConnection); + doReturn(null).when(spyConnection).disconnect(); + + // Call timedOut() using reflection + Method timedOutMethod = ClientToProxyConnection.class.getDeclaredMethod("timedOut"); + timedOutMethod.setAccessible(true); + timedOutMethod.invoke(spyConnection); + + // When server.lastReadTime > client.lastReadTime, the condition + // lastReadTime <= serverLastReadTime evaluates to true (1000 <= 2000) + // so disconnect() SHOULD be called - this is correct behavior + verify(spyConnection, times(1)).disconnect(); + } + + /** + * Demonstrates the bug condition logic mathematically. + * + *

This test shows exactly what happens with the buggy condition vs the fixed condition. Note: + * The actual fix also checks requestHasBeenWritten and currentRequest, but this test focuses on + * the core logic issue. + */ + @Test + public void demonstrateBugConditionLogic() { + // BUG DEMONSTRATION: + // When server.lastReadTime == 0 and client.lastReadTime > 0 + // and no request has been sent to server + + long clientLastRead = 1000L; // Client has read something + long serverLastRead = 0L; // Server has NEVER read (new connection) + boolean requestHasBeenWritten = false; // No request sent to server + Object currentRequest = null; // No pending request + + // Current buggy condition (simplified): + // lastReadTime <= currentServerConnection.lastReadTime + boolean buggyCondition = (clientLastRead <= serverLastRead); + + // The buggy condition evaluates to: 1000 <= 0 = FALSE + assertThat(buggyCondition).isFalse(); + + // This means super.timedOut() is NOT called - BUG! + // The client connection should be closed but it's not! + + // Fixed condition (actual fix): + // (server.lastReadTime == 0 && !requestHasBeenWritten && currentRequest == null) + // || lastReadTime <= server.lastReadTime + boolean fixedCondition = + (serverLastRead == 0 && !requestHasBeenWritten && currentRequest == null) + || (clientLastRead <= serverLastRead); + + // The fixed condition evaluates to: (true && true && true) || false = TRUE + assertThat(fixedCondition).isTrue(); + + // This means super.timedOut() IS called - CORRECT! + + // This demonstrates the bug clearly: + // - With buggy code: timeout handling is SKIPPED (wrong!) + // - With fixed code: timeout handling is EXECUTED (correct!) + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ConcurrentMapServerConnectionPoolTest.java b/src/test/java/org/littleshoot/proxy/impl/ConcurrentMapServerConnectionPoolTest.java new file mode 100644 index 00000000..1d6ad678 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ConcurrentMapServerConnectionPoolTest.java @@ -0,0 +1,475 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyInt; +import static org.mockito.Mockito.*; + +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.codec.http.DefaultHttpRequest; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpVersion; +import io.netty.handler.traffic.GlobalTrafficShapingHandler; +import java.lang.reflect.Field; +import java.net.InetSocketAddress; +import java.util.Queue; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentMap; +import java.util.concurrent.ScheduledExecutorService; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Tag; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.HostResolver; +import org.littleshoot.proxy.HttpFilters; +import org.littleshoot.proxy.HttpFiltersAdapter; + +class ConcurrentMapServerConnectionPoolTest { + + private DefaultHttpProxyServer mockProxyServer; + private GlobalTrafficShapingHandler mockTrafficHandler; + private ClientToProxyConnection mockClientConnection; + private HostResolver mockHostResolver; + private ConcurrentMapServerConnectionPool pool; + + @BeforeEach + void setUp() throws Exception { + mockProxyServer = mock(); + mockTrafficHandler = mock(); + mockClientConnection = mock(); + mockHostResolver = mock(); + + when(mockProxyServer.getChainProxyManager()).thenReturn(null); + when(mockProxyServer.getServerResolver()).thenReturn(mockHostResolver); + when(mockHostResolver.resolve(anyString(), anyInt())) + .thenReturn(new InetSocketAddress("127.0.0.1", 8080)); + when(mockProxyServer.getActivityTrackers()).thenReturn(java.util.Collections.emptyList()); + + when(mockClientConnection.flowContext()).thenReturn(mock()); + when(mockClientConnection.flowContextForServerConnection(any(ProxyToServerConnection.class))) + .thenReturn(mock()); + + pool = new ConcurrentMapServerConnectionPool(mockProxyServer, mockTrafficHandler); + } + + @AfterEach + void tearDown() { + pool.closeAll(); + } + + // ----------------------------------------------------------------------- + // Helpers + // ----------------------------------------------------------------------- + + /** Creates a mock ProxyToServerConnection with the given availability. */ + private ProxyToServerConnection createMockConnection( + boolean connected, boolean availableForNewRequest) throws Exception { + ProxyToServerConnection conn = mock(); + when(conn.isConnected()).thenReturn(connected); + when(conn.isAvailableForNewRequest()).thenReturn(availableForNewRequest); + when(conn.getServerHostAndPort()).thenReturn("example.com:80"); + return conn; + } + + /** + * Registers a connection in the pool so it appears as a tracked, connected connection. After + * this, the pool's {@code releaseConnection} can be called to add it to the available queue. + */ + @SuppressWarnings("unchecked") + private void registerInPool(ProxyToServerConnection conn, String poolKey) throws Exception { + + ConcurrentMap keys = getField(pool, "connectionKeys"); + keys.put(conn, poolKey); + + ConcurrentMap> connectionsByHost = + getField(pool, "connectionsByHostAndPort"); + connectionsByHost + .computeIfAbsent(poolKey, k -> new ConcurrentHashMap<>()) + .put(conn, Boolean.TRUE); + + ConcurrentMap counts = + getField(pool, "connectionCountByHostAndPort"); + counts + .computeIfAbsent(poolKey, k -> new java.util.concurrent.atomic.AtomicInteger(0)) + .incrementAndGet(); + } + + private static Field field(Class clazz, String name) throws Exception { + Class current = clazz; + while (current != null) { + try { + Field f = current.getDeclaredField(name); + f.setAccessible(true); + return f; + } catch (NoSuchFieldException e) { + current = current.getSuperclass(); + } + } + throw new NoSuchFieldException(name + " in " + clazz.getName()); + } + + @SuppressWarnings("unchecked") + private static T getField(Object obj, String name) throws Exception { + Field f = field(obj.getClass(), name); + return (T) f.get(obj); + } + + // ----------------------------------------------------------------------- + // PendingRequest tests (from original test, no reflection) + // ----------------------------------------------------------------------- + + @Test + void pendingRequestShouldStoreDataCorrectly() { + HttpRequest request = new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"); + PendingRequest pendingRequest = new PendingRequest(null, request, null); + assertThat(pendingRequest.getClientConnection()).isNull(); + assertThat(pendingRequest.getRequest()).isSameAs(request); + assertThat(pendingRequest.getFilters()).isNull(); + assertThat(pendingRequest.getTimestamp()).isGreaterThan(0); + } + + @Test + @Tag("slow-test") + void pendingRequestTimestampShouldBeRecent() { + long before = System.currentTimeMillis(); + PendingRequest pendingRequest = new PendingRequest(null, null, null); + long after = System.currentTimeMillis(); + assertThat(pendingRequest.getTimestamp()).isGreaterThanOrEqualTo(before); + assertThat(pendingRequest.getTimestamp()).isLessThanOrEqualTo(after); + } + + @Test + void poolShouldHaveDefaultMaxConnectionsPerHost() { + assertThat(ConcurrentMapServerConnectionPool.DEFAULT_MAX_CONNECTIONS_PER_HOST).isEqualTo(10); + } + + @Test + void poolShouldHaveDefaultMaxTotalConnections() { + assertThat(ConcurrentMapServerConnectionPool.DEFAULT_MAX_TOTAL_CONNECTIONS).isEqualTo(200); + } + + // ----------------------------------------------------------------------- + // releaseConnection (public API) + // ----------------------------------------------------------------------- + + @Test + @DisplayName("releaseConnection should add connected connection to available queue") + void releaseConnectionShouldAddToAvailableQueue() throws Exception { + ProxyToServerConnection conn = createMockConnection(true, true); + registerInPool(conn, "example.com:80:direct"); + + pool.releaseConnection(conn); + + Queue queue = + ((java.util.Map>) getField(pool, "availableConnectionsByHostAndPort")) + .get("example.com:80:direct"); + assertThat(queue).isNotNull().hasSize(1); + } + + @Test + @DisplayName("releaseConnection should call removeConnection for disconnected connections") + void releaseConnectionShouldRemoveDisconnectedConnections() throws Exception { + ProxyToServerConnection conn = createMockConnection(false, false); + registerInPool(conn, "example.com:80:direct"); + + pool.releaseConnection(conn); + + ConcurrentMap keys = getField(pool, "connectionKeys"); + assertThat(keys).doesNotContainKey(conn); + } + + @Test + @DisplayName("releaseConnection is no-op if connection is not in connectionKeys") + void releaseConnectionShouldBeNoOpIfNotTracked() { + ProxyToServerConnection conn = mock(); + pool.releaseConnection(conn); + // No exception expected + } + + @Test + @DisplayName("releaseConnection is no-op for null") + void releaseConnectionShouldBeNoOpForNull() { + pool.releaseConnection(null); + // No exception expected + } + + // ----------------------------------------------------------------------- + // removeConnection (public API) + // ----------------------------------------------------------------------- + + @Test + @DisplayName("removeConnection should clean connection from all maps") + void removeConnectionShouldCleanAllMaps() throws Exception { + ProxyToServerConnection conn = createMockConnection(true, true); + registerInPool(conn, "example.com:80:direct"); + pool.releaseConnection(conn); + + pool.removeConnection(conn); + + ConcurrentMap keys = getField(pool, "connectionKeys"); + assertThat(keys).doesNotContainKey(conn); + } + + @Test + @DisplayName("removeConnection is no-op for null") + void removeConnectionShouldBeNoOpForNull() { + pool.removeConnection(null); + } + + @Test + @DisplayName("removeConnection is no-op for untracked connection") + void removeConnectionShouldBeNoOpForUntrackedConnection() { + pool.removeConnection(mock(ProxyToServerConnection.class)); + } + + // ----------------------------------------------------------------------- + // borrowAvailableConnection — tested via getOrCreateConnection (public API) + // ----------------------------------------------------------------------- + + @Test + @DisplayName("getOrCreateConnection should return available connection when one exists") + void getOrCreateConnectionShouldReturnAvailableConnection() throws Exception { + ProxyToServerConnection conn = createMockConnection(true, true); + registerInPool(conn, "example.com:80:direct"); + pool.releaseConnection(conn); + + HttpRequest request = new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"); + HttpFilters filters = + new HttpFiltersAdapter(request) { + @Override + public InetSocketAddress proxyToServerResolutionStarted( + String resolvingServerHostAndPort) { + return null; + } + }; + ProxyToServerConnection result = + pool.getOrCreateConnection("example.com:80", null, mockClientConnection, filters, request); + + assertThat(result).isSameAs(conn); + } + + @Test + @DisplayName( + "borrowAvailableConnection should re-queue busy-but-connected connections " + + "and return an available one") + void borrowAvailableConnectionShouldRequeueBusyAndReturnAvailable() throws Exception { + ProxyToServerConnection busyConn = createMockConnection(true, false); + ProxyToServerConnection availableConn = createMockConnection(true, true); + registerInPool(busyConn, "example.com:80:direct"); + registerInPool(availableConn, "example.com:80:direct"); + pool.releaseConnection(busyConn); + pool.releaseConnection(availableConn); + + HttpRequest request = new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"); + HttpFilters filters = + new HttpFiltersAdapter(request) { + @Override + public InetSocketAddress proxyToServerResolutionStarted( + String resolvingServerHostAndPort) { + return null; + } + }; + ProxyToServerConnection result = + pool.getOrCreateConnection("example.com:80", null, mockClientConnection, filters, request); + + assertThat(result).isSameAs(availableConn); + + // The busy connection should still be in the available queue (re-queued, not lost) + Queue queue = + ((java.util.Map>) getField(pool, "availableConnectionsByHostAndPort")) + .get("example.com:80:direct"); + assertThat(queue).hasSize(1); + } + + @Test + @DisplayName("borrowAvailableConnection should not loop infinitely when all connections are busy") + void borrowAvailableConnectionShouldNotLoopWhenAllBusy() throws Exception { + ProxyToServerConnection busyConn = createMockConnection(true, false); + registerInPool(busyConn, "example.com:80:direct"); + pool.releaseConnection(busyConn); + + // getOrCreateConnection -> borrowAvailableConnection tries all connections in + // the queue, finds none available, returns null. Then getOrCreateConnection + // creates a new connection. + HttpRequest request = new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"); + HttpFilters filters = + new HttpFiltersAdapter(request) { + @Override + public InetSocketAddress proxyToServerResolutionStarted( + String resolvingServerHostAndPort) { + return null; + } + }; + ProxyToServerConnection result = + pool.getOrCreateConnection("example.com:80", null, mockClientConnection, filters, request); + + // A new connection was created (different from the busy one) + assertThat(result).isNotNull(); + assertThat(result).isNotSameAs(busyConn); + + // The busy connection should still be in the available queue + Queue queue = + ((java.util.Map>) getField(pool, "availableConnectionsByHostAndPort")) + .get("example.com:80:direct"); + assertThat(queue).hasSize(1); + } + + // ----------------------------------------------------------------------- + // closeAll (public API) + // ----------------------------------------------------------------------- + + @Test + @DisplayName("closeAll should shut down the eviction scheduler") + void closeAllShouldShutdownEvictionScheduler() throws Exception { + ScheduledExecutorService scheduler = getField(pool, "evictionScheduler"); + assertThat(scheduler.isShutdown()).isFalse(); + + pool.closeAll(); + + assertThat(scheduler.isShutdown()).isTrue(); + } + + @Test + @DisplayName("closeAll should clear all internal maps") + void closeAllShouldClearAllMaps() throws Exception { + ProxyToServerConnection conn = createMockConnection(true, true); + registerInPool(conn, "example.com:80:direct"); + pool.releaseConnection(conn); + + pool.closeAll(); + + assertThat((java.util.Map) getField(pool, "connectionsByHostAndPort")).isEmpty(); + assertThat((java.util.Map) getField(pool, "availableConnectionsByHostAndPort")).isEmpty(); + assertThat((java.util.Map) getField(pool, "connectionCountByHostAndPort")).isEmpty(); + assertThat((java.util.Map) getField(pool, "connectionKeys")).isEmpty(); + } + + @Test + @DisplayName("closeAll should be idempotent") + void closeAllShouldBeIdempotent() { + pool.closeAll(); + pool.closeAll(); + // No exception on second call + } + + // ----------------------------------------------------------------------- + // connectionKeys lifecycle (via public API getOrCreateConnection) + // ----------------------------------------------------------------------- + + @Test + @DisplayName("getOrCreateConnection should store pool key in connectionKeys on creation") + void getOrCreateConnectionShouldStorePoolKey() { + HttpRequest request = new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"); + HttpFilters filters = + new HttpFiltersAdapter(request) { + @Override + public InetSocketAddress proxyToServerResolutionStarted( + String resolvingServerHostAndPort) { + return null; + } + }; + ProxyToServerConnection conn = + pool.getOrCreateConnection("example.com:80", null, mockClientConnection, filters, request); + + assertThat(conn).isNotNull(); + + // The connection was created via createForPool, so connectionKeys was set inside + // getOrCreateConnection. Verify it is there. + ProxyToServerConnection finalConn = conn; // effectively final + Runnable check = + () -> { + try { + @SuppressWarnings("unchecked") + ConcurrentMap keys = getField(pool, "connectionKeys"); + assertThat(keys).containsKey(finalConn); + assertThat(keys.get(finalConn)).isEqualTo("example.com:80:direct"); + } catch (Exception e) { + throw new RuntimeException(e); + } + }; + check.run(); + } + + @Test + @DisplayName("getOrCreateConnection should create a new connection when none available") + void getOrCreateConnectionShouldCreateNewConnection() { + HttpRequest request = new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"); + HttpFilters filters = + new HttpFiltersAdapter(request) { + @Override + public InetSocketAddress proxyToServerResolutionStarted( + String resolvingServerHostAndPort) { + return null; + } + }; + ProxyToServerConnection conn = + pool.getOrCreateConnection("example.com:80", null, mockClientConnection, filters, request); + assertThat(conn).isNotNull(); + } + + // ----------------------------------------------------------------------- + // computePoolKey + // ----------------------------------------------------------------------- + + @Test + @DisplayName("computePoolKey should use :direct suffix for null chained proxy address") + void computePoolKeyForDirect() { + assertThat(pool.computePoolKey("example.com:80", null)).isEqualTo("example.com:80:direct"); + } + + @Test + @DisplayName("computePoolKey should include resolved chained proxy address") + void computePoolKeyForChainedProxy() { + InetSocketAddress proxyAddr = new InetSocketAddress("10.0.0.1", 3128); + assertThat(pool.computePoolKey("example.com:80", proxyAddr)) + .isEqualTo("example.com:80:10.0.0.1:3128"); + } + + @Test + @DisplayName("computePoolKey should use hostname for unresolved address") + void computePoolKeyForUnresolvedChainedProxy() { + InetSocketAddress proxyAddr = InetSocketAddress.createUnresolved("proxy.example.com", 3128); + assertThat(pool.computePoolKey("example.com:80", proxyAddr)) + .isEqualTo("example.com:80:proxy.example.com:3128"); + } + + // ----------------------------------------------------------------------- + // PendingRequest queue (public API) + // ----------------------------------------------------------------------- + + @Test + @DisplayName("drainPendingRequests should drain and remove pending requests") + void drainPendingRequestsShouldDrainAndRemove() { + EmbeddedChannel channel = new EmbeddedChannel(); + pool.registerPendingRequest( + channel, + mockClientConnection, + new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"), + null); + pool.drainPendingRequests(channel); + + assertThat(pool.peekPendingRequest(channel)).isNull(); + } + + @Test + @DisplayName("removePendingRequest should return oldest pending request") + void removePendingRequestShouldReturnOldest() { + EmbeddedChannel channel = new EmbeddedChannel(); + pool.registerPendingRequest( + channel, + mockClientConnection, + new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/first"), + null); + pool.registerPendingRequest( + channel, + mockClientConnection, + new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/second"), + null); + + PendingRequest first = pool.removePendingRequest(channel); + assertThat(first).isNotNull(); + assertThat(first.getRequest().uri()).isEqualTo("/first"); + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerBootstrapTest.java b/src/test/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerBootstrapTest.java new file mode 100644 index 00000000..3d340a80 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/DefaultHttpProxyServerBootstrapTest.java @@ -0,0 +1,56 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.lang.reflect.Field; +import java.util.Properties; +import org.junit.jupiter.api.Test; + +class DefaultHttpProxyServerBootstrapTest { + + private static Object getField(Object target, String name) throws Exception { + Field f = target.getClass().getDeclaredField(name); + f.setAccessible(true); + return f.get(target); + } + + @Test + void constructorParsesPoolSharedMitmConnectionsFromProperties() throws Exception { + Properties props = new Properties(); + props.setProperty("pool_shared_mitm_connections", "true"); + DefaultHttpProxyServerBootstrap bootstrap = new DefaultHttpProxyServerBootstrap(props); + assertThat(getField(bootstrap, "poolSharedMitmConnections")).isEqualTo(true); + } + + @Test + void constructorDefaultsPoolSharedMitmConnectionsToFalse() throws Exception { + Properties props = new Properties(); + DefaultHttpProxyServerBootstrap bootstrap = new DefaultHttpProxyServerBootstrap(props); + assertThat(getField(bootstrap, "poolSharedMitmConnections")).isEqualTo(false); + } + + @Test + void constructorParsesPoolPerRequestInMitmFromProperties() throws Exception { + Properties props = new Properties(); + props.setProperty("pool_per_request_in_mitm", "true"); + DefaultHttpProxyServerBootstrap bootstrap = new DefaultHttpProxyServerBootstrap(props); + assertThat(getField(bootstrap, "poolPerRequestInMitm")).isEqualTo(true); + } + + @Test + void constructorDefaultsPoolPerRequestInMitmToFalse() throws Exception { + Properties props = new Properties(); + DefaultHttpProxyServerBootstrap bootstrap = new DefaultHttpProxyServerBootstrap(props); + assertThat(getField(bootstrap, "poolPerRequestInMitm")).isEqualTo(false); + } + + @Test + void constructorParsesBothMitmPoolFlags() throws Exception { + Properties props = new Properties(); + props.setProperty("pool_shared_mitm_connections", "true"); + props.setProperty("pool_per_request_in_mitm", "true"); + DefaultHttpProxyServerBootstrap bootstrap = new DefaultHttpProxyServerBootstrap(props); + assertThat(getField(bootstrap, "poolSharedMitmConnections")).isEqualTo(true); + assertThat(getField(bootstrap, "poolPerRequestInMitm")).isEqualTo(true); + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ProxyConnectionHandlerAddedTest.java b/src/test/java/org/littleshoot/proxy/impl/ProxyConnectionHandlerAddedTest.java new file mode 100644 index 00000000..a4daec83 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ProxyConnectionHandlerAddedTest.java @@ -0,0 +1,111 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.*; + +import io.netty.channel.ChannelHandlerContext; +import io.netty.channel.ChannelPipeline; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.traffic.GlobalTrafficShapingHandler; +import java.lang.reflect.Field; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ChainedProxyManager; +import org.littleshoot.proxy.HttpFiltersSource; + +class ProxyConnectionHandlerAddedTest { + + private DefaultHttpProxyServer mockProxyServer; + private GlobalTrafficShapingHandler mockTrafficHandler; + + @BeforeEach + void setUp() { + mockProxyServer = mock(); + mockTrafficHandler = mock(); + when(mockProxyServer.getChainProxyManager()).thenReturn(mock(ChainedProxyManager.class)); + when(mockProxyServer.getFiltersSource()).thenReturn(mock(HttpFiltersSource.class)); + when(mockProxyServer.getMaxInitialLineLength()).thenReturn(8192); + when(mockProxyServer.getMaxHeaderSize()).thenReturn(16384); + when(mockProxyServer.getMaxChunkSize()).thenReturn(16384); + when(mockProxyServer.getIdleConnectionTimeout()).thenReturn(0); + when(mockProxyServer.isAcceptProxyProtocol()).thenReturn(false); + when(mockProxyServer.getProxyAlias()).thenReturn("test"); + when(mockProxyServer.isAllowRequestsToOriginServer()).thenReturn(true); + when(mockProxyServer.getActivityTrackers()).thenReturn(java.util.Collections.emptyList()); + } + + @Test + @DisplayName("handlerAdded should set ctx and channel before channelRegistered fires") + void handlerAddedShouldSetCtxAndChannel() throws Exception { + EmbeddedChannel channel = new EmbeddedChannel(); + ClientToProxyConnection clientConn = + new ClientToProxyConnection( + mockProxyServer, null, false, channel.pipeline(), mockTrafficHandler); + + // handlerAdded is invoked by Netty automatically when the handler is added to the pipeline. + // After that, ctx and channel should be non-null. + Field ctxField = ProxyConnection.class.getDeclaredField("ctx"); + ctxField.setAccessible(true); + ChannelHandlerContext ctx = (ChannelHandlerContext) ctxField.get(clientConn); + assertThat(ctx).as("ctx should be set after handlerAdded").isNotNull(); + assertThat(ctx.channel()).as("ctx.channel() should be the EmbeddedChannel").isSameAs(channel); + + Field channelField = ProxyConnection.class.getDeclaredField("channel"); + channelField.setAccessible(true); + Object connChannel = channelField.get(clientConn); + assertThat(connChannel).as("channel should be set after handlerAdded").isNotNull(); + assertThat(connChannel).isSameAs(channel); + + channel.finish(); + } + + @Test + @DisplayName("handlerAdded should set ctx before channelRegistered for ClientToProxyConnection") + void handlerAddedContextShouldBeAvailableInChannelRegistered() throws Exception { + // This test verifies the ordering guarantee: handlerAdded fires before channelRegistered + // when a handler is added to an active channel's pipeline. + // We use a pipeline that is already registered (via EmbeddedChannel) and add our handler + // to verify that the fields are correctly populated by handlerAdded. + + EmbeddedChannel channel = new EmbeddedChannel(); + ChannelPipeline pipeline = channel.pipeline(); + + ClientToProxyConnection clientConn = + new ClientToProxyConnection(mockProxyServer, null, false, pipeline, mockTrafficHandler); + + // The constructor calls initChannelPipeline which adds this handler to the pipeline. + // Since the EmbeddedChannel is already active, handlerAdded fires immediately + // before channelRegistered would fire. + + Field ctxField = ProxyConnection.class.getDeclaredField("ctx"); + ctxField.setAccessible(true); + ChannelHandlerContext ctx = (ChannelHandlerContext) ctxField.get(clientConn); + + assertThat(ctx).isNotNull(); + assertThat(ctx.pipeline()).isSameAs(pipeline); + + // Verify the channel is accessible through the context + assertThat(ctx.channel()).isSameAs(channel); + + channel.finish(); + } + + @Test + @DisplayName("handlerAdded should delegate to super.handlerAdded") + void handlerAddedShouldDelegateToSuper() throws Exception { + // We test this by verifying that no exception is thrown during normal pipeline setup + EmbeddedChannel channel = new EmbeddedChannel(); + ClientToProxyConnection clientConn = + new ClientToProxyConnection( + mockProxyServer, null, false, channel.pipeline(), mockTrafficHandler); + + // If super.handlerAdded wasn't called, the handler wouldn't be properly registered + // in the pipeline. Verify it's there. + assertThat(channel.pipeline().get("handler")) + .as("ClientToProxyConnection should be registered as 'handler' in the pipeline") + .isSameAs(clientConn); + + channel.finish(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionBugTest.java b/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionBugTest.java new file mode 100644 index 00000000..e4bcb19b --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionBugTest.java @@ -0,0 +1,386 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.*; + +import io.netty.buffer.ByteBuf; +import io.netty.channel.embedded.EmbeddedChannel; +import io.netty.handler.codec.haproxy.HAProxyProxiedProtocol; +import io.netty.handler.codec.http.DefaultHttpRequest; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpVersion; +import io.netty.handler.traffic.GlobalTrafficShapingHandler; +import java.lang.reflect.Field; +import java.lang.reflect.Method; +import java.net.InetSocketAddress; +import java.util.ArrayList; +import java.util.Collections; +import java.util.List; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ActivityTracker; +import org.littleshoot.proxy.ActivityTrackerAdapter; +import org.littleshoot.proxy.FlowContext; +import org.littleshoot.proxy.FullFlowContext; +import org.littleshoot.proxy.HostResolver; +import org.littleshoot.proxy.HttpFilters; +import org.littleshoot.proxy.extras.HAProxyMessageEncoder; +import org.littleshoot.proxy.extras.ProxyProtocolMessage; + +class ProxyToServerConnectionBugTest { + + private DefaultHttpProxyServer mockProxyServer; + private ClientToProxyConnection mockClientConnection; + private HttpFilters mockFilters; + private GlobalTrafficShapingHandler mockTrafficHandler; + private HostResolver mockHostResolver; + private FlowContext mockClientFlowContext; + private FullFlowContext mockFlowContext; + + @BeforeEach + void setup() throws Exception { + mockProxyServer = mock(); + mockClientConnection = mock(); + mockFilters = mock(); + mockTrafficHandler = mock(); + mockHostResolver = mock(); + + when(mockProxyServer.getServerResolver()).thenReturn(mockHostResolver); + when(mockHostResolver.resolve(any(), anyInt())) + .thenReturn(new InetSocketAddress("127.0.0.1", 8080)); + + mockClientFlowContext = mock(); + when(mockClientConnection.flowContext()).thenReturn(mockClientFlowContext); + mockFlowContext = mock(); + when(mockClientConnection.flowContextForServerConnection(any(ProxyToServerConnection.class))) + .thenReturn(mockFlowContext); + } + + private ProxyToServerConnection createConnection(List trackers) + throws Exception { + when(mockProxyServer.getActivityTrackers()).thenReturn(trackers); + return ProxyToServerConnection.create( + mockProxyServer, + mockClientConnection, + "localhost:8080", + mockFilters, + null, + mockTrafficHandler); + } + + private ProxyToServerConnection createConnectionWithPool( + ServerConnectionPool pool, List trackers) throws Exception { + when(mockProxyServer.getActivityTrackers()).thenReturn(trackers); + return ProxyToServerConnection.createForPool( + mockProxyServer, + pool, + mockClientConnection, + "localhost:8080", + null, + mockFilters, + new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/"), + mockTrafficHandler); + } + + // ============================================================ + // Bug 1: releaseToPool() must exist and release back to pool + // ============================================================ + + @Test + @DisplayName("releaseToPool should exist and call pool.releaseConnection") + void releaseToPoolShouldReleaseBackToPool() throws Exception { + ServerConnectionPool mockPool = mock(); + ProxyToServerConnection conn = createConnectionWithPool(mockPool, Collections.emptyList()); + assertThat(conn).isNotNull(); + + conn.setCurrentClientConnectionForRequest(mockClientConnection); + conn.releaseToPool(); + + verify(mockPool).releaseConnection(conn); + } + + @Test + @DisplayName("releaseToPool should be no-op when no pool is set") + void releaseToPoolShouldBeNoopWithoutPool() throws Exception { + ProxyToServerConnection conn = createConnection(Collections.emptyList()); + assertThat(conn).isNotNull(); + + conn.releaseToPool(); + } + + // ============================================================ + // Bug 2: clientConnected should fire before requestReceivedFromClient + // + // Tests the Netty pipeline ordering by creating a real + // ClientToProxyConnection with EmbeddedChannel and sending an HTTP request. + // ============================================================ + + @Test + @DisplayName( + "clientConnected should fire before requestReceivedFromClient when request arrives without PROXY header") + void clientConnectedShouldFireBeforeRequestReceived() throws Exception { + when(mockProxyServer.getFiltersSource()) + .thenReturn( + new org.littleshoot.proxy.HttpFiltersSource() { + @Override + public int getMaximumRequestBufferSizeInBytes() { + return 0; + } + + @Override + public int getMaximumResponseBufferSizeInBytes() { + return 0; + } + + @Override + public org.littleshoot.proxy.HttpFilters filterRequest( + io.netty.handler.codec.http.HttpRequest httpRequest, + io.netty.channel.ChannelHandlerContext ctx) { + return null; + } + }); + when(mockProxyServer.getChainProxyManager()) + .thenReturn(mock(org.littleshoot.proxy.ChainedProxyManager.class)); + when(mockProxyServer.getMaxInitialLineLength()).thenReturn(8192); + when(mockProxyServer.getMaxHeaderSize()).thenReturn(16384); + when(mockProxyServer.getMaxChunkSize()).thenReturn(16384); + when(mockProxyServer.getIdleConnectionTimeout()).thenReturn(0); + when(mockProxyServer.isAcceptProxyProtocol()).thenReturn(false); + when(mockProxyServer.getProxyAlias()).thenReturn("test-proxy"); + when(mockProxyServer.isAllowRequestsToOriginServer()).thenReturn(true); + + List eventOrder = new ArrayList<>(); + ActivityTracker tracker = + new ActivityTrackerAdapter() { + @Override + public void clientConnected(FlowContext flowContext) { + eventOrder.add("clientConnected"); + } + + @Override + public void requestReceivedFromClient(FlowContext flowContext, HttpRequest httpRequest) { + eventOrder.add("requestReceivedFromClient"); + } + }; + when(mockProxyServer.getActivityTrackers()).thenReturn(Collections.singletonList(tracker)); + + EmbeddedChannel channel = new EmbeddedChannel(); + ClientToProxyConnection clientConn = + new ClientToProxyConnection( + mockProxyServer, null, false, channel.pipeline(), mockTrafficHandler); + + DefaultHttpRequest request = + new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "http://example.com/"); + channel.writeInbound(request); + + assertThat(eventOrder) + .as( + "clientConnected should fire before requestReceivedFromClient for non-PROXY connections") + .containsSequence("clientConnected", "requestReceivedFromClient"); + + channel.finish(); + } + + // ============================================================ + // Bug 3: recordServerConnected/disconnected should use getClientConnection() + // ============================================================ + + @Test + @DisplayName( + "recordServerConnected should use getClientConnection() flowContext, not constructor client") + void recordServerConnectedShouldUseCurrentClient() throws Exception { + ClientToProxyConnection mockClient2 = mock(); + FullFlowContext mockFlowContext2Full = mock(); + when(mockClient2.flowContextForServerConnection(any(ProxyToServerConnection.class))) + .thenReturn(mockFlowContext2Full); + + ActivityTracker tracker = mock(ActivityTracker.class); + ProxyToServerConnection conn = + createConnectionWithPool(mock(), Collections.singletonList(tracker)); + assertThat(conn).isNotNull(); + + conn.setCurrentClientConnectionForRequest(mockClient2); + conn.recordServerConnected(); + + verify(tracker).serverConnected(eq(mockFlowContext2Full), any(InetSocketAddress.class)); + } + + @Test + @DisplayName( + "recordServerDisconnected should use getClientConnection() and clear via the correct client") + void recordServerDisconnectedShouldUseCurrentClient() throws Exception { + ClientToProxyConnection mockClient2 = mock(); + FullFlowContext mockFlowContext2Full = mock(); + when(mockClient2.flowContextForServerConnection(any(ProxyToServerConnection.class))) + .thenReturn(mockFlowContext2Full); + + ActivityTracker tracker = mock(ActivityTracker.class); + ProxyToServerConnection conn = + createConnectionWithPool(mock(), Collections.singletonList(tracker)); + assertThat(conn).isNotNull(); + + conn.setCurrentClientConnectionForRequest(mockClient2); + conn.recordServerDisconnected(); + + verify(tracker).serverDisconnected(eq(mockFlowContext2Full), any(InetSocketAddress.class)); + verify(mockClient2).clearFlowContextForServerConnection(conn); + } + + // ============================================================ + // Bug 5: markResponseComplete must dequeue and wire next pipelined request + // ============================================================ + + @Test + @DisplayName("markResponseComplete should dequeue pending request and keep connection in use") + void markResponseCompleteShouldDequeueAndWireNextPending() throws Exception { + ServerConnectionPool mockPool = mock(); + ProxyToServerConnection conn = createConnectionWithPool(mockPool, Collections.emptyList()); + EmbeddedChannel ch = new EmbeddedChannel(); + conn.channel = ch; + + HttpRequest pendingReq = new DefaultHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/next"); + ClientToProxyConnection pendingClient = mock(); + when(pendingClient.flowContext()).thenReturn(mock()); + when(pendingClient.flowContextForServerConnection(any(ProxyToServerConnection.class))) + .thenReturn(mock()); + HttpFilters pendingFilters = mock(); + PendingRequest pending = new PendingRequest(pendingClient, pendingReq, pendingFilters); + + when(mockPool.removePendingRequest(ch)).thenReturn(pending); + + invokeMarkResponseComplete(conn); + + assertThat(getField(conn, "currentHttpRequest")).isSameAs(pendingReq); + assertThat(getField(conn, "currentClientConnectionForRequest")).isSameAs(pendingClient); + assertThat(getField(conn, "currentFilters")).isSameAs(pendingFilters); + assertThat(getField(conn, "currentHttpResponse")).isNull(); + + verify(mockPool, never()).releaseConnection(conn); + verify(mockPool, never()).peekPendingRequest(any()); + } + + @Test + @DisplayName( + "markResponseComplete should release connection to pool when no pending requests remain") + void markResponseCompleteShouldReleaseToPoolWhenNoPending() throws Exception { + ServerConnectionPool mockPool = mock(); + ProxyToServerConnection conn = createConnectionWithPool(mockPool, Collections.emptyList()); + EmbeddedChannel ch = new EmbeddedChannel(); + conn.channel = ch; + + when(mockPool.removePendingRequest(ch)).thenReturn(null); + + invokeMarkResponseComplete(conn); + + assertThat(getField(conn, "currentHttpRequest")).isNull(); + assertThat(getField(conn, "currentHttpResponse")).isNull(); + + verify(mockPool).releaseConnection(conn); + } + + private static void invokeMarkResponseComplete(ProxyToServerConnection conn) throws Exception { + Method m = ProxyToServerConnection.class.getDeclaredMethod("markResponseComplete"); + m.setAccessible(true); + m.invoke(conn); + } + + private static Object getField(Object obj, String name) throws Exception { + Class clazz = obj.getClass(); + while (clazz != null) { + try { + Field f = clazz.getDeclaredField(name); + f.setAccessible(true); + return f.get(obj); + } catch (NoSuchFieldException e) { + clazz = clazz.getSuperclass(); + } + } + throw new NoSuchFieldException(name + " in " + obj.getClass().getName()); + } + + // ============================================================ + // Bug 4: SendProxyProtocolHeader always uses TCP4 + // + // Tests that HAProxyMessageEncoder correctly encodes TCP6 headers, + // and that SendProxyProtocolHeader selects the right protocol. + // ============================================================ + + @Test + @DisplayName("HAProxyMessageEncoder should produce valid PROXY TCP4 header for IPv4 addresses") + void encoderShouldProduceValidTcp4Header() throws Exception { + ProxyProtocolMessage msg = + new ProxyProtocolMessage( + io.netty.handler.codec.haproxy.HAProxyProtocolVersion.V1, + io.netty.handler.codec.haproxy.HAProxyCommand.PROXY, + HAProxyProxiedProtocol.TCP4, + "192.168.1.1", + "10.0.0.1", + 12345, + 443); + + EmbeddedChannel ch = new EmbeddedChannel(new HAProxyMessageEncoder()); + ch.writeOutbound(msg); + ByteBuf out = ch.readOutbound(); + String header = out.toString(io.netty.util.CharsetUtil.US_ASCII); + out.release(); + ch.finish(); + + assertThat(header).startsWith("PROXY TCP4 192.168.1.1 10.0.0.1 12345 443\r\n"); + } + + @Test + @DisplayName("HAProxyMessageEncoder should produce valid PROXY TCP6 header for IPv6 addresses") + void encoderShouldProduceValidTcp6Header() throws Exception { + ProxyProtocolMessage msg = + new ProxyProtocolMessage( + io.netty.handler.codec.haproxy.HAProxyProtocolVersion.V1, + io.netty.handler.codec.haproxy.HAProxyCommand.PROXY, + HAProxyProxiedProtocol.TCP6, + "2001:db8::1", + "2001:db8::2", + 12345, + 443); + + EmbeddedChannel ch = new EmbeddedChannel(new HAProxyMessageEncoder()); + ch.writeOutbound(msg); + ByteBuf out = ch.readOutbound(); + String header = out.toString(io.netty.util.CharsetUtil.US_ASCII); + out.release(); + ch.finish(); + + assertThat(header).startsWith("PROXY TCP6 2001:db8::1 2001:db8::2 12345 443\r\n"); + } + + @Test + @DisplayName("SendProxyProtocolHeader should select TCP6 when both client and server are IPv6") + void sendProxyProtocolHeaderShouldSelectTcp6ForIpv6() throws Exception { + ServerConnectionPool mockPool = mock(); + ProxyToServerConnection conn = createConnectionWithPool(mockPool, Collections.emptyList()); + assertThat(conn).isNotNull(); + + InetSocketAddress ipv6ClientAddr = new InetSocketAddress("2001:db8::1", 12345); + InetSocketAddress ipv6RemoteAddr = new InetSocketAddress("2001:db8::2", 443); + when(mockClientConnection.getHaProxyMessage()).thenReturn(null); + when(mockClientConnection.getClientAddress()).thenReturn(ipv6ClientAddr); + + conn.setRemoteAddress(ipv6RemoteAddr); + + EmbeddedChannel channel = new EmbeddedChannel(new HAProxyMessageEncoder()); + conn.channel = channel; + + conn.SendProxyProtocolHeader.execute(); + + ByteBuf out = channel.readOutbound(); + assertThat(out).as("PROXY protocol header should be written to the channel").isNotNull(); + String header = out.toString(io.netty.util.CharsetUtil.US_ASCII); + out.release(); + channel.finish(); + + assertThat(header) + .as("IPv6 addresses should produce a TCP6 PROXY protocol header") + .startsWith("PROXY TCP6 2001:db8:0:0:0:0:0:1 2001:db8:0:0:0:0:0:2 12345 443\r\n"); + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionTest.java b/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionTest.java new file mode 100644 index 00000000..13734938 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionTest.java @@ -0,0 +1,322 @@ +package org.littleshoot.proxy.impl; + +import static java.util.Locale.ROOT; +import static java.util.Objects.requireNonNull; +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.anyInt; +import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import io.netty.handler.traffic.GlobalTrafficShapingHandler; +import java.lang.reflect.Field; +import java.net.InetSocketAddress; +import java.net.UnknownHostException; +import java.util.List; +import java.util.Queue; +import javax.net.ssl.SSLEngine; +import javax.net.ssl.SSLException; +import javax.net.ssl.SSLHandshakeException; +import javax.net.ssl.SSLProtocolException; +import org.jspecify.annotations.NullMarked; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Nested; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.CsvSource; +import org.littleshoot.proxy.ActivityTracker; +import org.littleshoot.proxy.ChainedProxy; +import org.littleshoot.proxy.ChainedProxyManager; +import org.littleshoot.proxy.ChainedProxyType; +import org.littleshoot.proxy.FlowContext; +import org.littleshoot.proxy.FullFlowContext; +import org.littleshoot.proxy.HostResolver; +import org.littleshoot.proxy.HttpFilters; +import org.littleshoot.proxy.TransportProtocol; + +final class ProxyToServerConnectionTest { + + private final DefaultHttpProxyServer proxyServer = mock(); + private final ClientToProxyConnection clientConnection = mock(); + private final HttpFilters filters = mock(); + private final GlobalTrafficShapingHandler trafficHandler = mock(); + private final FlowContext flowContext = mock(); + private final FullFlowContext fullFlowContext = mock(); + private final InetSocketAddress proxyAddress = new InetSocketAddress("127.0.0.1", 9443); + private final InetSocketAddress hostAddress = new InetSocketAddress("127.0.0.1", 8080); + + @BeforeEach + void setup() throws UnknownHostException { + HostResolver hostResolver = mock(); + + when(proxyServer.getServerResolver()).thenReturn(hostResolver); + when(hostResolver.resolve(any(), anyInt())).thenReturn(hostAddress); + + when(clientConnection.flowContext()).thenReturn(flowContext); + when(clientConnection.flowContextForServerConnection(any())).thenReturn(fullFlowContext); + } + + @NullMarked + private ProxyToServerConnection createConnection(ActivityTracker... trackers) + throws UnknownHostException { + when(proxyServer.getActivityTrackers()).thenReturn(List.of(trackers)); + return requireNonNull( + ProxyToServerConnection.create( + proxyServer, clientConnection, "localhost:8080", filters, null, trafficHandler)); + } + + @Test + @DisplayName("disconnected should clear flow context even when ActivityTracker throws exception") + void disconnectedShouldClearFlowContextEvenWhenActivityTrackerThrowsException() throws Exception { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")) + .when(throwingTracker) + .serverDisconnected(any(), any()); + ProxyToServerConnection connection = createConnection(throwingTracker); + + connection.disconnected(); + + verify(throwingTracker).serverDisconnected(fullFlowContext, hostAddress); + verify(clientConnection).clearFlowContextForServerConnection(connection); + } + + @Test + @DisplayName("disconnected should clear flow context when no exception occurs") + void disconnectedShouldClearFlowContextWhenNoException() throws Exception { + ActivityTracker normalTracker = mock(); + ProxyToServerConnection connection = createConnection(normalTracker); + + connection.disconnected(); + + verify(normalTracker).serverDisconnected(fullFlowContext, hostAddress); + verify(clientConnection).clearFlowContextForServerConnection(connection); + } + + @Test + @DisplayName("serverConnected should notify all trackers even if one throws") + void serverConnectedShouldNotifyAllTrackersEvenIfOneThrows() throws Exception { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")) + .when(throwingTracker) + .serverConnected(any(), any()); + ActivityTracker succeedingTracker = mock(); + ProxyToServerConnection connection = createConnection(throwingTracker, succeedingTracker); + + connection.recordServerConnected(); + + verify(throwingTracker).serverConnected(fullFlowContext, hostAddress); + verify(succeedingTracker).serverConnected(fullFlowContext, hostAddress); + } + + @Test + @DisplayName( + "serverDisconnected should notify all trackers even if one throws and still clear context") + void serverDisconnectedShouldNotifyAllTrackersEvenIfOneThrows() throws Exception { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")) + .when(throwingTracker) + .serverDisconnected(any(), any()); + ActivityTracker succeedingTracker = mock(); + ProxyToServerConnection connection = createConnection(throwingTracker, succeedingTracker); + + connection.recordServerDisconnected(); + + verify(throwingTracker).serverDisconnected(fullFlowContext, hostAddress); + verify(succeedingTracker).serverDisconnected(fullFlowContext, hostAddress); + verify(clientConnection).clearFlowContextForServerConnection(connection); + } + + @Test + @DisplayName("connectionSaturated should notify all trackers even if one throws") + void connectionSaturatedShouldNotifyAllTrackersEvenIfOneThrows() throws Exception { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")) + .when(throwingTracker) + .connectionSaturated(any()); + ActivityTracker succeedingTracker = mock(); + ProxyToServerConnection connection = createConnection(throwingTracker, succeedingTracker); + + connection.recordConnectionSaturated(); + + verify(throwingTracker).connectionSaturated(fullFlowContext); + verify(succeedingTracker).connectionSaturated(fullFlowContext); + } + + @Test + @DisplayName("connectionWritable should notify all trackers even if one throws") + void connectionWritableShouldNotifyAllTrackersEvenIfOneThrows() throws Exception { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")).when(throwingTracker).connectionWritable(any()); + ActivityTracker succeedingTracker = mock(); + ProxyToServerConnection connection = createConnection(throwingTracker, succeedingTracker); + + connection.recordConnectionWritable(); + + verify(throwingTracker).connectionWritable(fullFlowContext); + verify(succeedingTracker).connectionWritable(fullFlowContext); + } + + @Test + @DisplayName("connectionTimedOut should notify all trackers even if one throws") + void connectionTimedOutShouldNotifyAllTrackersEvenIfOneThrows() throws Exception { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")).when(throwingTracker).connectionTimedOut(any()); + ActivityTracker succeedingTracker = mock(); + ProxyToServerConnection connection = createConnection(throwingTracker, succeedingTracker); + + connection.recordConnectionTimedOut(); + + verify(throwingTracker).connectionTimedOut(fullFlowContext); + verify(succeedingTracker).connectionTimedOut(fullFlowContext); + } + + @Test + @DisplayName("encrypted chained proxies should prefer peer-aware SSL engines") + void encryptedChainedProxiesShouldPreferPeerAwareSslEngines() throws Exception { + ChainedProxy chainedProxy = mock(); + SSLEngine peerAwareEngine = mock(); + when(chainedProxy.newSslEngine("127.0.0.1", 9443)).thenReturn(peerAwareEngine); + ProxyToServerConnection connection = createConnectionWithChainedProxy(chainedProxy); + + assertThat(connection.newChainedProxySslEngine()).isSameAs(peerAwareEngine); + + verify(chainedProxy).newSslEngine("127.0.0.1", 9443); + verify(chainedProxy, never()).newSslEngine(); + } + + @Test + @DisplayName("encrypted chained proxies should fall back to legacy SSL engines") + void encryptedChainedProxiesShouldFallBackToLegacySslEngines() throws Exception { + ChainedProxy chainedProxy = mock(); + SSLEngine legacyEngine = mock(); + when(chainedProxy.newSslEngine("127.0.0.1", 9443)).thenReturn(null); + when(chainedProxy.newSslEngine()).thenReturn(legacyEngine); + ProxyToServerConnection connection = createConnectionWithChainedProxy(chainedProxy); + + assertThat(connection.newChainedProxySslEngine()).isSameAs(legacyEngine); + + verify(chainedProxy).newSslEngine("127.0.0.1", 9443); + verify(chainedProxy).newSslEngine(); + } + + @Test + @DisplayName("connectionExceptionCaught should notify all trackers even if one throws") + void connectionExceptionCaughtShouldNotifyAllTrackersEvenIfOneThrows() throws Exception { + ActivityTracker throwingTracker = mock(); + doThrow(new RuntimeException("Test exception")) + .when(throwingTracker) + .connectionExceptionCaught(any(), any()); + ActivityTracker succeedingTracker = mock(); + ProxyToServerConnection connection = createConnection(throwingTracker, succeedingTracker); + RuntimeException cause = new RuntimeException("Test cause"); + + connection.recordConnectionExceptionCaught(cause); + + verify(throwingTracker).connectionExceptionCaught(fullFlowContext, cause); + verify(succeedingTracker).connectionExceptionCaught(fullFlowContext, cause); + } + + private ProxyToServerConnection createConnectionWithChainedProxy(ChainedProxy chainedProxy) + throws UnknownHostException { + ChainedProxyManager chainedProxyManager = mock(); + when(proxyServer.getChainProxyManager()).thenReturn(chainedProxyManager); + doAnswer( + invocation -> { + invocation.>getArgument(1).add(chainedProxy); + return null; + }) + .when(chainedProxyManager) + .lookupChainedProxies(any(), any(), any()); + + when(chainedProxy.getTransportProtocol()).thenReturn(TransportProtocol.TCP); + when(chainedProxy.getChainedProxyType()).thenReturn(ChainedProxyType.HTTP); + when(chainedProxy.getChainedProxyAddress()).thenReturn(proxyAddress); + + return createConnection(); + } + + /** + * Helper to set the private {@code disableSslForNonTls} field via reflection. + * + * @param connection the connection instance + * @param value the value to set + */ + private void setDisableSslForNonTls(ProxyToServerConnection connection, boolean value) + throws Exception { + Field field = ProxyToServerConnection.class.getDeclaredField("disableSslForNonTls"); + field.setAccessible(true); + field.setBoolean(connection, value); + } + + @Nested + class ShouldRetryWithoutSsl { + @ParameterizedTest + @CsvSource({ + "Remote host terminated the handshake", + "end of file", + "not an SSL/TLS record", + "Connection reset" + }) + void returnsTrue_forKnownErrorMessages(String errorMessage) throws Exception { + ProxyToServerConnection connection = createConnection(); + + assertThat(connection.shouldRetryWithoutSsl(new SSLHandshakeException(errorMessage))) + .isTrue(); + assertThat(connection.shouldRetryWithoutSsl(new SSLProtocolException(errorMessage))).isTrue(); + assertThat(connection.shouldRetryWithoutSsl(new SSLException(errorMessage))).isTrue(); + + assertThat( + connection.shouldRetryWithoutSsl( + new SSLHandshakeException(errorMessage.toUpperCase(ROOT)))) + .as("case-insensitive") + .isTrue(); + assertThat( + connection.shouldRetryWithoutSsl( + new SSLHandshakeException(errorMessage.toLowerCase(ROOT)))) + .as("case-insensitive") + .isTrue(); + } + + @Test + @DisplayName("should return false for null cause") + void returnsFalse_forNullCause() throws Exception { + ProxyToServerConnection connection = createConnection(); + + assertThat(connection.shouldRetryWithoutSsl(null)).isFalse(); + } + + @Test + @DisplayName("should return false for non-SSL exceptions") + void returnsFalse_forNonSslExceptions() throws Exception { + ProxyToServerConnection connection = createConnection(); + + assertThat(connection.shouldRetryWithoutSsl(new RuntimeException("some error"))).isFalse(); + } + + @Test + @DisplayName("should return false when already retried without SSL") + void returnsFalse_WhenAlreadyRetried() throws Exception { + ProxyToServerConnection connection = createConnection(); + + setDisableSslForNonTls(connection, true); // TODO + + SSLHandshakeException cause = + new SSLHandshakeException("Remote host terminated the handshake"); + assertThat(connection.shouldRetryWithoutSsl(cause)).isFalse(); + } + + @Test + @DisplayName("should return false for SSL exceptions with non-matching messages") + void returnsFalse_rorNonMatchingMessages() throws Exception { + ProxyToServerConnection connection = createConnection(); + + SSLHandshakeException cause = new SSLHandshakeException("certificate expired"); + assertThat(connection.shouldRetryWithoutSsl(cause)).isFalse(); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionUtilsTest.java b/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionUtilsTest.java index 2cda81b0..cbadeb1a 100644 --- a/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionUtilsTest.java +++ b/src/test/java/org/littleshoot/proxy/impl/ProxyToServerConnectionUtilsTest.java @@ -1,46 +1,43 @@ package org.littleshoot.proxy.impl; -import org.junit.Test; -import org.littleshoot.proxy.HostResolver; +import static org.mockito.Mockito.*; import java.net.UnknownHostException; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.HostResolver; -import static org.mockito.Mockito.*; - -/** - * Unit tests for static helper methods in {@link ProxyToServerConnection}. - */ -public class ProxyToServerConnectionUtilsTest { - @Test - public void testParseAddresses() throws UnknownHostException { - // mock out the proxy server and resolver; this test only verifies the addresses parse correctly - DefaultHttpProxyServer mockProxyServer = mock(DefaultHttpProxyServer.class); - HostResolver mockHostResolver = mock(HostResolver.class); +/** Unit tests for static helper methods in {@link ProxyToServerConnection}. */ +public final class ProxyToServerConnectionUtilsTest { + @Test + public void testParseAddresses() throws UnknownHostException { + // mock out the proxy server and resolver; this test only verifies the addresses parse correctly + DefaultHttpProxyServer mockProxyServer = mock(); + HostResolver mockHostResolver = mock(); - when(mockProxyServer.getServerResolver()).thenReturn(mockHostResolver); + when(mockProxyServer.getServerResolver()).thenReturn(mockHostResolver); - ProxyToServerConnection.addressFor("192.168.1.1", mockProxyServer); - verify(mockHostResolver).resolve("192.168.1.1", 80); + ProxyToServerConnection.addressFor("192.168.1.1", mockProxyServer); + verify(mockHostResolver).resolve("192.168.1.1", 80); - ProxyToServerConnection.addressFor("192.168.1.1:72", mockProxyServer); - verify(mockHostResolver).resolve("192.168.1.1", 72); + ProxyToServerConnection.addressFor("192.168.1.1:72", mockProxyServer); + verify(mockHostResolver).resolve("192.168.1.1", 72); - ProxyToServerConnection.addressFor("www.google.com", mockProxyServer); - verify(mockHostResolver).resolve("www.google.com", 80); + ProxyToServerConnection.addressFor("www.google.com", mockProxyServer); + verify(mockHostResolver).resolve("www.google.com", 80); - ProxyToServerConnection.addressFor("www.google.com:19650", mockProxyServer); - verify(mockHostResolver).resolve("www.google.com", 19650); + ProxyToServerConnection.addressFor("www.google.com:19650", mockProxyServer); + verify(mockHostResolver).resolve("www.google.com", 19650); - ProxyToServerConnection.addressFor("[::1]", mockProxyServer); - verify(mockHostResolver).resolve("::1", 80); + ProxyToServerConnection.addressFor("[::1]", mockProxyServer); + verify(mockHostResolver).resolve("::1", 80); - ProxyToServerConnection.addressFor("[::1]:56500", mockProxyServer); - verify(mockHostResolver).resolve("::1", 56500); + ProxyToServerConnection.addressFor("[::1]:56500", mockProxyServer); + verify(mockHostResolver).resolve("::1", 56500); - ProxyToServerConnection.addressFor("[a:b:c:d::1]", mockProxyServer); - verify(mockHostResolver).resolve("a:b:c:d::1", 80); + ProxyToServerConnection.addressFor("[a:b:c:d::1]", mockProxyServer); + verify(mockHostResolver).resolve("a:b:c:d::1", 80); - ProxyToServerConnection.addressFor("[a:b:c:d::1]:8650", mockProxyServer); - verify(mockHostResolver).resolve("a:b:c:d::1", 8650); - } + ProxyToServerConnection.addressFor("[a:b:c:d::1]:8650", mockProxyServer); + verify(mockHostResolver).resolve("a:b:c:d::1", 8650); + } } diff --git a/src/test/java/org/littleshoot/proxy/impl/ProxyUtilsTest.java b/src/test/java/org/littleshoot/proxy/impl/ProxyUtilsTest.java index 53be3450..4479ea09 100644 --- a/src/test/java/org/littleshoot/proxy/impl/ProxyUtilsTest.java +++ b/src/test/java/org/littleshoot/proxy/impl/ProxyUtilsTest.java @@ -1,280 +1,343 @@ package org.littleshoot.proxy.impl; -import io.netty.handler.codec.http.*; -import org.junit.Test; +import static io.netty.handler.codec.http.HttpHeaderNames.ACCEPT_ENCODING; +import static io.netty.handler.codec.http.HttpHeaderNames.TRANSFER_ENCODING; +import static io.netty.handler.codec.http.HttpResponseStatus.OK; +import static io.netty.handler.codec.http.HttpResponseStatus.SWITCHING_PROTOCOLS; +import static io.netty.handler.codec.http.HttpVersion.HTTP_1_1; +import static java.util.Collections.singletonList; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.impl.ProxyUtils.parseHostAndPort; +import io.netty.handler.codec.http.*; import java.util.ArrayList; import java.util.Arrays; import java.util.List; - -import static java.util.Collections.singletonList; -import static org.hamcrest.Matchers.*; -import static org.junit.Assert.*; - -/** - * Test for proxy utilities. - */ -public class ProxyUtilsTest { - - @Test - public void testParseHostAndPort() { - assertEquals("www.test.com:80", ProxyUtils.parseHostAndPort("http://www.test.com:80/test")); - assertEquals("www.test.com:80", ProxyUtils.parseHostAndPort("https://www.test.com:80/test")); - assertEquals("www.test.com:443", ProxyUtils.parseHostAndPort("https://www.test.com:443/test")); - assertEquals("www.test.com:80", ProxyUtils.parseHostAndPort("www.test.com:80/test")); - assertEquals("www.test.com", ProxyUtils.parseHostAndPort("http://www.test.com")); - assertEquals("www.test.com", ProxyUtils.parseHostAndPort("www.test.com")); - assertEquals("httpbin.org:443", ProxyUtils.parseHostAndPort("httpbin.org:443/get")); - } - - @Test - public void testAddNewViaHeader() { - String hostname = "hostname"; - - HttpMessage httpMessage = new DefaultFullHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/endpoint"); - ProxyUtils.addVia(httpMessage, hostname); - - List viaHeaders = httpMessage.headers().getAll(HttpHeaderNames.VIA); - assertThat(viaHeaders, hasSize(1)); - - String expectedViaHeader = "1.1 " + hostname; - assertEquals(expectedViaHeader, viaHeaders.get(0)); - } - - @Test - public void testCommaSeparatedHeaderValues() { - DefaultHttpMessage message; - List commaSeparatedHeaders; - - // test the empty headers case - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, empty()); - - // two headers present, but no values - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, ""); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, ""); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, empty()); - - // a single header value - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "chunked"); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, contains("chunked")); - - // a single header value with extra spaces - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, " chunked , "); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, contains("chunked")); - - // two comma-separated values in one header line - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "compress, gzip"); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, contains("compress", "gzip")); - - // two comma-separated values in one header line with a spurious ',' and space. see RFC 7230 section 7 - // for information on empty list items (not all of which are valid header-values). - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "compress, gzip, ,"); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, contains("compress", "gzip")); - - // two values in two separate header lines - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "gzip"); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "chunked"); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, contains("gzip", "chunked")); - - // multiple comma-separated values in two separate header lines - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "gzip, compress"); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "deflate, gzip"); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, contains("gzip", "compress", "deflate", "gzip")); - - // multiple comma-separated values in multiple header lines with spurious spaces, commas, - // and tabs (horizontal tabs are defined as optional whitespace in RFC 7230 section 3.2.3) - message = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, " gzip,compress,"); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "\tdeflate\t, gzip, "); - message.headers().add(HttpHeaderNames.TRANSFER_ENCODING, ",gzip,,deflate,\t, ,"); - commaSeparatedHeaders = ProxyUtils.getAllCommaSeparatedHeaderValues(HttpHeaderNames.TRANSFER_ENCODING, message); - assertThat(commaSeparatedHeaders, contains("gzip", "compress", "deflate", "gzip", "gzip", "deflate")); - } - - @Test - public void testIsResponseSelfTerminating() { - HttpResponse httpResponse; - boolean isResponseSelfTerminating; - - // test cases from the scenarios listed in RFC 2616, section 4.4 - // #1: 1.Any response message which "MUST NOT" include a message-body (such as the 1xx, 204, and 304 responses and any response to a HEAD request) is always terminated by the first empty line after the header fields, regardless of the entity-header fields present in the message. - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.CONTINUE); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.SWITCHING_PROTOCOLS); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.NO_CONTENT); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.RESET_CONTENT); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.NOT_MODIFIED); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - // #2: 2.If a Transfer-Encoding header field (section 14.41) is present and has any value other than "identity", then the transfer-length is defined by use of the "chunked" transfer-coding (section 3.6), unless the message is terminated by closing the connection. - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "chunked"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "gzip, chunked"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - // chunked encoding is not last, so not self terminating - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "chunked, gzip"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertFalse(isResponseSelfTerminating); - - // four encodings on two lines, chunked is not last, so not self terminating - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "gzip, chunked"); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "deflate, gzip"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertFalse(isResponseSelfTerminating); - - // three encodings on two lines, chunked is last, so self terminating - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "gzip"); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "deflate,chunked"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - // #3: 3.If a Content-Length header field (section 14.13) is present, its decimal value in OCTETs represents both the entity-length and the transfer-length. - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.CONTENT_LENGTH, "15"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - // continuing #3: If a message is received with both a Transfer-Encoding header field and a Content-Length header field, the latter MUST be ignored. - - // chunked is last Transfer-Encoding, so message is self-terminating - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "gzip, chunked"); - httpResponse.headers().add(HttpHeaderNames.CONTENT_LENGTH, "15"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertTrue(isResponseSelfTerminating); - - // chunked is not last Transfer-Encoding, so message is not self-terminating, since Content-Length is ignored - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - httpResponse.headers().add(HttpHeaderNames.TRANSFER_ENCODING, "gzip"); - httpResponse.headers().add(HttpHeaderNames.CONTENT_LENGTH, "15"); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertFalse(isResponseSelfTerminating); - - // without any of the above conditions, the message should not be self-terminating - // (multipart/byteranges is ignored, see note in method javadoc) - httpResponse = new DefaultHttpResponse(HttpVersion.HTTP_1_1, HttpResponseStatus.OK); - isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); - assertFalse(isResponseSelfTerminating); - +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +/** Test for proxy utilities. */ +@SuppressWarnings("deprecation") +public final class ProxyUtilsTest { + + @Test + public void testParseHostAndPort() { + assertThat(parseHostAndPort("http://www.test.com:80/test")).isEqualTo("www.test.com:80"); + assertThat(parseHostAndPort("https://www.test.com:80/test")).isEqualTo("www.test.com:80"); + assertThat(parseHostAndPort("https://www.test.com:443/test")).isEqualTo("www.test.com:443"); + assertThat(parseHostAndPort("www.test.com:80/test")).isEqualTo("www.test.com:80"); + assertThat(parseHostAndPort("http://www.test.com")).isEqualTo("www.test.com"); + assertThat(parseHostAndPort("www.test.com")).isEqualTo("www.test.com"); + assertThat(parseHostAndPort("httpbin.org:443/get")).isEqualTo("httpbin.org:443"); + assertThat(parseHostAndPort("")).isEqualTo(""); + assertThat(parseHostAndPort("invalid")).isEqualTo("invalid"); + assertThat(parseHostAndPort("invalid://")).isEqualTo("invalid:"); + } + + @Test + public void testAddNewViaHeader() { + HttpMessage httpMessage = new DefaultFullHttpRequest(HTTP_1_1, HttpMethod.GET, "/endpoint"); + ProxyUtils.addVia(httpMessage, "hostname"); + + List viaHeaders = httpMessage.headers().getAll(HttpHeaderNames.VIA); + assertThat(viaHeaders).containsExactly("1.1 hostname"); + } + + @Test + public void testGetHeaderValuesWhenHeadersAreEmpty() { + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).isEmpty(); + } + + @Test + public void testGetHeaderValuesWhenTwoHeadersWithNoValuesArePresent() { + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK); + message.headers().add(TRANSFER_ENCODING, ""); + message.headers().add(TRANSFER_ENCODING, ""); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).isEmpty(); + } + + @Test + public void testGetHeaderValuesWhenSingleHeaderValueIsPresent() { + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK); + message.headers().add(TRANSFER_ENCODING, "chunked"); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).containsExactly("chunked"); + } + + @Test + public void testGetHeaderValuesWhenSingleHeaderValueWithExtraSpacesIsPresent() { + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK, false); + message.headers().add(TRANSFER_ENCODING, " chunked , "); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).containsExactly("chunked"); + } + + @Test + public void testGetHeaderValuesWhenTwoCommaSeparatedValuesInOneHeaderLineArePresent() { + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK); + message.headers().add(TRANSFER_ENCODING, "compress, gzip"); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).containsExactly("compress", "gzip"); + } + + @Test + public void + testGetHeaderValuesWhenTwoCommaSeparatedValuesInOneHeaderLineWithSpuriousCommaAndSpaceArePresent() { + // two comma-separated values in one header line with a spurious ',' and space. see RFC 7230 + // section 7 + // for information on empty list items (not all of which are valid header-values). + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK); + message.headers().add(TRANSFER_ENCODING, "compress, gzip, ,"); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).containsExactly("compress", "gzip"); + } + + @Test + public void testGetHeaderValuesWhenTwoValuesInTwoSeparateHeaderLinesArePresent() { + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK); + message.headers().add(TRANSFER_ENCODING, "gzip"); + message.headers().add(TRANSFER_ENCODING, "chunked"); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).containsExactly("gzip", "chunked"); + } + + @Test + public void + testGetHeaderValuesWhenMultipleCommaSeparatedValuesInTwoSeparateHeaderLinesArePresent() { + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK); + message.headers().add(TRANSFER_ENCODING, "gzip, compress"); + message.headers().add(TRANSFER_ENCODING, "deflate, gzip"); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders).containsExactly("gzip", "compress", "deflate", "gzip"); + } + + @Test + public void + testGetHeaderValuesWhenMultipleCommaSeparatedValuesInMultipleSeparateHeaderLinesArePresent() { + // multiple comma-separated values in multiple header lines with spurious spaces, commas, + // and tabs (horizontal tabs are defined as optional whitespace in RFC 7230 section 3.2.3) + DefaultHttpMessage message = new DefaultHttpResponse(HTTP_1_1, OK, false); + message.headers().add(TRANSFER_ENCODING, " gzip,compress,"); + message.headers().add(TRANSFER_ENCODING, "\tdeflate\t, gzip, "); + message.headers().add(TRANSFER_ENCODING, ",gzip,,deflate,\t, ,"); + List commaSeparatedHeaders = + ProxyUtils.getAllCommaSeparatedHeaderValues(TRANSFER_ENCODING, message); + assertThat(commaSeparatedHeaders) + .containsExactly("gzip", "compress", "deflate", "gzip", "gzip", "deflate"); + } + + @Test + public void testIsResponseSelfTerminating() { + HttpResponse httpResponse; + boolean isResponseSelfTerminating; + + // test cases from the scenarios listed in RFC 2616, section 4.4 + // #1: 1.Any response message which "MUST NOT" include a message-body (such as the 1xx, 204, and + // 304 responses and any response to a HEAD request) is always terminated by the first empty + // line after the header fields, regardless of the entity-header fields present in the message. + httpResponse = new DefaultHttpResponse(HTTP_1_1, HttpResponseStatus.CONTINUE); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + httpResponse = new DefaultHttpResponse(HTTP_1_1, SWITCHING_PROTOCOLS); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + httpResponse = new DefaultHttpResponse(HTTP_1_1, HttpResponseStatus.NO_CONTENT); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + httpResponse = new DefaultHttpResponse(HTTP_1_1, HttpResponseStatus.RESET_CONTENT); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + httpResponse = new DefaultHttpResponse(HTTP_1_1, HttpResponseStatus.NOT_MODIFIED); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + // #2: 2.If a Transfer-Encoding header field (section 14.41) is present and has any value other + // than "identity", then the transfer-length is defined by use of the "chunked" transfer-coding + // (section 3.6), unless the message is terminated by closing the connection. + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(TRANSFER_ENCODING, "chunked"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(TRANSFER_ENCODING, "gzip, chunked"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + // chunked encoding is not last, so not self terminating + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(TRANSFER_ENCODING, "chunked, gzip"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isFalse(); + + // four encodings on two lines, chunked is not last, so not self terminating + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(TRANSFER_ENCODING, "gzip, chunked"); + httpResponse.headers().add(TRANSFER_ENCODING, "deflate, gzip"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isFalse(); + + // three encodings on two lines, chunked is last, so self terminating + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(TRANSFER_ENCODING, "gzip"); + httpResponse.headers().add(TRANSFER_ENCODING, "deflate,chunked"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + // #3: 3.If a Content-Length header field (section 14.13) is present, its decimal value in + // OCTETs represents both the entity-length and the transfer-length. + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(HttpHeaderNames.CONTENT_LENGTH, "15"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + // continuing #3: If a message is received with both a Transfer-Encoding header field and a + // Content-Length header field, the latter MUST be ignored. + + // chunked is last Transfer-Encoding, so message is self-terminating + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(TRANSFER_ENCODING, "gzip, chunked"); + httpResponse.headers().add(HttpHeaderNames.CONTENT_LENGTH, "15"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isTrue(); + + // chunked is not last Transfer-Encoding, so message is not self-terminating, since + // Content-Length is ignored + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + httpResponse.headers().add(TRANSFER_ENCODING, "gzip"); + httpResponse.headers().add(HttpHeaderNames.CONTENT_LENGTH, "15"); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isFalse(); + + // without any of the above conditions, the message should not be self-terminating + // (multipart/byteranges is ignored, see note in method javadoc) + httpResponse = new DefaultHttpResponse(HTTP_1_1, OK); + isResponseSelfTerminating = ProxyUtils.isResponseSelfTerminating(httpResponse); + assertThat(isResponseSelfTerminating).isFalse(); + } + + @Test + public void testAddNewViaHeaderToExistingViaHeader() { + HttpMessage httpMessage = new DefaultFullHttpRequest(HTTP_1_1, HttpMethod.GET, "/endpoint"); + httpMessage.headers().add(HttpHeaderNames.VIA, "1.1 otherproxy"); + ProxyUtils.addVia(httpMessage, "hostname"); + + List viaHeaders = httpMessage.headers().getAll(HttpHeaderNames.VIA); + assertThat(viaHeaders).containsExactly("1.1 otherproxy", "1.1 hostname"); + } + + @Test + @DisplayName("Incorrect header tokens") + public void testSplitCommaSeparatedHeaderValues_incorrect_header_tokens() { + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues("one")).containsExactly("one"); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues("one,two,three")) + .containsExactly("one", "two", "three"); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues("one, two, three")) + .containsExactly("one", "two", "three"); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues(" one,two, three ")) + .containsExactly("one", "two", "three"); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues("\t\tone ,\t two, three\t")) + .containsExactly("one", "two", "three"); + } + + @Test + @DisplayName("Expected no header tokens") + public void testSplitCommaSeparatedHeaderValues_expected_no_header_tokens() { + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues("")).isEmpty(); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues(",")).isEmpty(); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues(" ")).isEmpty(); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues("\t")).isEmpty(); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues(" \t \t ")).isEmpty(); + assertThat(ProxyUtils.splitCommaSeparatedHeaderValues(" , ,\t, ")).isEmpty(); + } + + @Test + void request_isSwitchingToWebSocketProtocol() { + HttpRequest request = new DefaultFullHttpRequest(HTTP_1_1, HttpMethod.GET, "/endpoint"); + request.headers().add(HttpHeaderNames.HOST, "echo.websocket.org"); + request.headers().add(HttpHeaderNames.ORIGIN, "https://tests.w"); + request.headers().add(HttpHeaderNames.SEC_WEBSOCKET_EXTENSIONS, "permessage-deflate"); + request.headers().add(HttpHeaderNames.SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ=="); + request.headers().add("Sec-Fetch-Mode", "websocket"); + request.headers().add(HttpHeaderNames.CONNECTION, HttpHeaderNames.UPGRADE); + request.headers().add(HttpHeaderNames.UPGRADE, "websocket"); + + assertThat(ProxyUtils.isSwitchingToWebSocketProtocol(request)).isTrue(); + } + + @Test + void response_isSwitchingToWebSocketProtocol() { + HttpResponse response = new DefaultFullHttpResponse(HTTP_1_1, SWITCHING_PROTOCOLS); + response.headers().add(HttpHeaderNames.UPGRADE, "websocket"); + response.headers().add(HttpHeaderNames.CONNECTION, HttpHeaderNames.UPGRADE); + response.headers().add(HttpHeaderNames.SEC_WEBSOCKET_ACCEPT, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + + assertThat(ProxyUtils.isSwitchingToWebSocketProtocol(response)).isTrue(); + } + + /** Verifies that 'sdch' is removed from the 'Accept-Encoding' header list. */ + @Test + public void testRemoveSdchEncoding() { + final List emptyList = new ArrayList<>(); + // Various cases where 'sdch' is not present within the accepted + // encodings list + assertRemoveSdchEncoding(singletonList(""), emptyList); + assertRemoveSdchEncoding(singletonList("gzip"), singletonList("gzip")); + + assertRemoveSdchEncoding( + Arrays.asList("gzip", "deflate", "br"), Arrays.asList("gzip", "deflate", "br")); + assertRemoveSdchEncoding( + singletonList("gzip, deflate, br"), singletonList("gzip, deflate, br")); + + // Various cases where 'sdch' is present within the accepted encodings + // list + assertRemoveSdchEncoding(singletonList("sdch"), emptyList); + assertRemoveSdchEncoding(singletonList("SDCH"), emptyList); + + assertRemoveSdchEncoding(Arrays.asList("sdch", "gzip"), singletonList("gzip")); + assertRemoveSdchEncoding(singletonList("sdch, gzip"), singletonList("gzip")); + + assertRemoveSdchEncoding( + Arrays.asList("gzip", "sdch", "deflate"), Arrays.asList("gzip", "deflate")); + assertRemoveSdchEncoding(singletonList("gzip, sdch, deflate"), singletonList("gzip, deflate")); + assertRemoveSdchEncoding(singletonList("gzip,deflate,sdch"), singletonList("gzip,deflate")); + + assertRemoveSdchEncoding( + Arrays.asList("gzip", "deflate, sdch", "br"), Arrays.asList("gzip", "deflate", "br")); + } + + /** + * Helper method that asserts that 'sdch' is removed from the 'Accept-Encoding' header. + * + * @param inputEncodings The input list that maps to the values of the 'Accept-Encoding' header + * that should be used as the basis for the assertion check. + * @param expectedEncodings The list containing the expected values of the 'Accept-Encoding' + * header after the 'sdch' encoding is removed. + */ + private void assertRemoveSdchEncoding( + List inputEncodings, List expectedEncodings) { + HttpHeaders headers = new DefaultHttpHeaders(); + + for (String encoding : inputEncodings) { + headers.add(HttpHeaderNames.ACCEPT_ENCODING, encoding); } - @Test - public void testAddNewViaHeaderToExistingViaHeader() { - String hostname = "hostname"; - - HttpMessage httpMessage = new DefaultFullHttpRequest(HttpVersion.HTTP_1_1, HttpMethod.GET, "/endpoint"); - httpMessage.headers().add(HttpHeaderNames.VIA, "1.1 otherproxy"); - ProxyUtils.addVia(httpMessage, hostname); - - List viaHeaders = httpMessage.headers().getAll(HttpHeaderNames.VIA); - assertThat(viaHeaders, hasSize(2)); - - assertEquals("1.1 otherproxy", viaHeaders.get(0)); - - String expectedViaHeader = "1.1 " + hostname; - assertEquals(expectedViaHeader, viaHeaders.get(1)); - } - - @Test - public void testSplitCommaSeparatedHeaderValues() { - assertThat("Incorrect header tokens", ProxyUtils.splitCommaSeparatedHeaderValues("one"), contains("one")); - assertThat("Incorrect header tokens", ProxyUtils.splitCommaSeparatedHeaderValues("one,two,three"), contains("one", "two", "three")); - assertThat("Incorrect header tokens", ProxyUtils.splitCommaSeparatedHeaderValues("one, two, three"), contains("one", "two", "three")); - assertThat("Incorrect header tokens", ProxyUtils.splitCommaSeparatedHeaderValues(" one,two, three "), contains("one", "two", "three")); - assertThat("Incorrect header tokens", ProxyUtils.splitCommaSeparatedHeaderValues("\t\tone ,\t two, three\t"), contains("one", "two", "three")); - - assertThat("Expected no header tokens", ProxyUtils.splitCommaSeparatedHeaderValues(""), empty()); - assertThat("Expected no header tokens", ProxyUtils.splitCommaSeparatedHeaderValues(","), empty()); - assertThat("Expected no header tokens", ProxyUtils.splitCommaSeparatedHeaderValues(" "), empty()); - assertThat("Expected no header tokens", ProxyUtils.splitCommaSeparatedHeaderValues("\t"), empty()); - assertThat("Expected no header tokens", ProxyUtils.splitCommaSeparatedHeaderValues(" \t \t "), empty()); - assertThat("Expected no header tokens", ProxyUtils.splitCommaSeparatedHeaderValues(" , ,\t, "), empty()); - } - - /** - * Verifies that 'sdch' is removed from the 'Accept-Encoding' header list. - */ - @Test - public void testRemoveSdchEncoding() { - final List emptyList = new ArrayList<>(); - // Various cases where 'sdch' is not present within the accepted - // encodings list - assertRemoveSdchEncoding(singletonList(""), emptyList); - assertRemoveSdchEncoding(singletonList("gzip"), singletonList("gzip")); - - assertRemoveSdchEncoding(Arrays.asList("gzip", "deflate", "br"), Arrays.asList("gzip", "deflate", "br")); - assertRemoveSdchEncoding(singletonList("gzip, deflate, br"), singletonList("gzip, deflate, br")); - - // Various cases where 'sdch' is present within the accepted encodings - // list - assertRemoveSdchEncoding(singletonList("sdch"), emptyList); - assertRemoveSdchEncoding(singletonList("SDCH"), emptyList); - - assertRemoveSdchEncoding(Arrays.asList("sdch", "gzip"), singletonList("gzip")); - assertRemoveSdchEncoding(singletonList("sdch, gzip"), singletonList("gzip")); - - assertRemoveSdchEncoding(Arrays.asList("gzip", "sdch", "deflate"), Arrays.asList("gzip", "deflate")); - assertRemoveSdchEncoding(singletonList("gzip, sdch, deflate"), singletonList("gzip, deflate")); - assertRemoveSdchEncoding(singletonList("gzip,deflate,sdch"), singletonList("gzip,deflate")); - - assertRemoveSdchEncoding(Arrays.asList("gzip", "deflate, sdch", "br"), Arrays.asList("gzip", "deflate", "br")); - } - - /** - * Helper method that asserts that 'sdch' is removed from the - * 'Accept-Encoding' header. - * - * @param inputEncodings The input list that maps to the values of the - * 'Accept-Encoding' header that should be used as the basis for the - * assertion check. - * @param expectedEncodings The list containing the expected values of the - * 'Accept-Encoding' header after the 'sdch' encoding is removed. - */ - private void assertRemoveSdchEncoding(List inputEncodings, List expectedEncodings) { - HttpHeaders headers = new DefaultHttpHeaders(); - - for (String encoding : inputEncodings) { - headers.add(HttpHeaderNames.ACCEPT_ENCODING, encoding); - } - - ProxyUtils.removeSdchEncoding(headers); - assertEquals(expectedEncodings, headers.getAll(HttpHeaderNames.ACCEPT_ENCODING)); - } + ProxyUtils.removeSdchEncoding(headers); + assertThat(headers.getAll(ACCEPT_ENCODING)).isEqualTo(expectedEncodings); + } } diff --git a/src/test/java/org/littleshoot/proxy/impl/RealChainedProxyAuthenticationTest.java b/src/test/java/org/littleshoot/proxy/impl/RealChainedProxyAuthenticationTest.java new file mode 100644 index 00000000..a773d5da --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/RealChainedProxyAuthenticationTest.java @@ -0,0 +1,402 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.*; + +import io.netty.handler.codec.http.HttpHeaders; +import io.netty.handler.codec.http.HttpMethod; +import io.netty.handler.codec.http.HttpRequest; +import io.netty.handler.codec.http.HttpResponse; +import io.netty.handler.codec.http.HttpResponseStatus; +import io.netty.handler.codec.http.HttpVersion; +import java.util.Base64; +import java.util.concurrent.atomic.AtomicBoolean; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ChainedProxy; +import org.littleshoot.proxy.ChainedProxyType; +import org.mockito.Mockito; + +/** + * Real unit tests that exercise the actual authentication code paths for all four chained proxy + * authentication scenarios. + */ +class RealChainedProxyAuthenticationTest { + + private ClientToProxyConnection clientToProxyConnection; + private ProxyToServerConnection proxyToServerConnection; + private AtomicBoolean authenticated; + + @BeforeEach + void setUp() { + // Mock the connections and dependencies + clientToProxyConnection = Mockito.mock(ClientToProxyConnection.class, CALLS_REAL_METHODS); + proxyToServerConnection = Mockito.mock(ProxyToServerConnection.class); + authenticated = new AtomicBoolean(false); + + // Set up the mock to return our test values + when(clientToProxyConnection.getAuthenticated()).thenReturn(authenticated); + + // Use reflection to set the currentServerConnection field directly since + // the shouldPreserveProxyAuthorizationForUpstream method accesses it directly + try { + java.lang.reflect.Field currentServerConnectionField = + ClientToProxyConnection.class.getDeclaredField("currentServerConnection"); + currentServerConnectionField.setAccessible(true); + currentServerConnectionField.set(clientToProxyConnection, proxyToServerConnection); + + // Initialize the LOG field to avoid NPE in logging calls + // LOG field is in the parent ProxyConnection class + java.lang.reflect.Field logField = ProxyConnection.class.getDeclaredField("LOG"); + logField.setAccessible(true); + logField.set(clientToProxyConnection, new ProxyConnectionLogger(clientToProxyConnection)); + } catch (Exception e) { + throw new RuntimeException("Failed to set up mock fields", e); + } + } + + /** + * Scenario 1: LittleProxy handles auth, next proxy does not + * + *

Client → [Proxy-Authorization: clientCreds] → LittleProxy (authenticates, removes header) → + * [No auth header] → Upstream Proxy (no auth needed) → ✅ Success + */ + @Test + void testScenario1_RealCodePath_LittleProxyAuthOnly() { + System.out.println( + "=== Real Test Scenario 1: LittleProxy handles auth, next proxy does not ==="); + + // Create a real request with client credentials + HttpRequest request = + new io.netty.handler.codec.http.DefaultHttpRequest( + HttpVersion.HTTP_1_1, HttpMethod.GET, "/test"); + request + .headers() + .set( + "Proxy-Authorization", + "Basic " + Base64.getEncoder().encodeToString("clientUser:clientPass".getBytes())); + + System.out.println("1. Client request: " + request.headers()); + + // Simulate client authentication (header would be removed by real auth logic) + request.headers().remove("Proxy-Authorization"); + + // Mock: upstream proxy doesn't require authentication + when(proxyToServerConnection.hasUpstreamChainedProxy()).thenReturn(true); + ChainedProxy noAuthProxy = createNoAuthProxy(); + when(proxyToServerConnection.getChainedProxy()).thenReturn(noAuthProxy); + + // Test the real shouldPreserveProxyAuthorizationForUpstream method + boolean shouldPreserve = clientToProxyConnection.shouldPreserveProxyAuthorizationForUpstream(); + + System.out.println("2. Should preserve Proxy-Authorization: " + shouldPreserve); + + // Verify: should NOT preserve (upstream proxy doesn't require auth) + assertThat(shouldPreserve) + .as("Should not preserve Proxy-Authorization when upstream proxy doesn't require auth") + .isFalse(); + + // Test the real stripHopByHopHeaders method + HttpHeaders headersCopy = request.headers().copy(); + clientToProxyConnection.stripHopByHopHeaders(headersCopy); + + System.out.println("3. After header processing: " + headersCopy); + + // Verify: no Proxy-Authorization header should be present + assertThat(headersCopy.contains("Proxy-Authorization")) + .as( + "No Proxy-Authorization header should be forwarded to upstream proxy that doesn't require auth") + .isFalse(); + + // Verify: shouldPreserve is false when no upstream auth is required + assertThat(shouldPreserve) + .as("Should not preserve Proxy-Authorization when upstream proxy doesn't require auth") + .isFalse(); + + System.out.println("✅ Scenario 1 real code path test passed"); + } + + /** + * Scenario 2: Both LittleProxy and next proxy handle auth + * + *

Client → [Proxy-Authorization: clientUser:clientPass] → LittleProxy (authenticates) → + * [Proxy-Authorization: upstreamUser:upstreamPass] → Upstream Proxy → ✅ Success + */ + @Test + void testScenario2_RealCodePath_BothHandleAuth() { + System.out.println("=== Real Test Scenario 2: Both proxies handle auth ==="); + + // Define different credentials for each proxy + String clientCredentials = "clientUser:clientPass"; + String upstreamCredentials = "upstreamUser:upstreamPass"; + + // Create a real request with CLIENT credentials for LittleProxy + HttpRequest request = + new io.netty.handler.codec.http.DefaultHttpRequest( + HttpVersion.HTTP_1_1, HttpMethod.GET, "/test"); + request + .headers() + .set( + "Proxy-Authorization", + "Basic " + Base64.getEncoder().encodeToString(clientCredentials.getBytes())); + + System.out.println("1. Client request with CLIENT credentials: " + request.headers()); + System.out.println(" Client credentials: " + clientCredentials); + + // Simulate client authentication (header would be removed by real auth logic) + request.headers().remove("Proxy-Authorization"); + System.out.println("2. After LittleProxy authentication: header removed"); + + // Mock: upstream proxy requires authentication with DIFFERENT credentials + when(proxyToServerConnection.hasUpstreamChainedProxy()).thenReturn(true); + ChainedProxy authRequiredProxy = createUpstreamAuthProxy(upstreamCredentials); + when(proxyToServerConnection.getChainedProxy()).thenReturn(authRequiredProxy); + + System.out.println(" Upstream credentials: " + upstreamCredentials); + + // Test the real shouldPreserveProxyAuthorizationForUpstream method + boolean shouldPreserve = clientToProxyConnection.shouldPreserveProxyAuthorizationForUpstream(); + + System.out.println("3. Should preserve Proxy-Authorization: " + shouldPreserve); + + // Verify: SHOULD preserve (upstream proxy requires auth) + assertThat(shouldPreserve) + .as("Should preserve Proxy-Authorization when upstream proxy requires auth") + .isTrue(); + + // Test the real stripHopByHopHeaders method + HttpHeaders headersCopy = request.headers().copy(); + clientToProxyConnection.stripHopByHopHeaders(headersCopy); + + System.out.println("4. After stripping hop-by-hop headers: " + headersCopy); + + // Verify: Proxy-Authorization header should be removed (was already removed by client auth + // simulation) + assertThat(headersCopy.contains("Proxy-Authorization")) + .as("Proxy-Authorization header should be removed after stripping hop-by-hop headers") + .isFalse(); + + // Test the real addUpstreamProxyAuthorization method + clientToProxyConnection.addUpstreamProxyAuthorization(headersCopy); + + System.out.println("5. After adding UPSTREAM credentials: " + headersCopy); + + // Verify: upstream credentials were added + assertThat(headersCopy.contains("Proxy-Authorization")) + .as("Proxy-Authorization header should be added with upstream credentials") + .isTrue(); + + // Verify: credentials are for upstream proxy, not client + String authHeader = headersCopy.get("Proxy-Authorization"); + String expectedUpstreamAuth = + "Basic " + Base64.getEncoder().encodeToString(upstreamCredentials.getBytes()); + + assertThat(authHeader) + .as("Should contain upstream credentials, not client credentials") + .isEqualTo(expectedUpstreamAuth) + .doesNotContain(clientCredentials); + + // Explicit verification that credentials are different + assertThat(authHeader) + .as("Upstream credentials should be different from client credentials") + .isNotEqualTo("Basic " + Base64.getEncoder().encodeToString(clientCredentials.getBytes())); + + System.out.println("✅ Scenario 2 real code path test passed - different credentials verified"); + } + + /** + * Scenario 3: LittleProxy does not handle auth, next proxy does + * + *

Client → [Proxy-Authorization: upstreamUser:upstreamPass] → LittleProxy (no auth, preserves + * header) → [Proxy-Authorization: upstreamUser:upstreamPass] → Upstream Proxy (authenticates) → ✅ + * Success + */ + @Test + void testScenario3_RealCodePath_NoLittleProxyAuth() { + System.out.println("=== Real Test Scenario 3: LittleProxy no auth, next proxy does ==="); + + // Create a real request with upstream credentials (no client auth) + HttpRequest request = + new io.netty.handler.codec.http.DefaultHttpRequest( + HttpVersion.HTTP_1_1, HttpMethod.GET, "/test"); + request + .headers() + .set( + "Proxy-Authorization", + "Basic " + Base64.getEncoder().encodeToString("upstreamUser:upstreamPass".getBytes())); + + System.out.println("1. Client request (no LittleProxy auth): " + request.headers()); + + // Mock: LittleProxy doesn't authenticate (authenticated remains false) + authenticated.set(false); + + // Mock: upstream proxy requires authentication + when(proxyToServerConnection.hasUpstreamChainedProxy()).thenReturn(true); + ChainedProxy authRequiredProxy = createAuthRequiredProxy(); + when(proxyToServerConnection.getChainedProxy()).thenReturn(authRequiredProxy); + + // Test the real shouldPreserveProxyAuthorizationForUpstream method + boolean shouldPreserve = clientToProxyConnection.shouldPreserveProxyAuthorizationForUpstream(); + + System.out.println("2. Should preserve Proxy-Authorization: " + shouldPreserve); + + // Verify: SHOULD preserve (upstream proxy requires auth) + assertThat(shouldPreserve) + .as("Should preserve Proxy-Authorization when upstream proxy requires auth") + .isTrue(); + + // Test the real stripHopByHopHeaders method - now always removes Proxy-Authorization + HttpHeaders headersCopy = request.headers().copy(); + clientToProxyConnection.stripHopByHopHeaders(headersCopy); + + System.out.println("3. After stripping hop-by-hop headers: " + headersCopy); + + // Verify: Proxy-Authorization header is removed (hop-by-hop headers are always stripped) + // The header is added back by addUpstreamProxyAuthorization when needed + assertThat(headersCopy.contains("Proxy-Authorization")) + .as("Proxy-Authorization header should be removed by stripHopByHopHeaders") + .isFalse(); + + // Since shouldPreserve is true, addUpstreamProxyAuthorization would add the header + // but we test that separately + System.out.println("✅ Scenario 3 real code path test passed"); + } + + /** + * Scenario 4: Neither proxy handles auth + * + *

Client → [No auth header] → LittleProxy (no auth) → [No auth header] → Upstream Proxy (no + * auth) → ✅ Success + */ + @Test + void testScenario4_RealCodePath_NoAuthAnywhere() { + System.out.println("=== Real Test Scenario 4: Neither proxy handles auth ==="); + + // Create a real request without any authentication + HttpRequest request = + new io.netty.handler.codec.http.DefaultHttpRequest( + HttpVersion.HTTP_1_1, HttpMethod.GET, "/test"); + // No Proxy-Authorization header + + System.out.println("1. Client request (no auth): " + request.headers()); + + // Mock: no upstream chained proxy + when(proxyToServerConnection.hasUpstreamChainedProxy()).thenReturn(false); + + // Test the real shouldPreserveProxyAuthorizationForUpstream method + boolean shouldPreserve = clientToProxyConnection.shouldPreserveProxyAuthorizationForUpstream(); + + System.out.println("2. Should preserve Proxy-Authorization: " + shouldPreserve); + + // Verify: should NOT preserve (no upstream proxy) + assertThat(shouldPreserve) + .as("Should not preserve Proxy-Authorization when no upstream proxy") + .isFalse(); + + // Test the real stripHopByHopHeaders method (normal behavior) + HttpHeaders headersCopy = request.headers().copy(); + clientToProxyConnection.stripHopByHopHeaders(headersCopy); + + System.out.println("3. After normal header processing: " + headersCopy); + + // Verify: no Proxy-Authorization header should be present + assertThat(headersCopy.contains("Proxy-Authorization")) + .as("No Proxy-Authorization header should be present when no auth is required") + .isFalse(); + + // Verify: shouldPreserve is false when no upstream proxy + assertThat(shouldPreserve) + .as("Should not preserve Proxy-Authorization when no upstream proxy") + .isFalse(); + + System.out.println("✅ Scenario 4 real code path test passed"); + } + + /** Test upstream 407 response handling */ + @Test + void testUpstream407ResponseHandling() { + System.out.println("=== Real Test: Upstream 407 Response Handling ==="); + + // Mock: upstream proxy requires authentication + when(proxyToServerConnection.hasUpstreamChainedProxy()).thenReturn(true); + ChainedProxy authRequiredProxy = createAuthRequiredProxy(); + when(proxyToServerConnection.getChainedProxy()).thenReturn(authRequiredProxy); + + // Test various response types + + // 1. Test 407 response from upstream proxy + HttpResponse upstream407 = + new io.netty.handler.codec.http.DefaultHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.PROXY_AUTHENTICATION_REQUIRED); + + boolean isUpstream407 = + clientToProxyConnection.handleUpstreamProxyAuthenticationRequired(upstream407); + + System.out.println("1. Upstream 407 response handled: " + isUpstream407); + assertThat(isUpstream407).as("Should identify and handle upstream 407 response").isTrue(); + + // 2. Test 200 response (should not be handled as 407) + HttpResponse successResponse = + new io.netty.handler.codec.http.DefaultHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.OK); + + boolean isSuccessHandled = + clientToProxyConnection.handleUpstreamProxyAuthenticationRequired(successResponse); + + System.out.println("2. Success response handled as 407: " + isSuccessHandled); + assertThat(isSuccessHandled).as("Should not handle success response as 407").isFalse(); + + // 3. Test 401 response (should not be handled as 407) + HttpResponse unauthorizedResponse = + new io.netty.handler.codec.http.DefaultHttpResponse( + HttpVersion.HTTP_1_1, HttpResponseStatus.UNAUTHORIZED); + + boolean isUnauthorizedHandled = + clientToProxyConnection.handleUpstreamProxyAuthenticationRequired(unauthorizedResponse); + + System.out.println("3. 401 response handled as 407: " + isUnauthorizedHandled); + assertThat(isUnauthorizedHandled).as("Should not handle 401 response as 407").isFalse(); + + System.out.println("✅ Upstream 407 response handling test passed"); + } + + // Helper methods to create test proxies + + private ChainedProxy createUpstreamAuthProxy(String credentials) { + ChainedProxy proxy = Mockito.mock(ChainedProxy.class); + when(proxy.getChainedProxyType()).thenReturn(ChainedProxyType.HTTP); + + // Parse credentials for username and password + String[] parts = credentials.split(":", 2); + String username = parts[0]; + String password = parts.length > 1 ? parts[1] : ""; + + when(proxy.getUsername()).thenReturn(username); + when(proxy.getPassword()).thenReturn(password); + return proxy; + } + + private ChainedProxy createAuthRequiredProxy() { + ChainedProxy proxy = Mockito.mock(ChainedProxy.class); + when(proxy.getChainedProxyType()).thenReturn(ChainedProxyType.HTTP); + when(proxy.getUsername()).thenReturn("upstreamUser"); + when(proxy.getPassword()).thenReturn("upstreamPass"); + return proxy; + } + + private ChainedProxy createNoAuthProxy() { + ChainedProxy proxy = Mockito.mock(ChainedProxy.class); + when(proxy.getChainedProxyType()).thenReturn(ChainedProxyType.HTTP); + when(proxy.getUsername()).thenReturn(null); + when(proxy.getPassword()).thenReturn(null); + return proxy; + } + + private ChainedProxy createSocksProxy() { + ChainedProxy proxy = Mockito.mock(ChainedProxy.class); + when(proxy.getChainedProxyType()).thenReturn(ChainedProxyType.SOCKS4); + when(proxy.getUsername()).thenReturn("socksUser"); + when(proxy.getPassword()).thenReturn("socksPass"); + return proxy; + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ServerConnectionPoolConfigTest.java b/src/test/java/org/littleshoot/proxy/impl/ServerConnectionPoolConfigTest.java new file mode 100644 index 00000000..39e42e0c --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ServerConnectionPoolConfigTest.java @@ -0,0 +1,124 @@ +package org.littleshoot.proxy.impl; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import java.time.Duration; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.ServerConnectionPoolType; + +class ServerConnectionPoolConfigTest { + + private final ServerConnectionPoolConfig config = new ServerConnectionPoolConfig(); + + @Test + void setPoolTypeRejectsNull() { + assertThatThrownBy(() -> config.setPoolType(null)) + .isInstanceOf(NullPointerException.class) + .hasMessageContaining("poolType"); + } + + @Test + void setPoolTypeAcceptsValidValue() { + config.setPoolType(ServerConnectionPoolType.CONCURRENT_MAP); + assertThat(config.getPoolType()).isEqualTo(ServerConnectionPoolType.CONCURRENT_MAP); + } + + @Test + void setMaxConnectionsPerHostRejectsZero() { + assertThatThrownBy(() -> config.setMaxConnectionsPerHost(0)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void setMaxConnectionsPerHostRejectsNegative() { + assertThatThrownBy(() -> config.setMaxConnectionsPerHost(-1)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void setMaxConnectionsPerHostAcceptsPositive() { + config.setMaxConnectionsPerHost(5); + assertThat(config.getMaxConnectionsPerHost()).isEqualTo(5); + } + + @Test + void setMaxConnectionsRejectsZero() { + assertThatThrownBy(() -> config.setMaxConnections(0)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void setMaxConnectionsRejectsNegative() { + assertThatThrownBy(() -> config.setMaxConnections(-1)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void setMaxConnectionsAcceptsPositive() { + config.setMaxConnections(100); + assertThat(config.getMaxConnections()).isEqualTo(100); + } + + @Test + void idleTimeoutIsNullable() { + assertThat(config.getIdleTimeout()).isNull(); + config.setIdleTimeout(Duration.ofSeconds(30)); + assertThat(config.getIdleTimeout()).isEqualTo(Duration.ofSeconds(30)); + config.setIdleTimeout(null); + assertThat(config.getIdleTimeout()).isNull(); + } + + @Test + void defaults() { + assertThat(config.isEnabled()).isFalse(); + assertThat(config.getPoolType()).isEqualTo(ServerConnectionPoolType.CONCURRENT_MAP); + assertThat(config.getMaxConnectionsPerHost()).isEqualTo(10); + assertThat(config.getMaxConnections()).isEqualTo(200); + assertThat(config.getIdleTimeout()).isNull(); + assertThat(config.isPoolSharedMitmConnections()).isFalse(); + assertThat(config.isPoolPerRequestInMitm()).isFalse(); + } + + @Test + void poolSharedMitmConnectionsDefaultsToFalse() { + assertThat(new ServerConnectionPoolConfig().isPoolSharedMitmConnections()).isFalse(); + } + + @Test + void poolSharedMitmConnectionsSetterAndGetter() { + assertThat(config.setPoolSharedMitmConnections(true)).isSameAs(config); + assertThat(config.isPoolSharedMitmConnections()).isTrue(); + config.setPoolSharedMitmConnections(false); + assertThat(config.isPoolSharedMitmConnections()).isFalse(); + } + + @Test + void poolPerRequestInMitmDefaultsToFalse() { + assertThat(new ServerConnectionPoolConfig().isPoolPerRequestInMitm()).isFalse(); + } + + @Test + void poolPerRequestInMitmSetterAndGetter() { + assertThat(config.setPoolPerRequestInMitm(true)).isSameAs(config); + assertThat(config.isPoolPerRequestInMitm()).isTrue(); + config.setPoolPerRequestInMitm(false); + assertThat(config.isPoolPerRequestInMitm()).isFalse(); + } + + @Test + void poolSharedMitmConnectionsRequiresEnabledPool() { + // poolSharedMitmConnections can be set regardless of enabled state; + // the actual behavior depends on the enabled flag at runtime. + config.setEnabled(false).setPoolSharedMitmConnections(true); + assertThat(config.isPoolSharedMitmConnections()).isTrue(); + } + + @Test + void poolPerRequestInMitmRequiresPoolSharedMitmConnections() { + // poolPerRequestInMitm requires poolSharedMitmConnections=true at runtime + // but the config itself doesn't enforce this invariant. + config.setPoolSharedMitmConnections(false).setPoolPerRequestInMitm(true); + assertThat(config.isPoolPerRequestInMitm()).isTrue(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/impl/ServerGroupTest.java b/src/test/java/org/littleshoot/proxy/impl/ServerGroupTest.java new file mode 100644 index 00000000..31bf289b --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/impl/ServerGroupTest.java @@ -0,0 +1,40 @@ +package org.littleshoot.proxy.impl; + +import static org.junit.jupiter.api.Assertions.*; + +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.HttpProxyServer; + +class ServerGroupTest { + private ServerGroup serverGroup; + + private void startAndStopProxyServer() { + HttpProxyServer proxyServer = + DefaultHttpProxyServer.bootstrap().withPort(0).withServerGroup(serverGroup).start(); + proxyServer.stop(); + } + + @Test + void autoStop() { + serverGroup = new ServerGroup("Test", 4, 4, 4); + startAndStopProxyServer(); + assertTrue(serverGroup.isStopped(), "serverGroup.isStopped"); + assertThrows(IllegalStateException.class, this::startAndStopProxyServer); + } + + @Test + void manualStop() { + serverGroup = new ServerGroup("Test", 4, 4, 4, false); + startAndStopProxyServer(); + assertFalse(serverGroup.isStopped(), "serverGroup.isStopped"); + startAndStopProxyServer(); + } + + @AfterEach + void shutdown() { + if (serverGroup != null) { + serverGroup.shutdown(false); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/test/EnableThreadDump.java b/src/test/java/org/littleshoot/proxy/test/EnableThreadDump.java new file mode 100644 index 00000000..08326944 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/test/EnableThreadDump.java @@ -0,0 +1,20 @@ +package org.littleshoot.proxy.test; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +/** + * Annotation to enable thread dump generation for a test class or test method. + * + *

When present on a test class, thread dumps will be generated for all test methods in that + * class. When present on a specific test method, thread dumps will only be generated for that + * method. + * + *

This annotation is only effective when the {@link ThreadDumpExtension} is registered (either + * via service loader or via {@code @ExtendWith(ThreadDumpExtension.class)}). + */ +@Target({ElementType.TYPE, ElementType.METHOD}) +@Retention(RetentionPolicy.RUNTIME) +public @interface EnableThreadDump {} diff --git a/src/test/java/org/littleshoot/proxy/test/HttpClientUtil.java b/src/test/java/org/littleshoot/proxy/test/HttpClientUtil.java index 81187c29..034bb1e8 100644 --- a/src/test/java/org/littleshoot/proxy/test/HttpClientUtil.java +++ b/src/test/java/org/littleshoot/proxy/test/HttpClientUtil.java @@ -1,5 +1,8 @@ package org.littleshoot.proxy.test; +import static java.nio.charset.StandardCharsets.UTF_8; + +import java.io.IOException; import org.apache.http.HttpEntity; import org.apache.http.HttpHost; import org.apache.http.client.methods.HttpGet; @@ -11,85 +14,86 @@ import org.apache.http.util.EntityUtils; import org.littleshoot.proxy.HttpProxyServer; -import java.io.IOException; -import java.nio.charset.Charset; - /** - * Utility methods for creating HTTP clients and sending requests to servers via a LittleProxy instance. + * Utility methods for creating HTTP clients and sending requests to servers via a LittleProxy + * instance. */ public class HttpClientUtil { - /** - * Creates a new HTTP client that uses the specified LittleProxy instance to perform a GET to the specified URL. - * The HTTP client is closed and discarded after the request is completed. - * - * @param url URL to post to - * @param proxyServer LittleProxy instance through which the GET will be proxied - * @return the HttpResponse object encapsulating the response from the server - */ - public static org.apache.http.HttpResponse performHttpGet(String url, HttpProxyServer proxyServer) { - CloseableHttpClient http = buildHttpClient(proxyServer); + public static org.apache.http.HttpResponse performLocalHttpGet( + int port, String path, HttpProxyServer proxyServer) { + return performHttpGet("http://localhost:" + port + path, proxyServer); + } - HttpGet get = new HttpGet(url); + /** + * Creates a new HTTP client that uses the specified LittleProxy instance to perform a GET to the + * specified URL. The HTTP client is closed and discarded after the request is completed. + * + * @param url URL to post to + * @param proxyServer LittleProxy instance through which the GET will be proxied + * @return the HttpResponse object encapsulating the response from the server + */ + public static org.apache.http.HttpResponse performHttpGet( + String url, HttpProxyServer proxyServer) { + CloseableHttpClient http = buildHttpClient(proxyServer); - return performHttpRequest(http, get); - } + HttpGet get = new HttpGet(url); - /** - * Creates a new HTTP client that uses the specified LittleProxy instance to perform a POST of the specified size to - * to the URL. The POST body will consist of a meaningless UTF-8-encoded String (currently, the letter 'q'). The - * HTTP client is closed and discarded after the request is completed. - * - * @param url URL to post to - * @param postSizeInBytes size of the POST body - * @param proxyServer LittleProxy instance through which the POST will be proxied - * @return the HttpResponse from the server - */ - public static org.apache.http.HttpResponse performHttpPost(String url, int postSizeInBytes, HttpProxyServer proxyServer) { - CloseableHttpClient httpClient = buildHttpClient(proxyServer); + return performHttpRequest(http, get); + } - StringBuilder sb = new StringBuilder(); - for (int i = 0; i < postSizeInBytes; i++) { - sb.append('q'); - } + /** + * Creates a new HTTP client that uses the specified LittleProxy instance to perform a POST of the + * specified size to the URL. The POST body will consist of a meaningless UTF-8-encoded String + * (currently, the letter 'q'). The HTTP client is closed and discarded after the request is + * completed. + * + * @param url URL to post to + * @param postSizeInBytes size of the POST body + * @param proxyServer LittleProxy instance through which the POST will be proxied + */ + public static void performHttpPost(String url, int postSizeInBytes, HttpProxyServer proxyServer) { + CloseableHttpClient httpClient = buildHttpClient(proxyServer); - HttpPost post = new HttpPost(url); - post.setEntity(new StringEntity(sb.toString(), Charset.forName("UTF-8"))); + HttpPost post = new HttpPost(url); + post.setEntity(new StringEntity("q".repeat(Math.max(0, postSizeInBytes)), UTF_8)); - return performHttpRequest(httpClient, post); - } + performHttpRequest(httpClient, post); + } - /** - * Returns a new HttpClient that is configured to use the specified LittleProxy instance. - * - * @param proxyServer LittleProxy instance through which requests will be proxied - * @return new HttpClient - */ - private static CloseableHttpClient buildHttpClient(HttpProxyServer proxyServer) { - return HttpClients.custom() - .setProxy(new HttpHost("127.0.0.1", proxyServer.getListenAddress().getPort())) - .build(); - } + /** + * Returns a new HttpClient that is configured to use the specified LittleProxy instance. + * + * @param proxyServer LittleProxy instance through which requests will be proxied + * @return new HttpClient + */ + private static CloseableHttpClient buildHttpClient(HttpProxyServer proxyServer) { + return HttpClients.custom() + .setProxy(new HttpHost("127.0.0.1", proxyServer.getListenAddress().getPort())) + .build(); + } - /** - * Performs the specified request using the HTTP client. Consumes the response entity and shuts down the HTTP client - * after the request has been made. - * - * @param httpClient HTTP client to use - * @param request HTTP request to perform - * @return the HttpResponse from the server - */ - private static org.apache.http.HttpResponse performHttpRequest(CloseableHttpClient httpClient, HttpUriRequest request) { - try { - org.apache.http.HttpResponse hr = httpClient.execute(request); - HttpEntity responseEntity = hr.getEntity(); - EntityUtils.consume(responseEntity); + /** + * Performs the specified request using the HTTP client. Consumes the response entity and shuts + * down the HTTP client after the request has been made. + * + * @param httpClient HTTP client to use + * @param request HTTP request to perform + * @return the HttpResponse from the server + */ + private static org.apache.http.HttpResponse performHttpRequest( + CloseableHttpClient httpClient, HttpUriRequest request) { + try { + org.apache.http.HttpResponse hr = httpClient.execute(request); + HttpEntity responseEntity = hr.getEntity(); + EntityUtils.consume(responseEntity); - httpClient.close(); + httpClient.close(); - return hr; - } catch (IOException e) { - // this is test code; just let all exceptions bubble up the stack, which will cause the test to fail - throw new RuntimeException("Unable to perform HTTP request", e); - } + return hr; + } catch (IOException e) { + // this is a test code; just let all exceptions bubble up the stack, which will cause the test + // to fail + throw new RuntimeException("Unable to perform HTTP request", e); } + } } diff --git a/src/test/java/org/littleshoot/proxy/test/ServerErrorTest.java b/src/test/java/org/littleshoot/proxy/test/ServerErrorTest.java index 088cf687..0d511a57 100644 --- a/src/test/java/org/littleshoot/proxy/test/ServerErrorTest.java +++ b/src/test/java/org/littleshoot/proxy/test/ServerErrorTest.java @@ -1,9 +1,7 @@ package org.littleshoot.proxy.test; -import org.junit.After; -import org.junit.Test; -import org.littleshoot.proxy.HttpProxyServer; -import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import static org.assertj.core.api.Assertions.assertThat; +import static org.littleshoot.proxy.test.HttpClientUtil.performLocalHttpGet; import java.io.BufferedReader; import java.io.IOException; @@ -11,81 +9,87 @@ import java.io.PrintWriter; import java.net.ServerSocket; import java.net.Socket; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.Test; +import org.littleshoot.proxy.HttpProxyServer; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; -import static org.junit.Assert.assertEquals; - -public class ServerErrorTest { - private HttpProxyServer proxyServer; +public final class ServerErrorTest { + private HttpProxyServer proxyServer; - @Test - public void testInvalidServerResponse() throws IOException { - proxyServer = DefaultHttpProxyServer.bootstrap() - .withPort(0) - .start(); + @Test + public void testInvalidServerResponse() throws IOException { + proxyServer = DefaultHttpProxyServer.bootstrap().withPort(0).start(); - // we have to create our own socket here, since any proper http server (jetty, mockserver, etc.) won't allow us to - // send invalid responses. - try (ServerSocket socket = createServerWithBadResponse()) { - org.apache.http.HttpResponse response = HttpClientUtil.performHttpGet("http://localhost:" + socket.getLocalPort(), proxyServer); + // we have to create our own socket here, since any proper http server (jetty, + // wiremock, etc.) won't allow us to + // send invalid responses. + try (ServerSocket socket = createServerWithBadResponse()) { + org.apache.http.HttpResponse response = + performLocalHttpGet(socket.getLocalPort(), "/", proxyServer); - assertEquals("Expected to receive a 502 Bad Gateway after server responded with invalid response", 502, response.getStatusLine().getStatusCode()); - } + assertThat(response.getStatusLine().getStatusCode()) + .as("Expected to receive a 502 Bad Gateway after server responded with invalid response") + .isEqualTo(502); } + } - @After - public void tearDown() { - if (proxyServer != null) { - proxyServer.abort(); - } + @AfterEach + void tearDown() { + if (proxyServer != null) { + proxyServer.abort(); + } + } + + /** + * Creates a ServerSocket that will read an HTTP request and response with an invalid response. + * NOTE: the ServerSocket must be closed after the response is consumed. + */ + private static ServerSocket createServerWithBadResponse() { + final ServerSocket serverSocket; + try { + serverSocket = new ServerSocket(0); + } catch (IOException e) { + throw new RuntimeException(e); } - /** - * Creates a ServerSocket that will read an HTTP request and response with an invalid response. - * NOTE: the ServerSocket must be closed after the response is consumed. - */ - private static ServerSocket createServerWithBadResponse() { - final ServerSocket serverSocket; - try { - serverSocket = new ServerSocket(0); - } catch (IOException e) { - throw new RuntimeException(e); - } - - Runnable server = () -> { - try { - Socket socket = serverSocket.accept(); - - try (PrintWriter out = new PrintWriter(socket.getOutputStream(), true); - BufferedReader in = new BufferedReader(new InputStreamReader(socket.getInputStream()))) { - while (!in.readLine().isEmpty()) { - // read the request up to the double-CRLF - } - - // write a a response with an invalid HTTP version - out.write("HTTP/1.12312312312312411231231231 200 OK\r\n" + - "Connection: close\r\n" + - "Content-Length: 0\r\n" + - "\r\n"); - out.flush(); - } - } catch (IOException e) { - throw new RuntimeException(e); + Runnable server = + () -> { + try { + Socket socket = serverSocket.accept(); + + try (PrintWriter out = new PrintWriter(socket.getOutputStream(), true); + BufferedReader in = + new BufferedReader(new InputStreamReader(socket.getInputStream()))) { + while (!in.readLine().isEmpty()) { + // read the request up to the double-CRLF + } + + // write a response with an invalid HTTP version + out.write( + "HTTP/1.12312312312312411231231231 200 OK\r\n" + + "Connection: close\r\n" + + "Content-Length: 0\r\n" + + "\r\n"); + out.flush(); } + } catch (IOException e) { + throw new RuntimeException(e); + } }; - // start the server in a separate thread - Thread serverThread = new Thread(server); - serverThread.setDaemon(true); - serverThread.start(); - - // wait for the server to start - try { - Thread.sleep(500); - } catch (InterruptedException e) { - throw new RuntimeException(e); - } + // start the server in a separate thread + Thread serverThread = new Thread(server); + serverThread.setDaemon(true); + serverThread.start(); - return serverSocket; + // wait for the server to start + try { + Thread.sleep(500); + } catch (InterruptedException e) { + throw new RuntimeException(e); } + return serverSocket; + } } diff --git a/src/test/java/org/littleshoot/proxy/test/SocketClientUtil.java b/src/test/java/org/littleshoot/proxy/test/SocketClientUtil.java index ad4653d3..4374974e 100644 --- a/src/test/java/org/littleshoot/proxy/test/SocketClientUtil.java +++ b/src/test/java/org/littleshoot/proxy/test/SocketClientUtil.java @@ -1,6 +1,6 @@ package org.littleshoot.proxy.test; -import org.littleshoot.proxy.HttpProxyServer; +import static java.nio.charset.StandardCharsets.UTF_8; import java.io.EOFException; import java.io.IOException; @@ -10,98 +10,101 @@ import java.net.Socket; import java.net.SocketException; import java.net.SocketTimeoutException; -import java.nio.charset.Charset; +import org.littleshoot.proxy.HttpProxyServer; -/** - * Utilities for interacting with the proxy server using sockets. - */ +/** Utilities for interacting with the proxy server using sockets. */ public class SocketClientUtil { - /** - * Writes and flushes the UTF-8 encoded contents of a String to a socket. - * - * @param string string to write - * @param socket socket to write to - */ - public static void writeStringToSocket(String string, Socket socket) throws IOException { - OutputStream out = socket.getOutputStream(); - out.write(string.getBytes(Charset.forName("UTF-8"))); - out.flush(); - } + /** + * Writes and flushes the UTF-8 encoded contents of a String to a socket. + * + * @param string string to write + * @param socket socket to write to + */ + public static void writeStringToSocket(String string, Socket socket) throws IOException { + OutputStream out = socket.getOutputStream(); + out.write(string.getBytes(UTF_8)); + out.flush(); + } - /** - * Reads all available data from the socket and returns a String containing that content, interpreted in the - * UTF-8 charset. - * - * @param socket socket to read UTF-8 bytes from - * @return String containing the contents of whatever was read from the socket - * @throws EOFException if the socket has been closed - */ - public static String readStringFromSocket(Socket socket) throws IOException { - InputStream in = socket.getInputStream(); - byte[] bytes = new byte[10000]; - int bytesRead = in.read(bytes); - if (bytesRead == -1) { - throw new EOFException("Unable to read from socket. The socket is closed."); - } - - return new String(bytes, 0, bytesRead, Charset.forName("UTF-8")); + /** + * Reads all available data from the socket and returns a String containing that content, + * interpreted in the UTF-8 charset. + * + * @param socket socket to read UTF-8 bytes from + * @return String containing the contents of whatever was read from the socket + * @throws EOFException if the socket has been closed + */ + public static String readStringFromSocket(Socket socket) throws IOException { + InputStream in = socket.getInputStream(); + byte[] bytes = new byte[10000]; + int bytesRead = in.read(bytes); + if (bytesRead == -1) { + throw new EOFException("Unable to read from socket. The socket is closed."); } - /** - * Determines if the socket can be written to. This method tests the writability of the socket by writing to the socket, - * so it should only be used immediately before closing the socket. - * - * @param socket socket to test - * @return true if the socket is open and can be written to, otherwise false - */ - public static boolean isSocketReadyToWrite(Socket socket) throws IOException { - OutputStream out = socket.getOutputStream(); - try { - for(int i = 0; i < 500; ++i) { - out.write(0); - out.flush(); - } - } catch (SocketException e) { - return false; - } + return new String(bytes, 0, bytesRead, UTF_8); + } - return true; + /** + * Determines if the socket can be written to. This method tests the writability of the socket by + * writing to the socket, so it should only be used immediately before closing the socket. + * + * @param socket socket to test + * @return true if the socket is open and can be written to, otherwise false + */ + public static boolean isSocketReadyToWrite(Socket socket) throws IOException { + OutputStream out = socket.getOutputStream(); + try { + for (int i = 0; i < 500; ++i) { + // CR (0x0D) is skipped by Netty's HttpRequestDecoder in both strict and lenient modes; + // null bytes (0x00) trigger InvalidLineSeparatorException in strict mode (Netty ≥4.2.15). + out.write('\r'); + out.flush(); + } + } catch (SocketException e) { + return false; } - /** - * Determines if the socket can be read from. This method tests the readability of the socket by attempting to read - * a byte from the socket. If successful, the byte will be lost, so this method should only be called immediately - * before closing the socket. - * - * @param socket socket to test - * @return true if the socket is open and can be read from, otherwise false - */ - public static boolean isSocketReadyToRead(Socket socket) throws IOException { - InputStream in = socket.getInputStream(); - try { - int readByte = in.read(); + return true; + } - // we just lost that byte but it doesn't really matter for testing purposes - return readByte != -1; - } catch (SocketException e) { - // the socket couldn't be read, perhaps because the connection was reset or some other error. it cannot be read. - return false; - } catch (SocketTimeoutException e) { - // the read timed out, which means the socket is still connected but there's no data on it - return true; - } - } + /** + * Determines if the socket can be read from. This method tests the readability of the socket by + * attempting to read a byte from the socket. If successful, the byte will be lost, so this method + * should only be called immediately before closing the socket. + * + * @param socket socket to test + * @return true if the socket is open and can be read from, otherwise false + */ + public static boolean isSocketReadyToRead(Socket socket) throws IOException { + InputStream in = socket.getInputStream(); + try { + int readByte = in.read(); - /** - * Opens a socket to the specified proxy server with a 3s timeout. The socket should be closed after it has been used. - * - * @param proxyServer proxy server to open the socket to - * @return the new socket - */ - public static Socket getSocketToProxyServer(HttpProxyServer proxyServer) throws IOException { - Socket socket = new Socket(); - socket.connect(new InetSocketAddress("localhost", proxyServer.getListenAddress().getPort()), 1000); - socket.setSoTimeout(3000); - return socket; + // we just lost that byte, but it doesn't really matter for testing purposes + return readByte != -1; + } catch (SocketException e) { + // the socket couldn't be read, perhaps because the connection was reset or some other error. + // it cannot be read. + return false; + } catch (SocketTimeoutException e) { + // the read timed out, which means the socket is still connected but there's no data on it + return true; } + } + + /** + * Opens a socket to the specified proxy server with a 3s timeout. The socket should be closed + * after it has been used. + * + * @param proxyServer proxy server to open the socket to + * @return the new socket + */ + public static Socket getSocketToProxyServer(HttpProxyServer proxyServer) throws IOException { + Socket socket = new Socket(); + socket.connect( + new InetSocketAddress("localhost", proxyServer.getListenAddress().getPort()), 1000); + socket.setSoTimeout(3000); + return socket; + } } diff --git a/src/test/java/org/littleshoot/proxy/test/SocketUtil.java b/src/test/java/org/littleshoot/proxy/test/SocketUtil.java new file mode 100644 index 00000000..e69de29b diff --git a/src/test/java/org/littleshoot/proxy/test/ThreadDumpExtension.java b/src/test/java/org/littleshoot/proxy/test/ThreadDumpExtension.java new file mode 100644 index 00000000..89d2f38f --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/test/ThreadDumpExtension.java @@ -0,0 +1,169 @@ +package org.littleshoot.proxy.test; + +import static java.nio.charset.StandardCharsets.UTF_8; +import static java.util.concurrent.Executors.newScheduledThreadPool; +import static java.util.concurrent.TimeUnit.MILLISECONDS; +import static org.slf4j.LoggerFactory.getLogger; + +import java.io.File; +import java.io.IOException; +import java.lang.management.ManagementFactory; +import java.lang.management.ThreadInfo; +import java.lang.management.ThreadMXBean; +import java.lang.reflect.Method; +import java.util.concurrent.ScheduledExecutorService; +import org.apache.commons.io.FileUtils; +import org.junit.jupiter.api.Timeout; +import org.junit.jupiter.api.extension.AfterAllCallback; +import org.junit.jupiter.api.extension.AfterEachCallback; +import org.junit.jupiter.api.extension.BeforeAllCallback; +import org.junit.jupiter.api.extension.BeforeEachCallback; +import org.junit.jupiter.api.extension.ExtensionContext; +import org.junit.jupiter.api.extension.ExtensionContext.Namespace; +import org.opentest4j.TestAbortedException; +import org.slf4j.Logger; + +public class ThreadDumpExtension + implements BeforeAllCallback, AfterAllCallback, BeforeEachCallback, AfterEachCallback { + private static final Namespace NAMESPACE = Namespace.create("ThreadDumpExtension"); + private static final int INITIAL_DELAY_MS = 8000; + private static final int DELAY_MS = 1000; + + /** + * Checks if thread dump is enabled for the given context. + * + *

Thread dump is enabled if: + * + *

    + *
  • The test class has {@link EnableThreadDump} annotation, OR + *
  • The test method has {@link EnableThreadDump} annotation + *
+ */ + private boolean isThreadDumpEnabled(ExtensionContext context) { + return context + .getTestClass() + .map(clazz -> clazz.isAnnotationPresent(EnableThreadDump.class)) + .orElse(false) + || context + .getTestMethod() + .map(method -> method.isAnnotationPresent(EnableThreadDump.class)) + .orElse(false); + } + + @Override + public void beforeAll(ExtensionContext context) { + if (!isThreadDumpEnabled(context)) { + return; + } + getLogger(context.getDisplayName()).info("Starting tests ({})", memory()); + } + + @Override + public void afterAll(ExtensionContext context) { + if (!isThreadDumpEnabled(context)) { + return; + } + getLogger(context.getDisplayName()) + .info("Finished tests - {} ({})", verdict(context), memory()); + } + + @Override + public void beforeEach(ExtensionContext context) { + if (!isThreadDumpEnabled(context)) { + return; + } + Logger log = logger(context); + log.info("starting {} ({})...", context.getDisplayName(), memory()); + ScheduledExecutorService executor = newScheduledThreadPool(1); + executor.scheduleWithFixedDelay( + () -> takeThreadDump(log), initialDelayMs(context), DELAY_MS, MILLISECONDS); + + context.getStore(NAMESPACE).put("executor", executor); + + String originalThreadName = Thread.currentThread().getName(); + String newThreadName = + String.format( + "%s-%s-%s", + originalThreadName, + context.getTestClass().map(Class::getName).orElse(""), + context.getDisplayName()); + Thread.currentThread().setName(newThreadName); + context.getStore(NAMESPACE).put("originalThreadName", originalThreadName); + } + + private int initialDelayMs(ExtensionContext context) { + return context + .getTestMethod() + .map(method -> method.getAnnotation(Timeout.class)) + .map(timeout -> (int) timeout.unit().toMillis(timeout.value()) - 3000) + .map(timeout -> Math.max(timeout, 1000)) + .orElse(INITIAL_DELAY_MS); + } + + @Override + public void afterEach(ExtensionContext context) { + if (!isThreadDumpEnabled(context)) { + return; + } + logger(context) + .info("finished {} - {} ({})", context.getDisplayName(), verdict(context), memory()); + String originalThreadName = context.getStore(NAMESPACE).get("originalThreadName", String.class); + if (originalThreadName != null) { + Thread.currentThread().setName(originalThreadName); + } + ScheduledExecutorService executor = + (ScheduledExecutorService) context.getStore(NAMESPACE).remove("executor"); + if (executor != null) { + executor.shutdown(); + } + } + + private static Logger logger(ExtensionContext context) { + return getLogger( + context.getRequiredTestClass().getSimpleName() + + '.' + + context.getTestMethod().map(Method::getName).orElse("?")); + } + + private String verdict(ExtensionContext context) { + return context.getExecutionException().isPresent() + ? (context.getExecutionException().get() instanceof TestAbortedException + ? "skipped" + : "NOK") + : "OK"; + } + + private String memory() { + long freeMemory = Runtime.getRuntime().freeMemory(); + long maxMemory = Runtime.getRuntime().maxMemory(); + long totalMemory = Runtime.getRuntime().totalMemory(); + long usedMemory = totalMemory - freeMemory; + return "memory used:" + + mb(usedMemory) + + ", free:" + + mb(freeMemory) + + ", total:" + + mb(totalMemory) + + ", max:" + + mb(maxMemory); + } + + private long mb(long bytes) { + return bytes / 1024 / 1024; + } + + private void takeThreadDump(Logger log) { + StringBuilder threadDump = new StringBuilder(); + ThreadMXBean threadMXBean = ManagementFactory.getThreadMXBean(); + for (ThreadInfo threadInfo : threadMXBean.dumpAllThreads(true, true, 50)) { + threadDump.append(threadInfo.toString()); + } + File dump = new File("target/thread-dump-" + System.currentTimeMillis() + ".txt"); + try { + FileUtils.writeStringToFile(dump, threadDump.toString(), UTF_8); + log.info("Saved thread dump to file {}", dump); + } catch (IOException e) { + log.error("Failed to save thread dump to file {}", dump, e); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/websockets/WebSocketClient.java b/src/test/java/org/littleshoot/proxy/websockets/WebSocketClient.java new file mode 100644 index 00000000..6ca27e8a --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/websockets/WebSocketClient.java @@ -0,0 +1,200 @@ +package org.littleshoot.proxy.websockets; + +import static java.util.Objects.requireNonNull; +import static java.util.concurrent.TimeUnit.MILLISECONDS; + +import io.netty.bootstrap.Bootstrap; +import io.netty.channel.*; +import io.netty.channel.nio.NioEventLoopGroup; +import io.netty.channel.socket.SocketChannel; +import io.netty.channel.socket.nio.NioSocketChannel; +import io.netty.example.http.websocketx.client.WebSocketClientHandler; +import io.netty.handler.codec.http.EmptyHttpHeaders; +import io.netty.handler.codec.http.HttpClientCodec; +import io.netty.handler.codec.http.HttpObjectAggregator; +import io.netty.handler.codec.http.websocketx.*; +import io.netty.handler.logging.LogLevel; +import io.netty.handler.logging.LoggingHandler; +import io.netty.handler.proxy.HttpProxyHandler; +import io.netty.handler.ssl.SslContext; +import io.netty.handler.ssl.SslContextBuilder; +import io.netty.handler.ssl.SslHandler; +import io.netty.handler.ssl.util.InsecureTrustManagerFactory; +import java.net.InetSocketAddress; +import java.net.URI; +import java.time.Duration; +import java.util.Optional; +import java.util.Set; +import java.util.concurrent.BlockingQueue; +import java.util.concurrent.LinkedBlockingQueue; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.locks.ReadWriteLock; +import java.util.concurrent.locks.ReentrantReadWriteLock; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** Simple WebSocket client for use in unit tests that sends and receives text frames. */ +public class WebSocketClient { + private static final Duration CONNECT_TIMEOUT = Duration.ofSeconds(5); + private static final int MAX_AGGREGATOR_CONTENT_LENGTH = 65536; + private static final int MAX_PAYLOAD_FRAME_LENGTH = 1280000; + private static final Logger logger = LoggerFactory.getLogger(WebSocketClient.class); + + // According to RFC-6455 (https://tools.ietf.org/html/rfc6455#section-3) + // only the ws and wss schemes should be used, but some applications incorrectly + // use http/https anyway, so we allow both for testing those edge cases + private static final Set SECURE_SCHEMES = Set.of("wss", "https"); + + private final ReadWriteLock lock = new ReentrantReadWriteLock(); + private final BlockingQueue receivedMessages = new LinkedBlockingQueue<>(); + private Channel channel; + private EventLoopGroup group; + + public void open( + final URI uri, final Duration connectTimeout, final Optional httpProxy) + throws TimeoutException { + logger.info( + "{} connecting to {} via proxy {}", + getClass().getSimpleName(), + uri, + httpProxy.map(Object::toString).orElse("(none)")); + final boolean isSecure = SECURE_SCHEMES.contains(uri.getScheme()); + boolean connectionComplete = false; + lock.writeLock().lock(); + try { + if (channel != null) { + throw new IllegalStateException("Client already open"); + } + + group = new NioEventLoopGroup(); + + final WebSocketFrameReader handler = + new WebSocketFrameReader( + WebSocketClientHandshakerFactory.newHandshaker( + uri, + WebSocketVersion.V13, + null, + false, + EmptyHttpHeaders.INSTANCE, + MAX_PAYLOAD_FRAME_LENGTH, + true, + false, + -1, + httpProxy.isPresent())); + + final Bootstrap bootstrap = + new Bootstrap() + .group(group) + .channel(NioSocketChannel.class) + .option(ChannelOption.CONNECT_TIMEOUT_MILLIS, (int) CONNECT_TIMEOUT.toMillis()) + .handler(new WebSocketClientChannelInitializer(handler, uri, httpProxy)); + + /* + * If using a proxy we connect directly to the proxy for non-secure schemes. For + * secure schemes we add a proxy handler that uses HTTP CONNECT, so the + * bootstrap just needs to connect as if it's communicating directly with the + * origin server. + */ + final ChannelFuture connectFuture; + if (httpProxy.isPresent() && !isSecure) { + connectFuture = + bootstrap.connect(httpProxy.get().getHostString(), httpProxy.get().getPort()); + } else { + connectFuture = bootstrap.connect(uri.getHost(), uri.getPort()); + } + + if (!connectFuture.awaitUninterruptibly(connectTimeout.toMillis())) { + throw new TimeoutException("Connection timed out after " + connectTimeout); + } + channel = connectFuture.channel(); + if (!handler.handshakeFuture().awaitUninterruptibly(connectTimeout.toMillis())) { + throw new TimeoutException("Handshake timed out after " + connectTimeout); + } + connectionComplete = true; + } finally { + if (!connectionComplete) { + close(); + } + lock.writeLock().unlock(); + } + } + + public ChannelFuture send(final String value) { + lock.readLock().lock(); + try { + return channel.writeAndFlush(new TextWebSocketFrame(value)); + } finally { + lock.readLock().unlock(); + } + } + + public String waitForResponse(final Duration timeout) throws InterruptedException { + return receivedMessages.poll(timeout.toMillis(), MILLISECONDS); + } + + public void close() { + lock.writeLock().lock(); + try { + if (channel == null) { + return; + } + channel.writeAndFlush(new CloseWebSocketFrame()); + channel.close(); + group.shutdownGracefully(); + channel = null; + group = null; + } finally { + lock.writeLock().unlock(); + } + } + + private static class WebSocketClientChannelInitializer extends ChannelInitializer { + private final WebSocketClientHandler handler; + private final URI uri; + private final Optional httpProxy; + + public WebSocketClientChannelInitializer( + final WebSocketClientHandler handler, + final URI uri, + final Optional httpProxy) { + this.handler = requireNonNull(handler); + this.uri = requireNonNull(uri); + this.httpProxy = requireNonNull(httpProxy); + } + + @Override + public void initChannel(final SocketChannel ch) throws Exception { + ChannelPipeline pipeline = ch.pipeline(); + if (SECURE_SCHEMES.contains(uri.getScheme())) { + httpProxy.ifPresent(proxyAddress -> pipeline.addFirst(new HttpProxyHandler(proxyAddress))); + final SslContext sslContext = + SslContextBuilder.forClient() + .trustManager(InsecureTrustManagerFactory.INSTANCE) + .build(); + final SslHandler sslHandler = + sslContext.newHandler(ch.alloc(), uri.getHost(), uri.getPort()); + sslHandler.setHandshakeTimeoutMillis(CONNECT_TIMEOUT.toMillis()); + pipeline.addLast("ssl-handler", sslHandler); + } + pipeline.addLast("logging", new LoggingHandler(LogLevel.DEBUG)); + pipeline.addLast("http-codec", new HttpClientCodec()); + pipeline.addLast("http-aggregator", new HttpObjectAggregator(MAX_AGGREGATOR_CONTENT_LENGTH)); + pipeline.addLast("ws-handler", handler); + } + } + + private class WebSocketFrameReader extends WebSocketClientHandler { + + public WebSocketFrameReader(final WebSocketClientHandshaker handshaker) { + super(handshaker); + } + + @Override + public void channelRead0(final ChannelHandlerContext ctx, final Object msg) throws Exception { + if (msg instanceof TextWebSocketFrame textFrame) { + receivedMessages.offer(textFrame.text()); + } + super.channelRead0(ctx, msg); + } + } +} diff --git a/src/test/java/org/littleshoot/proxy/websockets/WebSocketClientServerTest.java b/src/test/java/org/littleshoot/proxy/websockets/WebSocketClientServerTest.java new file mode 100644 index 00000000..c10d9712 --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/websockets/WebSocketClientServerTest.java @@ -0,0 +1,143 @@ +package org.littleshoot.proxy.websockets; + +import static org.assertj.core.api.Assertions.assertThat; + +import java.net.InetSocketAddress; +import java.net.URI; +import java.time.Duration; +import java.util.Optional; +import java.util.concurrent.TimeoutException; +import org.junit.jupiter.api.*; +import org.littleshoot.proxy.HttpProxyServer; +import org.littleshoot.proxy.HttpProxyServerBootstrap; +import org.littleshoot.proxy.extras.TestMitmManager; +import org.littleshoot.proxy.impl.DefaultHttpProxyServer; +import org.littleshoot.proxy.test.EnableThreadDump; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +@Tag("slow-test") +@EnableThreadDump +public final class WebSocketClientServerTest { + private static final Duration CONNECT_TIMEOUT = Duration.ofSeconds(5); + private static final Duration RESPONSE_TIMEOUT = Duration.ofSeconds(5); + private static final int MAX_CONNECTION_ATTEMPTS = 5; + private static final long TEST_TIMEOUT_SECONDS = 60L; + private static final Logger logger = LoggerFactory.getLogger(WebSocketClientServerTest.class); + private HttpProxyServer proxy; + private final WebSocketServer server = new WebSocketServer(); + private final WebSocketClient client = new WebSocketClient(); + + @AfterEach + void tearDown() throws Exception { + client.close(); + server.stop(); + if (proxy != null) { + proxy.stop(); + proxy = null; + } + } + + private void startProxy(final boolean withSsl) { + final HttpProxyServerBootstrap bootstrap = + DefaultHttpProxyServer.bootstrap().withTransparent(true).withPort(0); + if (withSsl) { + bootstrap.withManInTheMiddle(new TestMitmManager()); + } + proxy = bootstrap.start(); + } + + @Disabled("Only useful for debugging issues with the proxy tests") + @Test + @Timeout(TEST_TIMEOUT_SECONDS) + public void directInsecureConnection() throws Exception { + testIntegration(false); + } + + @Disabled("Only useful for debugging issues with the proxy tests") + @Test + @Timeout(TEST_TIMEOUT_SECONDS) + public void directSecureConnection() throws Exception { + testIntegration(true); + } + + @Test + @Timeout(TEST_TIMEOUT_SECONDS) + public void proxiedInsecureConnectionWsScheme() throws Exception { + testIntegration(false, true, "ws"); + } + + @Test + @Timeout(TEST_TIMEOUT_SECONDS) + public void proxiedInsecureConnectionHttpScheme() throws Exception { + testIntegration(false, true, "http"); + } + + @Test + @Timeout(TEST_TIMEOUT_SECONDS) + public void proxiedSecureConnectionWssScheme() throws Exception { + testIntegration(true, true, "wss"); + } + + @Test + @Timeout(TEST_TIMEOUT_SECONDS) + public void proxiedSecureConnectionHttpsScheme() throws Exception { + testIntegration(true, true, "https"); + } + + private void testIntegration(final boolean withSsl) throws Exception { + testIntegration(withSsl, false, withSsl ? "wss" : "ws"); + } + + private void testIntegration(final boolean withSsl, final boolean withProxy, final String scheme) + throws Exception { + final InetSocketAddress serverAddress = server.start(withSsl, CONNECT_TIMEOUT); + if (withProxy) { + startProxy(withSsl); + } + + final URI serverUri = + URI.create( + scheme + + "://" + + serverAddress.getHostString() + + ":" + + serverAddress.getPort() + + WebSocketServer.WEBSOCKET_PATH); + + openClient(serverUri, withProxy); + + final String request = "test 1 test 2 test 3 test 4"; + assertThat(client.send(request).awaitUninterruptibly(RESPONSE_TIMEOUT.toMillis())) + .as("Timed out waiting for message to be sent after %s s.", RESPONSE_TIMEOUT) + .isTrue(); + final String response = client.waitForResponse(RESPONSE_TIMEOUT); + assertThat(response).isEqualTo(request.toUpperCase()); + } + + private void openClient(final URI uri, final boolean withProxy) throws InterruptedException { + final Optional proxyAddress = + Optional.ofNullable(proxy) + .filter(httpProxy -> withProxy) + .map(HttpProxyServer::getListenAddress); + int connectionAttempt = 0; + boolean connected = false; + while (!connected && connectionAttempt++ < MAX_CONNECTION_ATTEMPTS) { + try { + client.open(uri, CONNECT_TIMEOUT, proxyAddress); + connected = true; + } catch (TimeoutException e) { + logger.warn( + "Connection attempt {} of {} : {}", + connectionAttempt, + MAX_CONNECTION_ATTEMPTS, + e.getMessage(), + e); + Thread.sleep(CONNECT_TIMEOUT.toMillis() / 2); + } + } + assertThat(connected) + .as("Connection timed out after " + MAX_CONNECTION_ATTEMPTS + " attempts") + .isTrue(); + } +} diff --git a/src/test/java/org/littleshoot/proxy/websockets/WebSocketServer.java b/src/test/java/org/littleshoot/proxy/websockets/WebSocketServer.java new file mode 100644 index 00000000..0bda2a2c --- /dev/null +++ b/src/test/java/org/littleshoot/proxy/websockets/WebSocketServer.java @@ -0,0 +1,120 @@ +package org.littleshoot.proxy.websockets; + +import static java.util.Objects.requireNonNull; + +import io.netty.bootstrap.ServerBootstrap; +import io.netty.channel.*; +import io.netty.channel.nio.NioEventLoopGroup; +import io.netty.channel.socket.SocketChannel; +import io.netty.channel.socket.nio.NioServerSocketChannel; +import io.netty.example.http.websocketx.server.WebSocketFrameHandler; +import io.netty.handler.codec.http.HttpObjectAggregator; +import io.netty.handler.codec.http.HttpServerCodec; +import io.netty.handler.codec.http.websocketx.WebSocketServerProtocolHandler; +import io.netty.handler.logging.LogLevel; +import io.netty.handler.logging.LoggingHandler; +import io.netty.handler.ssl.SslContext; +import io.netty.handler.ssl.SslContextBuilder; +import io.netty.handler.ssl.util.SelfSignedCertificate; +import java.net.InetSocketAddress; +import java.security.cert.CertificateException; +import java.time.Duration; +import java.util.Optional; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.locks.Lock; +import java.util.concurrent.locks.ReentrantLock; +import javax.net.ssl.SSLException; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * Simple WebSocket server for use in unit tests that receives text frames and echoes them back + * after converting to upper case. + */ +public class WebSocketServer { + static final String WEBSOCKET_PATH = "/websocket"; + private static final int MAX_AGGREGATOR_CONTENT_LENGTH = 65536; + private static final Logger logger = LoggerFactory.getLogger(WebSocketServer.class); + private final Lock lock = new ReentrantLock(); + private EventLoopGroup bossGroup; + private EventLoopGroup workerGroup; + private Channel channel; + + public InetSocketAddress start(final boolean ssl, final Duration bindTimeout) + throws CertificateException, SSLException, TimeoutException { + lock.lock(); + try { + if (bossGroup != null) { + throw new IllegalStateException("Server already started"); + } + + final Optional sslCtx; + if (ssl) { + final SelfSignedCertificate ssc = new SelfSignedCertificate(); + sslCtx = + Optional.of(SslContextBuilder.forServer(ssc.certificate(), ssc.privateKey()).build()); + } else { + sslCtx = Optional.empty(); + } + + bossGroup = new NioEventLoopGroup(1); + workerGroup = new NioEventLoopGroup(); + + final ServerBootstrap bootstrap = + new ServerBootstrap() + .group(bossGroup, workerGroup) + .channel(NioServerSocketChannel.class) + .handler(new LoggingHandler(LogLevel.DEBUG)) + .childHandler(new WebSocketServerInitializer(sslCtx)); + + final ChannelFuture bindFuture = bootstrap.bind("localhost", 0); + if (!bindFuture.awaitUninterruptibly(bindTimeout.toMillis())) { + throw new TimeoutException("Bind timed out after " + bindTimeout); + } + channel = bindFuture.channel(); + final InetSocketAddress serverAddress = (InetSocketAddress) channel.localAddress(); + logger.info("{} listening on {}", getClass().getSimpleName(), serverAddress); + return serverAddress; + } finally { + lock.unlock(); + } + } + + public void stop() throws InterruptedException { + lock.lock(); + try { + if (bossGroup == null) { + return; + } + channel.close().sync(); + bossGroup.shutdownGracefully(); + workerGroup.shutdownGracefully(); + channel = null; + bossGroup = null; + workerGroup = null; + } finally { + lock.unlock(); + } + } + + private static class WebSocketServerInitializer extends ChannelInitializer { + private final Optional sslCtx; + + public WebSocketServerInitializer(final Optional sslCtx) { + this.sslCtx = requireNonNull(sslCtx); + } + + @Override + public void initChannel(final SocketChannel channel) { + final ChannelPipeline pipeline = channel.pipeline(); + sslCtx + .map(ctx -> ctx.newHandler(channel.alloc())) + .ifPresent(handler -> pipeline.addLast("ssl", handler)); + pipeline.addLast("http-codec", new HttpServerCodec()); + pipeline.addLast("http-aggregator", new HttpObjectAggregator(MAX_AGGREGATOR_CONTENT_LENGTH)); + pipeline.addLast( + "ws-protocol", new WebSocketServerProtocolHandler(WEBSOCKET_PATH, null, true)); + pipeline.addLast("ws-frame", new WebSocketFrameHandler()); + } + } +} diff --git a/src/test/resources/META-INF/services/org.junit.jupiter.api.extension.Extension b/src/test/resources/META-INF/services/org.junit.jupiter.api.extension.Extension new file mode 100644 index 00000000..1e1ddd12 --- /dev/null +++ b/src/test/resources/META-INF/services/org.junit.jupiter.api.extension.Extension @@ -0,0 +1 @@ +org.littleshoot.proxy.test.ThreadDumpExtension \ No newline at end of file diff --git a/src/test/resources/junit-platform.properties b/src/test/resources/junit-platform.properties new file mode 100644 index 00000000..789650ae --- /dev/null +++ b/src/test/resources/junit-platform.properties @@ -0,0 +1,7 @@ +junit.jupiter.extensions.autodetection.enabled=true +junit.jupiter.execution.timeout.default=30s +junit.jupiter.execution.parallel.enabled=true +junit.jupiter.execution.parallel.mode.default=concurrent +junit.jupiter.execution.parallel.mode.classes.default=concurrent +junit.jupiter.execution.parallel.config.strategy=dynamic +junit.jupiter.execution.parallel.config.dynamic.factor=1 \ No newline at end of file diff --git a/src/test/resources/littleproxy.properties b/src/test/resources/littleproxy.properties new file mode 100644 index 00000000..fa94aa83 --- /dev/null +++ b/src/test/resources/littleproxy.properties @@ -0,0 +1,2 @@ +# Idle connections are disconnected after X seconds of inactivity +idle_connection_timeout=50 \ No newline at end of file diff --git a/src/test/resources/log4j.xml b/src/test/resources/log4j.xml deleted file mode 100644 index 384002f6..00000000 --- a/src/test/resources/log4j.xml +++ /dev/null @@ -1,33 +0,0 @@ - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - \ No newline at end of file diff --git a/src/test/resources/log4j2.xml b/src/test/resources/log4j2.xml new file mode 100644 index 00000000..3f5956e7 --- /dev/null +++ b/src/test/resources/log4j2.xml @@ -0,0 +1,29 @@ + + + + target/log.txt + %-6r %d{ISO8601} %-5p [%t] %c{2} (%F:%L).%M() - %m%n + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/tagRelease.bash b/tagRelease.bash deleted file mode 100755 index a9002a36..00000000 --- a/tagRelease.bash +++ /dev/null @@ -1,12 +0,0 @@ -#!/usr/bin/env bash - -ARGS=1 # One arg to script expected. - -if [ $# -ne "$ARGS" ] -then - echo "Must include the version number" - exit 1 -fi - -RELEASE_VERSION=$1 -svn copy "http://svn.littleshoot.org/svn/littleproxy/trunk" "http://svn.littleshoot.org/svn/littleproxy/tags/littleproxy-${RELEASE_VERSION}" -m "Tag for LittleProxy release ${RELEASE_VERSION}"