diff --git a/.spelling b/.spelling index 0367bbc12..407e4f054 100644 --- a/.spelling +++ b/.spelling @@ -258,3 +258,4 @@ subkey 380 aks AKS +worstCase \ No newline at end of file diff --git a/Makefile b/Makefile index 5a309195e..11b6b5457 100644 --- a/Makefile +++ b/Makefile @@ -334,9 +334,9 @@ check-license-headers: .PHONY: check-fmtprints check-fmtprints: SHELL:=/bin/sh -check-fmtprints: # fails if there are any fmt.Print* calls outside of the 3 approved files +check-fmtprints: # fails if there are any fmt.Print* calls outside of the approved files @cd pkg && \ - fmtprints=$$(git grep -n fmt.Print | grep -v 'appinfo/usage/usage.go' | grep -v '^daemon/'); \ + fmtprints=$$(git grep -n fmt.Print | grep -v 'appinfo/usage/usage.go' | grep -v '^daemon/' | grep -v '^lb/example_test.go:'); \ count=0; \ if [ -n "$$fmtprints" ]; then \ count="$$(echo "$$fmtprints" | wc -l | tr -d '[:space:]')" ; \ diff --git a/deploy/kube/configmap.yaml b/deploy/kube/configmap.yaml index 44b980755..823844d0c 100644 --- a/deploy/kube/configmap.yaml +++ b/deploy/kube/configmap.yaml @@ -1149,6 +1149,10 @@ data: # provider: alb # alb: # # mechanism defines the ALB pool member selection mechanism. + # # rr (round_robin), p2c (power_of_two_choices), lc (least_connections), lt (least_time) + # # and hrw (highest_random_weight) select one pool member per request; fr, fgr, nlm and + # # tsm fan each request out to the pool; ur routes by user. race (connect_race) and mirror + # # (udp_mirror) serve only tcp/tls and udp stream listeners. See /docs/alb.md. # # values are rr, fr, fgr, nlm, or tsm. see the docs for detailed descriptions of each # # rr - standard round robin # # fr - fanout and return the First Response regardless of status code @@ -1169,6 +1173,12 @@ data: # # - foo-01.example.com # # - name: foo-02.example.com # # weight: 3 # receives 3 of every 4 requests + # # + # # a backup entry stands by: it is used only while no other pool member is healthy. e.g.: + # # pool: + # # - foo-01.example.com + # # - name: foo-02.example.com + # # backup: true # # healthy_floor is the minimum health status for a Backend to be considered healthy in the pool # # 1 indicates only backends positively reporting as healthy are included @@ -1177,6 +1187,10 @@ data: # # default is 0 # healthy_floor: 0 + # # propagate_health makes this ALB report itself unavailable, to any ALB that has it as a pool + # # member, while it has no healthy member of its own. default is false: it keeps its share + # propagate_health: false + # # max_capture_bytes overrides the backend-level max_capture_bytes for this ALB's fanout members. # # Set this when the ALB's expected response shape differs from the backend default. When 0 (the # # default), the parent Backend's max_capture_bytes is used, falling back to 268435456 (256 MiB). @@ -1184,9 +1198,47 @@ data: # fgr: # First Good Response mechanism options, only applicable when mechanism is set to fgr # # status_codes is a list of status codes considered 'good' when using the fgr mechanism # # when this is not set, any response code < 400 is considered good. Use this setting to - # # provide an explicit list. + # # provide an explicit list. An entry is a single code or an inclusive range, e.g.: + # # status_codes: [ { start: 200, end: 299 }, 304 ] # status_codes: [ 200 ] # this would consider only 200 OK's good, and not 204, 302, etc. + # hrw: # Highest Random Weight mechanism options, only applicable when mechanism is set to hrw + # # key is what a client's affinity follows: client_ip (default), host, header:, + # # cookie: or query:. Requests sharing a key reach the same pool member. + # # A stream listener can read client_ip, sni (the TLS server name) when it is tls, and + # # proxy_tlv: (a PROXY protocol v2 TLV, e.g. proxy_tlv:0xEA) when it accepts the + # # PROXY protocol. A native protocol listener, such as mysql, can read client_ip and user. + # key: client_ip + # # ipv6_prefix is how many leading bits of an IPv6 client address form a client_ip key. + # # default is 64, which keeps a client that rotates its privacy address on one member + # ipv6_prefix: 64 + + # lt: # Least Time mechanism options, only applicable when mechanism is set to lt + # # decay is how quickly a member's latency average follows a change. default is 10s + # decay: 10s + # # status_codes lists the response codes that count as a good answer; any other records + # # a latency penalty, so a member that fails fast never looks fast. Entries are single + # # codes or inclusive ranges. default is every code except 502, 503 and 504 + # status_codes: [ { start: 100, end: 501 }, { start: 505, end: 599 } ] + # # signal is what is timed. The default depends on the listener: first_write on http (the + # # first byte sent to the client), connect on tcp and tls (or first_byte), first_reply on udp + # signal: first_write + + # stream: # only for an ALB that serves a tcp, tls or udp listener. See /docs/alb.md + # # connect_retries is how many other pool members a tcp or tls connection is offered when + # # it cannot connect to the one it was given. default is 0: the connection is refused + # connect_retries: 0 + # # passive_health ejects a member after repeated failures to connect to it, without + # # waiting for a health check. It is off unless this block is present + # passive_health: + # failures: 3 # consecutive failed connects that eject a member + # eject: 30s # how long it stays out + # max_ejected_percent: 50 # never more of the pool than this, and never its last member + # # race_width is how many pool members the race mechanism connects to at once, 2 to 8. + # # default is every member, up to 4. connect_retries and passive_health do not apply to + # # the race and mirror mechanisms + # # race_width: 2 + # # ALB Pool Autodiscovery - see /docs/alb-autodiscovery.md for more information # # # # An ALB's pool membership can be discovered and kept current at runtime from diff --git a/deploy/kube/crds/trickstercachepolicies.yaml b/deploy/kube/crds/trickstercachepolicies.yaml index 6a3f3f6da..5e64fc9fc 100644 --- a/deploy/kube/crds/trickstercachepolicies.yaml +++ b/deploy/kube/crds/trickstercachepolicies.yaml @@ -168,6 +168,24 @@ spec: enum: - probe - provider + loadBalancing: + description: >- + The mechanism that spreads traffic across a Service's endpoints in the + endpoint routing mode; rr unless set. The weights between a rule's + backendRefs are always apportioned by rr. + type: string + enum: + - rr + - p2c + - lc + - lt + - hrw + loadBalancingKey: + description: >- + What the hrw mechanism keeps on one endpoint: client_ip (the default), + host, header:, cookie: or query: for HTTP routes, and + client_ip or sni for TCP, TLS and UDP routes. + type: string resultHeader: description: >- Whether the X-Trickster-Result response header reaches the client; diff --git a/docs/alb.md b/docs/alb.md index 1dedb8cf7..0fc79487f 100644 --- a/docs/alb.md +++ b/docs/alb.md @@ -5,11 +5,17 @@ Trickster 2.x provides an Application Load Balancer that is easy to configure an | Mechanism | Config | Provides | Description | |-----|-----|-----|----| | Round Robin | rr | Scaling | a basic, stateless round robin between healthy pool members | +| Power of Two Choices | p2c | Scaling | draws two healthy pool members at random and routes to the one with fewer requests in flight | +| Least Connections | lc | Scaling | routes to the healthy pool member with the fewest requests in flight | +| Least Time | lt | Speed | routes to the healthy pool member with the lowest response latency, scaled by its requests in flight | +| Highest Random Weight | hrw | Affinity | consistently routes each client, tenant or other key to the same healthy pool member | | Time Series Merge | tsm | Federation | uses scatter/gather to collect and merge data from multiple replica tsdb sources | | First Response | fr | Speed | fans a request out to multiple backends, and returns the first response received | | First Good Response | fgr | Speed | fans a request out to multiple backends, and returns the first response received with a status code < 400 | | Newest Last‑Modified | nlm | Freshness | fans a request out to multiple backends, and returns the response with the newest Last-Modified header | | User Router | ur | Control | Inspects the credentials in the Request and routes it based on the Username | +| Connect Race | race | Speed | connects a `tcp` or `tls` stream connection to several pool members at once and relays over the first to connect | +| UDP Mirror | mirror | Replication | copies every datagram of a `udp` stream session to every healthy pool member, and answers from the first | ## Integration with Backends @@ -29,7 +35,7 @@ Each mechanism has its own use cases and pitfalls. Be sure to read about each on A basic **Round Robin** rotates through a pool of healthy backends used to service client requests. Each time a client request is made to Trickster, the round robiner will identify the next healthy backend in the rotation schedule and route the request to it. -The Trickster ALB is intended to support stateless workloads, and currently does not support Sticky Sessions or other advanced ALB capabilities. +The Trickster ALB is intended to support stateless workloads, and currently does not support Sticky Sessions. For affinity without session state, see [Highest Random Weight](#highest-random-weight). #### Weighted Round Robin @@ -42,15 +48,13 @@ pool: weight: 3 # receives 3 of every 4 requests ``` -Apportionment is exact: over any `totalWeight` consecutive requests against a stable healthy pool, each member is selected exactly `weight` times. Weights also carry through from autodiscovery sources that convey them (DNS SRV record weights, member-file `weight` fields); see [ALB Autodiscovery](./alb-autodiscovery.md). +Apportionment is exact: over any `totalWeight` consecutive requests against a stable healthy pool, each member is selected exactly `weight` times. A heavier member's turns are spread through the rotation rather than taken back to back: two members weighted 3 and 2 are served `A B A A B`, not `A A A B B`. Weights also carry through from autodiscovery sources that convey them (DNS SRV record weights, member-file `weight` fields); see [ALB Autodiscovery](./alb-autodiscovery.md). -The legacy workaround of repeating a member name multiple times in the pool list still functions, but explicit weights replace it and are preferred. - -Weights apply to mechanisms that select a single member per request (round robin). Fan-out mechanisms (fr, fgr, nlm, tsm) dispatch to every healthy member regardless of weight. +Weights apply to every mechanism that selects a single member per request, though only Round Robin makes them an exact guarantee; see [Weights and the Selection Mechanisms](#weights-and-the-selection-mechanisms). Fan-out mechanisms (fr, fgr, nlm, tsm) dispatch to every healthy member regardless of weight. #### More About Our Round Robin Mechanism -Trickster's Round Robin Mechanism works by maintaining an atomic uint64 counter that increments each time a request is received by the ALB. With uniform weights, the ALB performs a modulo operation on the request's counter value, with the denominator being the count of healthy backends in the pool; the resulting value, ranging from `0` to `len(healthy_pool) - 1`, indicates the assigned backend based on the counter and current pool size. With mixed weights, the modulo denominator becomes the pool's total weight, and each member owns a contiguous `weight`-sized span of that rotation. Selection remains lock- and allocation-free in both forms. +Trickster's Round Robin Mechanism works by maintaining an atomic uint64 counter that increments each time a request is received by the ALB. With uniform weights, the ALB performs a modulo operation on the request's counter value, with the denominator being the count of healthy backends in the pool; the resulting value, ranging from `0` to `len(healthy_pool) - 1`, indicates the assigned backend based on the counter and current pool size. With mixed weights, the modulo denominator becomes the pool's total weight, and the result indexes a rotation schedule, computed once each time the set of healthy members changes, in which every member's turns are evenly spaced. Selection remains lock- and allocation-free in both forms. The counter starts at a random value, so a fleet of Trickster replicas started together does not send its first requests to the same pool member. #### Example Round Robin Configuration @@ -92,6 +96,180 @@ Here is the visual representation of this configuration: +### Power of Two Choices + +The **Power of Two Choices** (p2c) mechanism draws two healthy pool members at random and routes the request to whichever has fewer requests in flight. It keeps load nearly as even as inspecting every member would, at a cost that does not grow with the size of the pool, and unlike Round Robin it reacts to a member that has become slow: requests pile up there, so it loses more of its draws. It is a good default for large pools and for requests whose cost varies widely. + +```yaml +backends: + api: + provider: alb + alb: + mechanism: p2c # or power_of_two_choices + pool: [ node01, node02, node03 ] +``` + +### Least Connections + +The **Least Connections** (lc) mechanism routes each request to the healthy pool member with the fewest requests in flight. It reads every member on every request, so it suits small pools; Power of Two Choices approximates it for large ones. While several members are tied, as all are when the pool is idle, they take turns. + +```yaml +backends: + api: + provider: alb + alb: + mechanism: lc # or least_connections + pool: [ node01, node02 ] +``` + +### Least Time + +The **Least Time** (lt) mechanism routes each request to the healthy pool member that has been answering fastest. Each member's score is its latency average multiplied by one more than its requests in flight, and the lowest score wins, so a fast member is preferred until it is busy enough that a slower one would answer sooner. + +Latency is the time from routing a request to the first byte of its response. The average rises at once when a member slows down and falls gradually, over `lt.decay`, as it recovers. A few details keep the ranking honest: + +* A member with no requests yet, such as one just discovered, is scored as its fastest peer. It shares that peer's requests until its own first response ranks it, so it is neither flooded as the apparent fastest nor left waiting for a turn that an idle pool would never give it. +* A failed request is recorded as a long latency rather than a short one, so a member that returns errors in a millisecond never looks fast. By default a response of `502`, `503` or `504` is a failure, as is one that was never completed; `lt.status_codes` sets the codes that count as a good answer instead, as bare codes, inclusive ranges, or both. A request the client abandoned counts neither way. +* A member's average fades while it is passed over, so one ranked last on an old measurement or a past failure is tried again rather than ignored for good. +* Averages survive a configuration reload and autodiscovery membership changes. + +When the pool is idle and its members are equally loaded, every request goes to the fastest member. That is the mechanism working as intended; choose `p2c` or `lc` to spread idle traffic instead. + +```yaml +backends: + video: + provider: alb + alb: + mechanism: lt # or least_time + pool: [ edge01, edge02 ] + lt: + decay: 10s # default + status_codes: [ { start: 200, end: 499 } ] # optional; default is every code but 502, 503 and 504 +``` + +### Highest Random Weight + +The **Highest Random Weight** (hrw) mechanism, also known as rendezvous hashing, routes every request that shares a key to the same healthy pool member. Use it to keep a client or tenant on one member's warm cache. When a member leaves the pool, only the keys it owned move, each to a different remaining member; when it returns, exactly those keys move back. The mapping depends only on the key and the members' names, so every Trickster replica agrees on it and a restart does not change it. + +`hrw.key` selects what is hashed: + +| Key | Follows | +|-----|-----| +| `client_ip` (default) | the client's IP address, after [trusted proxy](./configuring.md) resolution. The port is never part of the key. IPv6 addresses are keyed on their leading `hrw.ipv6_prefix` bits (default `64`), because privacy addressing changes the rest of a client's address over time. | +| `host` | the request's host name, without its port and without regard to case | +| `header:` | the first value of the named request header | +| `cookie:` | the value of the named cookie | +| `query:` | the value of the named query string parameter, as written in the URL | +| `sni` | the TLS server name the client offered; only for an ALB that serves a `tls` [stream listener](#load-balancing-stream-listeners) | +| `proxy_tlv:` | the value of a [PROXY protocol](./configuring.md) version 2 TLV, such as `proxy_tlv:0xEA` for an AWS VPC endpoint ID; only for an ALB that serves a `tcp` or `tls` stream listener with `proxy_protocol` enabled. The type is one byte, written in decimal or `0x` hex. | +| `user` | the user name a session authenticated as; only for an ALB that serves a [native protocol listener](#load-balancing-native-protocol-sessions) | + +A request that lacks the configured key, such as one without the header, has no affinity to preserve and is routed to a member at random. + +hrw balances keys, not requests. With many keys of similar volume the members' loads even out, but a single very busy key is always served by one member. + +```yaml +backends: + tenants: + provider: alb + alb: + mechanism: hrw # or highest_random_weight + pool: [ cache01, cache02, cache03 ] + hrw: + key: header:X-Tenant +``` + +### Load Balancing Stream Listeners + +The mechanisms that select one member, `rr`, `p2c`, `lc`, `lt` and `hrw`, also balance the connections of a `tcp` or `tls` [stream listener](./configuring.md) and the sessions of a `udp` one. They are the same mechanisms with the same weights; only what they measure differs. The mechanisms that fan a request out (`fr`, `fgr`, `nlm`, `tsm`) and the User Router need an HTTP request, and are refused on a stream listener. Two more mechanisms, [`race` and `mirror`](#connect-race-and-udp-mirror), serve only stream listeners. + +| | http | tcp and tls | udp | +|-----|-----|-----|-----| +| unit of work | a request | a connection | a session: one client address and port | +| in flight (`p2c`, `lc`, `lt`) | requests being served | connections open | sessions open | +| latency (`lt.signal`) | `first_write`: the first byte sent to the client | `connect` (default): the time to connect to the member; or `first_byte`: the member's first byte | `first_reply`: the member's first datagram | +| `hrw.key` | `client_ip`, `host`, `header:`, `cookie:`, `query:` | `client_ip`; `sni` on a `tls` listener; `proxy_tlv:` with `proxy_protocol` | `client_ip` | + +A connection or session stays on the member it was given until it ends, whatever the mechanism. `client_ip` is the address a [PROXY protocol](./configuring.md) header names when the listener trusts one. + +Two settings apply only to an ALB that selects one member on a stream listener, under `alb.stream`: + +```yaml +backends: + pg: + provider: alb + listener_names: [ postgres ] + alb: + mechanism: p2c + pool: [ pg1, pg2, pg3 ] + stream: + connect_retries: 1 # default 0 + passive_health: # off unless present + failures: 3 # default + eject: 30s # default + max_ejected_percent: 50 # default +``` + +* `connect_retries` is how many other members a `tcp` or `tls` connection is offered when it cannot connect to the one it was given. All attempts share the listener's `stream.connect_timeout`. The default, 0, refuses the connection, which is what keeps a weighted split exact: a member under the reserved `.invalid` domain exists to refuse its share, and is never retried past. `udp` has no connect to fail, so it is not retried. +* `passive_health` takes a member out of the pool after `failures` consecutive failed connects, for `eject`, without waiting for a health check. Only failures to reach the member count, never anything it sent. At most `max_ejected_percent` of the pool is out at once, and the last live member is never ejected. When `eject` ends the member returns, unless its health check has it down. On `udp`, where nothing connects, a member that answers a datagram with a port-unreachable is what counts as a failure. + +Health checks work as they do for HTTP pools: a `tcp://` member with a `healthcheck.interval` is probed by opening a connection to it. A `udp://` member has no generic probe; rely on [autodiscovery](./alb-autodiscovery.md) readiness or on `passive_health`. + +#### Connect Race and UDP Mirror + +Two mechanisms exist only for stream listeners, because they commit one flow to several members at once. Both use every healthy member that has an address to dial, ignore `weight`, and take neither `connect_retries` nor `passive_health`. An ALB that uses one must be mapped to the stream listener directly; it cannot be a member of another ALB's pool. + +The **Connect Race** (`race`, or `connect_race`) mechanism serves `tcp` and `tls` listeners. Each client connection is connected to several members at once, within the listener's `stream.connect_timeout`; the first member to connect carries the connection, and the other attempts are closed. A member that is down or slow to accept costs the client nothing. `stream.race_width` is how many members are raced, from 2 to 8; the default is every member, up to 4. In a pool wider than the race, each connection starts one member further along, so the connects are shared across the pool. A race opens and discards connections on the members that lose, so use it where a connect is cheap for the member. + +The **UDP Mirror** (`mirror`, or `udp_mirror`) mechanism serves `udp` listeners. Every datagram a client sends is copied to every healthy member, which suits one-way protocols such as statsd, syslog and NetFlow. The first healthy member in pool order answers the session: only its replies are relayed to the client, and the other members' replies are discarded. A mirror pool has at most 8 members. + +```yaml +backends: + statsd: + provider: alb + listener_names: [ statsd ] + alb: + mechanism: mirror + pool: [ statsd1, statsd2 ] + db: + provider: alb + listener_names: [ postgres ] + alb: + mechanism: race + pool: [ pg1, pg2, pg3 ] + stream: + race_width: 2 +``` + +### Load Balancing Native Protocol Sessions + +A listener that speaks a backend's own wire protocol, such as a `mysql` [listener](./mysql.md), can map to an ALB that uses `rr`, `p2c`, `lc` or `hrw`. Each client session is committed to one pool member once it authenticates, and stays there until it ends; a session is the unit of work, so `p2c` and `lc` compare members by their open sessions. `hrw.key` is `client_ip` or `user`. `lt` is not available, since a session reports no latency to rank members by. Every pool member must be a backend of the listener's own provider, listed directly, and [autodiscovery](./alb-autodiscovery.md) is not supported. The ALB authenticates the listener's clients, so it carries the `authenticator_name` that a [User Router](#user-router) on the same listener would. + +```yaml +backends: + replicas: + provider: alb + listener_names: [ mysql ] + authenticator_name: mysql-clients + alb: + mechanism: lc + pool: [ replica1, replica2 ] +``` + +### Weights and the Selection Mechanisms + +Every mechanism that selects one member per request honors the pool `weight`, but what a weight promises differs: + +| Mechanism | A weight is | Guarantee | +|-----|-----|-----| +| rr | a share of requests | exact: `weight` of every `totalWeight` consecutive requests | +| p2c | a capacity | proportional on average: heavier members are drawn more often and compared by requests in flight per unit of weight | +| hrw | a share of the keys | proportional on average over many keys; changing a weight moves as few keys as possible | +| lc | a capacity | members are kept at equal requests in flight per unit of weight | +| lt | a bias | the score is divided by the weight, which shifts load under contention; an idle pool still sends every request to its best-scoring member | + +If you need a guaranteed split, use `rr`. + ### Time Series Merge The **Time Series Merge** mechanism supports both High Availability and federation. Each physical backend represents one logical data shard. Set the backend-level `replica_group` option to the same value on physical backends that are HA replicas of that shard. TSM first coalesces those replicas, using configured pool order to resolve overlapping points and later replicas to fill gaps, and then reduces the distinct logical shards. @@ -280,7 +458,7 @@ This mechanism is useful in applications such as live internet television. Consi #### Custom Good Status Codes List -By default, fgr will return the first response with a status code < 400. However, you can optionally provide an explicit list of good status codes using the `fgr.status_codes` configuration setting, as shown in the example below. When set, Trickster will return the first response to be returned that has a status code found in the configured list. +By default, fgr will return the first response with a status code < 400. However, you can optionally provide an explicit list of good status codes using the `fgr.status_codes` configuration setting, as shown in the example below. When set, Trickster will return the first response to be returned that has a status code found in the configured list. An entry may be a single code or an inclusive range, and the two forms mix: `status_codes: [ { start: 200, end: 299 }, 304 ]`. #### First Good Response Configuration Example @@ -525,6 +703,48 @@ Backends that do not have a [health check interval](./health#example+health+chec Setting `healthy_floor` below `0` admits members the probe has confirmed `unavailable`, not just members in the transient `unknown` state. If your goal is to keep traffic flowing during the cold-start window before the first probes complete, lower the pool members' `recovery_threshold` so they transition out of `unknown` faster -- don't lower the floor. When `healthy_floor < 0` Trickster emits a startup warning and sets the `trickster_alb_pool_admits_failing{backend_name}` gauge to `1`. +### Backup Pool Members + +A pool member marked `backup: true` stands by: it receives traffic only while no other member of the pool is in the healthy pool. As soon as one of the other members returns, the backup members stand down again. This applies to every mechanism that has a pool, and to stream and native protocol listeners as well as HTTP. A pool must have at least one member that is not a backup, unless its other members come from [autodiscovery](./alb-autodiscovery.md); discovered members are never backups. + +```yaml +backends: + db: + provider: alb + listener_names: [ postgres ] + alb: + mechanism: rr + healthy_floor: 1 + pool: + - primary + - name: standby + backup: true +``` + +Failover depends on the ALB learning that its other members are down, so give them a [health check interval](./health#example+health+check+configuration+for+use+in+alb) or, on a stream listener, `stream.passive_health`. While an ALB with backup members is dispatching to them, the `trickster_alb_pool_on_backup{backend_name}` gauge is `1`, and a warning is logged when it fails over. + +### ALBs as Pool Members + +An ALB that is a member of another ALB's pool is always treated as `available`, even when its own healthy pool is empty: it keeps its share of the outer pool's traffic and fails it. That is deliberate for weighted splits, where an empty member must not shift its share onto its siblings. Set `propagate_health: true` on the inner ALB to change that. It then reports `unavailable` to the pools it belongs to while it has no healthy member, and `available` otherwise, so an outer pool whose `healthy_floor` excludes `unavailable` members sends that share to its other members instead. + +```yaml +backends: + region-east: + provider: alb + alb: + mechanism: rr + propagate_health: true + pool: [ east1, east2 ] + global: + provider: alb + alb: + mechanism: rr + pool: + - region-east + - name: region-west + backup: true +``` + ### Example ALB Configuration Routing Only To Known Healthy Backends ```yaml diff --git a/docs/configuring.md b/docs/configuring.md index 76f3fefae..718f05348 100644 --- a/docs/configuring.md +++ b/docs/configuring.md @@ -181,12 +181,17 @@ it; every stream listener uses `port` and `address` alone. A stream listener's backend is a `reverseproxy` (`rp`) backend, whose `origin_url` supplies the host and port to dial and nothing more (the scheme -may be `tcp://` or `udp://`), or an `alb` backend using the `rr` mechanism, -whose pool members are such backends. Each connection or session is -committed to one pool member chosen by weighted round robin, as an HTTP -request is; a member that cannot be dialed refuses its share rather than -passing it to a sibling, so it is health checks or discovery readiness that -take a dead member out of rotation. A member whose origin host is under the +may be `tcp://` or `udp://`), or an `alb` backend whose pool members are such +backends. The ALB may use any mechanism that commits to one member: `rr`, +`p2c`, `lc`, `lt` or `hrw` (see [Load Balancing Stream +Listeners](./alb.md#load-balancing-stream-listeners)). Each connection or +session is committed to one pool member for its whole life, chosen as an HTTP +request's is. By default a member that cannot be dialed refuses its share +rather than passing it to a sibling, so it is health checks, discovery +readiness or passive health that take a dead member out of rotation; set +`alb.stream.connect_retries` to try another member instead. A `tcp://` member +with a `healthcheck.interval` is probed by opening a connection to it; a +`udp://` member has no generic probe. A member whose origin host is under the reserved `.invalid` domain, which can never resolve, refuses its share without a lookup, which is how a share that must be refused is expressed. A discovery-backed ALB works too, and a `scheme` of `tcp` or `udp` on its diff --git a/docs/developer/environment/trickster-config/trickster.yaml b/docs/developer/environment/trickster-config/trickster.yaml index 7dac00d0a..4d6c58e3e 100644 --- a/docs/developer/environment/trickster-config/trickster.yaml +++ b/docs/developer/environment/trickster-config/trickster.yaml @@ -116,9 +116,8 @@ backends: alb: mechanism: rr pool: - - prom1 - - prom1 - - prom1 + - name: prom1 + weight: 3 - sim1 mysql1: provider: mysql diff --git a/docs/kubernetes-cache-policy.md b/docs/kubernetes-cache-policy.md index 18c045f0f..e6fcbdf9b 100644 --- a/docs/kubernetes-cache-policy.md +++ b/docs/kubernetes-cache-policy.md @@ -100,6 +100,8 @@ spec: | `cors.mode` | `preserve`, `merge`, `replace`, `disable` | how origin CORS headers combine with the configured ones | | `cors.headers` | a map of headers | the CORS headers `merge` and `replace` apply | | `healthMode` | `probe`, `provider` | how discovered members are judged healthy in the endpoint routing mode | +| `loadBalancing` | `rr`, `p2c`, `lc`, `lt`, `hrw` | how traffic is spread across a Service's endpoints in the endpoint routing mode; `rr` unless set. See [the ALB mechanisms](./alb.md) | +| `loadBalancingKey` | `client_ip`, `host`, `header:`, `cookie:`, `query:` | what `hrw` keeps on one endpoint; `client_ip` unless set | | `resultHeader` | `Expose`, `Hide` | whether `X-Trickster-Result` reaches the client; see below | In a header map a name prefixed with `-` deletes the header and one prefixed with `+` diff --git a/docs/kubernetes-gateway.md b/docs/kubernetes-gateway.md index f9b0493c7..387562504 100644 --- a/docs/kubernetes-gateway.md +++ b/docs/kubernetes-gateway.md @@ -63,6 +63,8 @@ data: | `tracing_name`, `req_rewriter_name`, `authenticator_name` | a configured tracer, request rewriter or authenticator | | `timeout` | a duration with a unit, such as `30s` | | `health_mode` | `probe` or `provider`, for generated discovery-backed ALBs | +| `load_balancing` | `rr`, `p2c`, `lc`, `lt` or `hrw`: how traffic is spread across a Service's endpoints in the endpoint routing mode | +| `load_balancing_key` | what `hrw` keeps on one endpoint: `client_ip`, or `sni` on a TLSRoute; an HTTP route may also use `host`, `header:`, `cookie:` or `query:` | Unlike an Ingress annotation, a GatewayClass may set the operator-tier names, because a GatewayClass is cluster-scoped infrastructure and whoever can write @@ -465,7 +467,12 @@ unresolved reference, keeping its weight and refusing its share. Every connection, and every UDP client's session, is committed to one backendRef by weighted round robin, and in the endpoint routing mode to one -of that Service's ready endpoints in turn; a member that cannot be dialed +of that Service's ready endpoints, in turn unless the GatewayClass's +`load_balancing` parameter names another mechanism. The split between +backendRefs is always round robin, because Gateway API makes those weights an +exact apportionment; the mechanism applies within each Service. A key that a +route's listener cannot read, such as `sni` on a TCPRoute, is ignored in +favor of the client address. A member that cannot be dialed refuses the connection rather than passing it to a sibling, an unresolved reference refuses its share without a lookup, and it is the Service's readiness that takes an endpoint out of rotation. The @@ -473,7 +480,7 @@ readiness that takes an endpoint out of rotation. The judged by readiness, since no probe speaks the protocol they carry. A stream route caches nothing, and a `TricksterCachePolicy` cannot target one; `kubernetes.defaults` and a GatewayClass's parameters reach it only for -`routing_mode`. +`routing_mode`, `load_balancing` and `load_balancing_key`. The controller watches the three kinds only in a cluster whose experimental Gateway API channel serves them (`gateway.networking.k8s.io/v1alpha2`); see diff --git a/docs/kubernetes-ingress.md b/docs/kubernetes-ingress.md index 2cf956760..365ac7c92 100644 --- a/docs/kubernetes-ingress.md +++ b/docs/kubernetes-ingress.md @@ -159,6 +159,8 @@ leave an operator believing a setting is in force when it is not. | `trickstercache.org/use-regex` | `true`, `false` | compiles this object's `ImplementationSpecific` paths as anchored regular expressions | | `trickstercache.org/rewrite-target` | a path | rewrites the matched path on the way upstream | | `trickstercache.org/health-mode` | `probe`, `provider` | how discovered members are judged healthy in the endpoint routing mode | +| `trickstercache.org/load-balancing` | `rr`, `p2c`, `lc`, `lt`, `hrw` | how traffic is spread across a Service's endpoints in the endpoint routing mode; `rr` unless set | +| `trickstercache.org/load-balancing-key` | `client_ip`, `host`, `header:`, `cookie:`, `query:` | what `hrw` keeps on one endpoint | Durations require a unit: `600` is rejected, `600s` is not. @@ -305,6 +307,7 @@ independently. These are the equivalents: | `response-headers` | `responseHeaders`, a map, or a `ResponseHeaderModifier` filter | | `cors-mode`, `cors-headers` | `cors.mode`, `cors.headers` | | `health-mode` | `healthMode` | +| `load-balancing`, `load-balancing-key` | `loadBalancing`, `loadBalancingKey` | | `use-regex` | an HTTPRoute path match of type `RegularExpression` | | `rewrite-target` | a `URLRewrite` filter, whose `ReplacePrefixMatch` replaces the matched prefix and `ReplaceFullPath` the whole path | diff --git a/docs/metrics.md b/docs/metrics.md index cab1635c5..20cbc5d14 100644 --- a/docs/metrics.md +++ b/docs/metrics.md @@ -115,6 +115,25 @@ The following metrics are available for polling with any Trickster configuration * `protocol` - `tcp`, `tls` or `udp` * `direction` - `in` from the client to the backend, `out` from the backend to the client +* `trickster_proxy_stream_member_connections_total` (Counter) - The number of connections and UDP sessions a stream listener committed to an ALB pool member. + * labels: + * `listener_name` - the name of the configured listener + * `protocol` - `tcp`, `tls` or `udp` + * `backend_name` - the name of the pool member backend + * `result` - `proxied`, `dial_failed` (the member could not be connected to) or `unreachable` (a `udp` member answered a datagram with a port-unreachable) + +* `trickster_proxy_stream_member_active_connections` (Gauge) - The number of connections and UDP sessions open to an ALB pool member. + * labels: + * `listener_name` - the name of the configured listener + * `protocol` - `tcp`, `tls` or `udp` + * `backend_name` - the name of the pool member backend + +* `trickster_proxy_stream_member_connect_duration_seconds` (Histogram) - The time taken to connect to an ALB pool member. + * labels: + * `listener_name` - the name of the configured listener + * `protocol` - `tcp`, `tls` or `udp` + * `backend_name` - the name of the pool member backend + * `trickster_proxy_query_range_rejected_total` (Counter) - Trickster total number of queries rejected due to exceeding the `max_query_range` limit. * labels: * `backend` - the name of the configured backend rejecting the query @@ -193,9 +212,20 @@ The following metrics are available for polling with any Trickster configuration * `backend_name` - the name of the configured ALB backend * `trickster_alb_pool_floor_reset` (Gauge) - 1 when an ALB pool's `healthy_floor` was reset to 0 at startup because pool members have no health check and could never reach the configured floor, 0 otherwise. See [alb.md](./alb.md#health-based-backend-selection). +* `trickster_alb_pool_on_backup` (Gauge) - 1 while an ALB pool that has `backup` members is dispatching to them because no other member is available, 0 otherwise. Present only for pools with backup members. See [alb.md](./alb.md#backup-pool-members). * labels: * `backend_name` - the name of the configured ALB backend +* `trickster_alb_member_inflight` (Gauge) - Current number of requests in flight to an ALB pool member. Exported for the mechanisms that track it (`p2c`, `lc`, `lt`), for requests and for stream connections and sessions alike; read when the metrics endpoint is scraped, at no cost to request routing. + * labels: + * `alb_name` - the name of the configured ALB backend + * `member` - the name of the pool member backend + +* `trickster_alb_member_ejections_total` (Counter) - The number of times `alb.stream.passive_health` took a pool member out of selection after repeated connect failures. + * labels: + * `alb_name` - the name of the configured ALB backend + * `member` - the name of the pool member backend + The following metrics are available when [ALB Autodiscovery](./alb-autodiscovery.md) is configured: * `trickster_alb_discovery_members` (Gauge) - Current number of discovered ALB pool members diff --git a/docs/mysql.md b/docs/mysql.md index d20edef01..f07343d93 100644 --- a/docs/mysql.md +++ b/docs/mysql.md @@ -74,8 +74,9 @@ backends: recovery_threshold: 2 ``` -Exactly one direct MySQL backend or one supported MySQL User Router ALB maps to -a MySQL listener. The origin URL must use the `mysql` scheme and include an +Exactly one direct MySQL backend, one supported MySQL User Router ALB, or one +[session-balancing ALB](#balancing-sessions-across-replicas) maps to a MySQL +listener. The origin URL must use the `mysql` scheme and include an origin username. Percent-encode reserved username, password, and database characters. Configuration stringification and the sanitized management configuration redact an embedded origin password, but the source configuration @@ -376,6 +377,53 @@ The verified username and selected terminal remain in cache identity. Route metrics use configured router/backend names and bounded outcomes, never the username. +## Balancing sessions across replicas + +A MySQL listener can also map to an ALB that balances its sessions across a +pool of direct MySQL backends, such as read replicas: + +```yaml +backends: + replica-1: + provider: mysql + authenticator_name: app-clients + origin_url: mysql://app_ro:REDACTED@replica-1.example:3306/analytics + healthcheck: + interval: 5s + + replica-2: + provider: mysql + authenticator_name: app-clients + origin_url: mysql://app_ro:REDACTED@replica-2.example:3306/analytics + healthcheck: + interval: 5s + + replicas: + provider: alb + listener_names: [mysql-replicas] + authenticator_name: app-clients + alb: + mechanism: lc # rr, p2c, lc or hrw + pool: + - replica-1 + - name: replica-2 + weight: 2 +``` + +The ALB owns the downstream authentication exchange, admission, and TLS, as a +User Router does, and each pool member owns its origin credentials, cache, +health, and query policy. A session is committed to one member after it +authenticates and stays there until it ends. `lc` and `p2c` compare members by +their open sessions; `hrw` keeps a client address (`hrw.key: client_ip`, the +default) or a user name (`hrw.key: user`) on one member. `weight`, `backup` +members, and `healthy_floor` apply as they do for any ALB. `lt`, the fanout +mechanisms, nested ALBs, mixed providers, and autodiscovery are configuration +errors. See +[Load Balancing Native Protocol Sessions](./alb.md#load-balancing-native-protocol-sessions). + +A change of mechanism or weight applies to new sessions on reload; a change to +the set of pool members restarts the listener. + ## Metrics, logs, and health Important metrics include: diff --git a/docs/roadmap.md b/docs/roadmap.md index 5d76dc869..bd920b955 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -32,7 +32,7 @@ The roadmap for Trickster in 2026 focuses on delivering Trickster versions 2.0, - [ ] Support Access Control Lists (ACLs) for IPv4 and IPv6 - [ ] Support Rate Limiting w/ bucketing on configurable request attributes - [ ] Support ALB Sticky Sessions - - [ ] Support L4 Load Balancing + - [x] Support L4 Load Balancing - [ ] Support Media over Quick (MoQ) Relaying - [ ] Improved support for accelerating distributed Mimir deployments diff --git a/examples/conf/example.full.yaml b/examples/conf/example.full.yaml index 0da62d866..efb5e126d 100644 --- a/examples/conf/example.full.yaml +++ b/examples/conf/example.full.yaml @@ -1138,6 +1138,10 @@ backends: # provider: alb # alb: # # mechanism defines the ALB pool member selection mechanism. +# # rr (round_robin), p2c (power_of_two_choices), lc (least_connections), lt (least_time) +# # and hrw (highest_random_weight) select one pool member per request; fr, fgr, nlm and +# # tsm fan each request out to the pool; ur routes by user. race (connect_race) and mirror +# # (udp_mirror) serve only tcp/tls and udp stream listeners. See /docs/alb.md. # # values are rr, fr, fgr, nlm, or tsm. see the docs for detailed descriptions of each # # rr - standard round robin # # fr - fanout and return the First Response regardless of status code @@ -1158,6 +1162,12 @@ backends: # # - foo-01.example.com # # - name: foo-02.example.com # # weight: 3 # receives 3 of every 4 requests +# # +# # a backup entry stands by: it is used only while no other pool member is healthy. e.g.: +# # pool: +# # - foo-01.example.com +# # - name: foo-02.example.com +# # backup: true # # healthy_floor is the minimum health status for a Backend to be considered healthy in the pool # # 1 indicates only backends positively reporting as healthy are included @@ -1166,6 +1176,10 @@ backends: # # default is 0 # healthy_floor: 0 +# # propagate_health makes this ALB report itself unavailable, to any ALB that has it as a pool +# # member, while it has no healthy member of its own. default is false: it keeps its share +# propagate_health: false + # # max_capture_bytes overrides the backend-level max_capture_bytes for this ALB's fanout members. # # Set this when the ALB's expected response shape differs from the backend default. When 0 (the # # default), the parent Backend's max_capture_bytes is used, falling back to 268435456 (256 MiB). @@ -1173,9 +1187,47 @@ backends: # fgr: # First Good Response mechanism options, only applicable when mechanism is set to fgr # # status_codes is a list of status codes considered 'good' when using the fgr mechanism # # when this is not set, any response code < 400 is considered good. Use this setting to -# # provide an explicit list. +# # provide an explicit list. An entry is a single code or an inclusive range, e.g.: +# # status_codes: [ { start: 200, end: 299 }, 304 ] # status_codes: [ 200 ] # this would consider only 200 OK's good, and not 204, 302, etc. +# hrw: # Highest Random Weight mechanism options, only applicable when mechanism is set to hrw +# # key is what a client's affinity follows: client_ip (default), host, header:, +# # cookie: or query:. Requests sharing a key reach the same pool member. +# # A stream listener can read client_ip, sni (the TLS server name) when it is tls, and +# # proxy_tlv: (a PROXY protocol v2 TLV, e.g. proxy_tlv:0xEA) when it accepts the +# # PROXY protocol. A native protocol listener, such as mysql, can read client_ip and user. +# key: client_ip +# # ipv6_prefix is how many leading bits of an IPv6 client address form a client_ip key. +# # default is 64, which keeps a client that rotates its privacy address on one member +# ipv6_prefix: 64 + +# lt: # Least Time mechanism options, only applicable when mechanism is set to lt +# # decay is how quickly a member's latency average follows a change. default is 10s +# decay: 10s +# # status_codes lists the response codes that count as a good answer; any other records +# # a latency penalty, so a member that fails fast never looks fast. Entries are single +# # codes or inclusive ranges. default is every code except 502, 503 and 504 +# status_codes: [ { start: 100, end: 501 }, { start: 505, end: 599 } ] +# # signal is what is timed. The default depends on the listener: first_write on http (the +# # first byte sent to the client), connect on tcp and tls (or first_byte), first_reply on udp +# signal: first_write + +# stream: # only for an ALB that serves a tcp, tls or udp listener. See /docs/alb.md +# # connect_retries is how many other pool members a tcp or tls connection is offered when +# # it cannot connect to the one it was given. default is 0: the connection is refused +# connect_retries: 0 +# # passive_health ejects a member after repeated failures to connect to it, without +# # waiting for a health check. It is off unless this block is present +# passive_health: +# failures: 3 # consecutive failed connects that eject a member +# eject: 30s # how long it stays out +# max_ejected_percent: 50 # never more of the pool than this, and never its last member +# # race_width is how many pool members the race mechanism connects to at once, 2 to 8. +# # default is every member, up to 4. connect_retries and passive_health do not apply to +# # the race and mirror mechanisms +# # race_width: 2 + # # ALB Pool Autodiscovery - see /docs/alb-autodiscovery.md for more information # # # # An ALB's pool membership can be discovered and kept current at runtime from diff --git a/integration/alb_strategies_test.go b/integration/alb_strategies_test.go new file mode 100644 index 000000000..68e8cad61 --- /dev/null +++ b/integration/alb_strategies_test.go @@ -0,0 +1,203 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package integration + +import ( + "context" + "fmt" + "io" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/integration/internal/portutil" + "github.com/trickstercache/trickster/v2/integration/promstub" + + "github.com/stretchr/testify/require" +) + +// strategyALB runs Trickster with one ALB of the given mechanism over a pool of stubs. +// extra is appended, already indented, under the alb block; slow names pool members that +// answer with simulated latency. +type strategyALB struct { + mech string + extra string + stubs []*flappingStub + slow map[int]string + front int +} + +func startStrategyALB(t *testing.T, mech, extra string, poolSize int, slow map[int]string) *strategyALB { + t.Helper() + ports, release := portutil.Reserve(t, 3) + a := &strategyALB{mech: mech, extra: extra, slow: slow, front: ports[0]} + a.stubs = make([]*flappingStub, poolSize) + for i := range a.stubs { + a.stubs[i] = newFlappingStub(t, fmt.Sprintf("p%d", i), true) + } + var sb strings.Builder + sb.WriteString(promstub.Preamble(ports[0], ports[1], ports[2])) + sb.WriteString("backends:\n") + for i, s := range a.stubs { + sb.WriteString(promstub.BackendStanza(fmt.Sprintf("prom%d", i), s.URL())) + if d, ok := slow[i]; ok { + fmt.Fprintf(&sb, " latency_min: %s\n latency_max: %s\n", d, d) + } + } + sb.WriteString(" alb-strategy:\n provider: alb\n alb:\n") + fmt.Fprintf(&sb, " mechanism: %s\n", mech) + sb.WriteString(" healthy_floor: 1\n") + sb.WriteString(extra) + sb.WriteString(" pool:\n") + for i := range a.stubs { + fmt.Fprintf(&sb, " - prom%d\n", i) + } + cfgPath := filepath.Join(t.TempDir(), "trickster.yaml") + require.NoError(t, os.WriteFile(cfgPath, []byte(sb.String()), 0o644)) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + release() + runTrickster(t, ctx, "-config", cfgPath) + waitForTrickster(t, fmt.Sprintf("127.0.0.1:%d", ports[1])) + healthURL := fmt.Sprintf("http://127.0.0.1:%d/trickster/health", ports[1]) + for i := range a.stubs { + requireHealthState(t, healthURL, fmt.Sprintf("prom%d", i), "available", 10*time.Second) + } + return a +} + +// status sends one instant query, distinct per n so no cache answers it, and returns the +// response code +func (a *strategyALB) status(t *testing.T, n int, header http.Header) int { + t.Helper() + q := url.Values{"query": {fmt.Sprintf(`up{n="%d"}`, n)}} + u := fmt.Sprintf("http://127.0.0.1:%d/alb-strategy/api/v1/query?%s", a.front, q.Encode()) + req, err := http.NewRequest(http.MethodGet, u, nil) + require.NoError(t, err) + for k, v := range header { + req.Header[k] = v + } + resp, err := (&http.Client{Timeout: 10 * time.Second}).Do(req) + require.NoError(t, err) + _, _ = io.Copy(io.Discard, resp.Body) + _ = resp.Body.Close() + return resp.StatusCode +} + +// query is status for a pool whose every member is up: anything but a 200 fails the test +func (a *strategyALB) query(t *testing.T, n int, header http.Header) { + t.Helper() + require.Equal(t, http.StatusOK, a.status(t, n, header), "%s request %d", a.mech, n) +} + +func (a *strategyALB) hits() []int64 { + out := make([]int64, len(a.stubs)) + for i, s := range a.stubs { + out[i] = s.hits.Load() + } + return out +} + +// every strategy serves every request from a healthy pool, reaches all of its members, and +// stops routing to a member once its health check fails +func TestALBStrategiesServeAndFollowHealth(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster per mechanism; skipping in -short mode") + } + for _, mech := range []string{"p2c", "lc", "lt", "hrw", "rr"} { + t.Run(mech, func(t *testing.T) { + // every test client is 127.0.0.1, so the mechanism that routes by key is given one + // that varies + extra := "" + if mech == "hrw" { + extra = " hrw:\n key: header:X-Client\n" + } + a := startStrategyALB(t, mech, extra, 3, nil) + for i := range 60 { + a.query(t, i, http.Header{"X-Client": {fmt.Sprintf("client-%d", i)}}) + } + if mech != "lt" { + // lt rightly favors whichever member answered fastest; the rest spread + for i, n := range a.hits() { + require.NotZero(t, n, "%s never routed to pool member %d of 3: %v", mech, i, a.hits()) + } + } + + // until its health check notices, the downed member still answers its share, with + // errors; after that it must get nothing and every request must succeed + a.stubs[0].setUp(false) + n := 1000 + require.Eventually(t, func() bool { + start := a.stubs[0].hits.Load() + clean := true + for range 12 { + n++ + clean = a.status(t, n, http.Header{"X-Client": {fmt.Sprintf("client-%d", n)}}) == http.StatusOK && clean + } + return clean && a.stubs[0].hits.Load() == start + }, 15*time.Second, 200*time.Millisecond, + "%s kept routing to a member whose health check fails", mech) + }) + } +} + +// with hrw keyed on a header, a tenant's requests all reach one member, and tenants spread +func TestALBHRWStickiness(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + a := startStrategyALB(t, "hrw", " hrw:\n key: header:X-Tenant\n", 4, nil) + owners := make(map[int]bool) + n := 0 + for tenant := range 12 { + before := a.hits() + for range 5 { + a.query(t, n, http.Header{"X-Tenant": {fmt.Sprintf("tenant-%d", tenant)}}) + n++ + } + var served []int + for i, h := range a.hits() { + if h != before[i] { + served = append(served, i) + require.EqualValues(t, 5, h-before[i], "tenant-%d was split across members", tenant) + } + } + require.Len(t, served, 1, "tenant-%d reached members %v", tenant, served) + owners[served[0]] = true + } + require.GreaterOrEqual(t, len(owners), 3, "12 tenants landed on only %d of 4 members", len(owners)) +} + +// lt learns which member answers faster and sends it the bulk of sequential traffic +func TestALBLeastTimePrefersTheFasterMember(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + a := startStrategyALB(t, "lt", "", 2, map[int]string{0: "150ms"}) + for i := range 40 { + a.query(t, i, nil) + } + hits := a.hits() + require.Greater(t, hits[1], 3*hits[0], + "the member without simulated latency took %d of 40, the slow one %d", hits[1], hits[0]) + require.NotZero(t, hits[0], "the slow member was never tried") +} diff --git a/integration/stream_lb_test.go b/integration/stream_lb_test.go new file mode 100644 index 000000000..c0c7513b6 --- /dev/null +++ b/integration/stream_lb_test.go @@ -0,0 +1,510 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package integration + +import ( + "bufio" + "context" + "fmt" + "net" + "os" + "os/signal" + "path/filepath" + "slices" + "strings" + "sync" + "sync/atomic" + "syscall" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/integration/internal/portutil" + "github.com/trickstercache/trickster/v2/integration/promstub" + + "github.com/stretchr/testify/require" +) + +// streamEcho is a tcp backend that answers each line with its own name and counts the +// connections it accepted; it can be taken down and brought back on the same port +type streamEcho struct { + name string + addr string + mu sync.Mutex + ln net.Listener + accepted atomic.Int64 +} + +func newStreamEcho(t *testing.T, name string) *streamEcho { + t.Helper() + e := &streamEcho{name: name} + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + e.addr = ln.Addr().String() + e.serve(ln) + t.Cleanup(e.stop) + return e +} + +func (e *streamEcho) serve(ln net.Listener) { + e.mu.Lock() + e.ln = ln + e.mu.Unlock() + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + e.accepted.Add(1) + go func() { + defer conn.Close() + r := bufio.NewReader(conn) + for { + line, err := r.ReadString('\n') + if err != nil { + return + } + _, _ = fmt.Fprintf(conn, "%s:%s\n", e.name, strings.TrimSpace(line)) + } + }() + } + }() +} + +func (e *streamEcho) stop() { + e.mu.Lock() + defer e.mu.Unlock() + if e.ln != nil { + _ = e.ln.Close() + e.ln = nil + } +} + +func (e *streamEcho) restart(t *testing.T) { + t.Helper() + require.Eventually(t, func() bool { + ln, err := net.Listen("tcp", e.addr) + if err != nil { + return false + } + e.serve(ln) + return true + }, 5*time.Second, 50*time.Millisecond, "could not listen on %s again", e.addr) +} + +func udpEchoBackend(t *testing.T, name string) string { + t.Helper() + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = pc.Close() }) + go func() { + buf := make([]byte, 1500) + for { + n, from, err := pc.ReadFrom(buf) + if err != nil { + return + } + _, _ = pc.WriteTo([]byte(name+":"+string(buf[:n])), from) + } + }() + return pc.LocalAddr().String() +} + +// streamLB is a running Trickster with one stream listener bound to one ALB +type streamLB struct { + cfgPath string + addr string + protocol string + ports []int +} + +// streamConfig writes the configuration: members are name -> address, and alb is the body of +// the ALB's alb block below its pool, already indented; the members named as backups stand by +func (s *streamLB) write(t *testing.T, members [][2]string, memberExtra, alb string, backups ...string) { + t.Helper() + var sb strings.Builder + // the stream listener belongs with the preamble's own, ahead of its other sections + relay := fmt.Sprintf("listeners:\n relay:\n address: 127.0.0.1\n protocol: %s\n port: %d\n"+ + " stream:\n connect_timeout: 2s\n", s.protocol, s.ports[3]) + sb.WriteString(strings.Replace(promstub.Preamble(s.ports[0], s.ports[1], s.ports[2]), "listeners:\n", relay, 1)) + sb.WriteString("backends:\n") + sb.WriteString(" none:\n provider: rp\n origin_url: http://127.0.0.1:1\n") + for _, m := range members { + fmt.Fprintf(&sb, " %s:\n provider: rp\n origin_url: %s://%s\n listener_names: [relay]\n", + m[0], map[bool]string{true: "udp", false: "tcp"}[s.protocol == "udp"], m[1]) + sb.WriteString(memberExtra) + } + sb.WriteString(" lb:\n provider: alb\n listener_names: [relay]\n alb:\n") + sb.WriteString(alb) + sb.WriteString(" pool:\n") + for _, m := range members { + if slices.Contains(backups, m[0]) { + fmt.Fprintf(&sb, " - {name: %s, backup: true}\n", m[0]) + continue + } + fmt.Fprintf(&sb, " - %s\n", m[0]) + } + require.NoError(t, os.WriteFile(s.cfgPath, []byte(sb.String()), 0o644)) +} + +func startStreamLB(t *testing.T, protocol string, members [][2]string, memberExtra, alb string, + backups ...string, +) *streamLB { + t.Helper() + ports, release := portutil.Reserve(t, 4) + s := &streamLB{ + cfgPath: filepath.Join(t.TempDir(), "trickster.yaml"), protocol: protocol, ports: ports, + addr: fmt.Sprintf("127.0.0.1:%d", ports[3]), + } + s.write(t, members, memberExtra, alb, backups...) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + release() + runTrickster(t, ctx, "-config", s.cfgPath) + waitForTrickster(t, fmt.Sprintf("127.0.0.1:%d", ports[1])) + return s +} + +// ask opens a connection, sends one line and returns which member answered, or "" when the +// connection was refused; the connection is returned open +func (s *streamLB) ask(t *testing.T) (string, net.Conn) { + t.Helper() + conn, err := net.DialTimeout("tcp", s.addr, 2*time.Second) + require.NoError(t, err) + _ = conn.SetDeadline(time.Now().Add(5 * time.Second)) + if _, err := conn.Write([]byte("hi\n")); err != nil { + _ = conn.Close() + return "", nil + } + reply, err := bufio.NewReader(conn).ReadString('\n') + if err != nil { + _ = conn.Close() + return "", nil + } + name, _, _ := strings.Cut(reply, ":") + return name, conn +} + +func (s *streamLB) askAndClose(t *testing.T) string { + t.Helper() + name, conn := s.ask(t) + if conn != nil { + _ = conn.Close() + } + return name +} + +func tcpMembers(t *testing.T, names ...string) ([][2]string, map[string]*streamEcho) { + t.Helper() + members := make([][2]string, len(names)) + echoes := make(map[string]*streamEcho, len(names)) + for i, n := range names { + echoes[n] = newStreamEcho(t, n) + members[i] = [2]string{n, echoes[n].addr} + } + return members, echoes +} + +// every mechanism that commits to one member balances tcp connections across the whole pool +func TestStreamLBMechanisms(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster per mechanism; skipping in -short mode") + } + for _, mech := range []string{"rr", "p2c", "lc", "lt"} { + t.Run(mech, func(t *testing.T) { + members, _ := tcpMembers(t, "a", "b", "c") + s := startStreamLB(t, "tcp", members, "", " mechanism: "+mech+"\n") + seen := make(map[string]int) + var held []net.Conn + for range 60 { + name, conn := s.ask(t) + require.NotEmpty(t, name, "%s refused a connection with every member up", mech) + seen[name]++ + // keep some work in flight, which is what the load-aware mechanisms balance + held = append(held, conn) + if len(held) > 6 { + _ = held[0].Close() + held = held[1:] + } + } + for _, c := range held { + _ = c.Close() + } + for _, n := range []string{"a", "b", "c"} { + require.NotZero(t, seen[n], "%s never connected to %s: %v", mech, n, seen) + } + if mech == "rr" { + require.Equal(t, map[string]int{"a": 20, "b": 20, "c": 20}, seen) + } + }) + } +} + +// with connections held open, least connections keeps the members level +func TestStreamLBLeastConnectionsHoldsLevel(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + members, _ := tcpMembers(t, "a", "b", "c") + s := startStreamLB(t, "tcp", members, "", " mechanism: lc\n") + held := make(map[string]int) + var conns []net.Conn + for range 30 { + name, conn := s.ask(t) + require.NotEmpty(t, name) + held[name]++ + conns = append(conns, conn) + } + defer func() { + for _, c := range conns { + _ = c.Close() + } + }() + require.Equal(t, map[string]int{"a": 10, "b": 10, "c": 10}, held) +} + +// hrw keeps a client on one member: every connection here comes from 127.0.0.1 +func TestStreamLBHRWKeepsAClientOnOneMember(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + members, _ := tcpMembers(t, "a", "b", "c", "d") + s := startStreamLB(t, "tcp", members, "", " mechanism: hrw\n") + first := s.askAndClose(t) + require.NotEmpty(t, first) + for range 20 { + require.Equal(t, first, s.askAndClose(t), "one client address reached two members") + } +} + +// a udp session is committed to one member for life, and sessions are spread across the pool +func TestStreamLBUDPSessions(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + members := [][2]string{{"a", udpEchoBackend(t, "a")}, {"b", udpEchoBackend(t, "b")}} + s := startStreamLB(t, "udp", members, "", " mechanism: p2c\n") + seen := make(map[string]int) + for i := range 12 { + conn, err := net.Dial("udp", s.addr) + require.NoError(t, err) + var owner string + for j := range 3 { + _, _ = conn.Write([]byte("ping")) + _ = conn.SetReadDeadline(time.Now().Add(3 * time.Second)) + buf := make([]byte, 32) + n, err := conn.Read(buf) + require.NoError(t, err, "session %d datagram %d", i, j) + name, _, _ := strings.Cut(string(buf[:n]), ":") + if j > 0 { + require.Equal(t, owner, name, "session %d moved between members", i) + } + owner = name + } + seen[owner]++ + _ = conn.Close() + } + require.NotZero(t, seen["a"], "%v", seen) + require.NotZero(t, seen["b"], "%v", seen) +} + +// a member that goes down is taken out by its connect probe, and returns when it is back +func TestStreamLBConnectProbe(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + members, echoes := tcpMembers(t, "a", "b") + probe := " healthcheck:\n interval: 100ms\n timeout: 500ms\n failure_threshold: 1\n recovery_threshold: 1\n" + s := startStreamLB(t, "tcp", members, probe, " mechanism: rr\n healthy_floor: 1\n") + healthURL := fmt.Sprintf("http://127.0.0.1:%d/trickster/health", s.ports[1]) + requireHealthState(t, healthURL, "a", "available", 10*time.Second) + requireHealthState(t, healthURL, "b", "available", 10*time.Second) + + echoes["a"].stop() + requireHealthState(t, healthURL, "a", "unavailable", 10*time.Second) + for range 10 { + require.Equal(t, "b", s.askAndClose(t), "a connection was offered to the member that is down") + } + echoes["a"].restart(t) + requireHealthState(t, healthURL, "a", "available", 10*time.Second) + require.Eventually(t, func() bool { return s.askAndClose(t) == "a" }, 5*time.Second, 50*time.Millisecond, + "the recovered member never took a connection") +} + +// with no probe, a dead member refuses its share until passive health ejects it; with +// retries, its share is served by the others from the start +func TestStreamLBPassiveHealthAndRetries(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + t.Run("default refuses the dead member's share", func(t *testing.T) { + members, echoes := tcpMembers(t, "a", "b") + s := startStreamLB(t, "tcp", members, "", " mechanism: rr\n") + echoes["a"].stop() + var refused int + for range 10 { + if s.askAndClose(t) == "" { + refused++ + } + } + require.Equal(t, 5, refused, "a dead member's share is refused, not shed") + }) + t.Run("passive health ejects it", func(t *testing.T) { + members, echoes := tcpMembers(t, "a", "b", "c") + s := startStreamLB(t, "tcp", members, "", + " mechanism: rr\n stream:\n passive_health: {failures: 2, eject: 1h}\n") + echoes["a"].stop() + var refused int + for range 30 { + if s.askAndClose(t) == "" { + refused++ + } + } + require.Equal(t, 2, refused, "the member was ejected after two failed connects, and refused no more") + }) + t.Run("connect retries serve its share elsewhere", func(t *testing.T) { + members, echoes := tcpMembers(t, "a", "b", "c") + s := startStreamLB(t, "tcp", members, "", " mechanism: rr\n stream:\n connect_retries: 1\n") + echoes["a"].stop() + before := echoes["b"].accepted.Load() + echoes["c"].accepted.Load() + for range 15 { + require.NotEmpty(t, s.askAndClose(t), "a connection was refused although retries are on") + } + require.EqualValues(t, 15, echoes["b"].accepted.Load()+echoes["c"].accepted.Load()-before) + }) +} + +// a reload changes where new connections go and leaves the open ones where they are +func TestStreamLBReloadMidTraffic(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + signal.Reset(syscall.SIGHUP) + members, _ := tcpMembers(t, "a", "b") + s := startStreamLB(t, "tcp", members[:1], "", " mechanism: rr\n") + name, held := s.ask(t) + require.Equal(t, "a", name) + defer held.Close() + + s.write(t, members[1:], "", " mechanism: lc\n") + // the daemon reloads only a config whose file is newer than the one it loaded + future := time.Now().Add(2 * time.Second) + require.NoError(t, os.Chtimes(s.cfgPath, future, future)) + require.NoError(t, syscall.Kill(os.Getpid(), syscall.SIGHUP)) + require.Eventually(t, func() bool { return s.askAndClose(t) == "b" }, 15*time.Second, 100*time.Millisecond, + "new connections never reached the reloaded pool") + + _ = held.SetDeadline(time.Now().Add(5 * time.Second)) + _, err := held.Write([]byte("still\n")) + require.NoError(t, err) + reply, err := bufio.NewReader(held).ReadString('\n') + require.NoError(t, err) + require.Equal(t, "a:still\n", reply, "a connection open across the reload moved or broke") +} + +// a backup member takes connections only while no other member is available +func TestStreamLBBackupMember(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + members, echoes := tcpMembers(t, "primary", "standby") + probe := " healthcheck:\n interval: 100ms\n timeout: 500ms\n failure_threshold: 1\n recovery_threshold: 1\n" + s := startStreamLB(t, "tcp", members, probe, " mechanism: rr\n healthy_floor: 1\n", "standby") + healthURL := fmt.Sprintf("http://127.0.0.1:%d/trickster/health", s.ports[1]) + requireHealthState(t, healthURL, "primary", "available", 10*time.Second) + requireHealthState(t, healthURL, "standby", "available", 10*time.Second) + for range 10 { + require.Equal(t, "primary", s.askAndClose(t), "the standby took a connection while the primary was up") + } + echoes["primary"].stop() + requireHealthState(t, healthURL, "primary", "unavailable", 10*time.Second) + for range 5 { + require.Equal(t, "standby", s.askAndClose(t)) + } + echoes["primary"].restart(t) + requireHealthState(t, healthURL, "primary", "available", 10*time.Second) + require.Eventually(t, func() bool { return s.askAndClose(t) == "primary" }, 5*time.Second, 50*time.Millisecond) + for range 5 { + require.Equal(t, "primary", s.askAndClose(t), "the standby kept taking connections after the primary returned") + } +} + +// a race connects to its members together, so a dead one costs a client nothing +func TestStreamLBConnectRace(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + members, echoes := tcpMembers(t, "a", "b", "c") + s := startStreamLB(t, "tcp", members, "", " mechanism: race\n") + for range 10 { + require.NotEmpty(t, s.askAndClose(t)) + } + echoes["a"].stop() + echoes["b"].stop() + for range 10 { + require.Equal(t, "c", s.askAndClose(t), "the one live member did not win the race") + } +} + +// every member of a mirror receives every datagram, and only the first answers the client +func TestStreamLBUDPMirror(t *testing.T) { + if testing.Short() { + t.Skip("starts Trickster; skipping in -short mode") + } + var mu sync.Mutex + got := map[string][]string{} + sink := func(name string) string { + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = pc.Close() }) + go func() { + buf := make([]byte, 1500) + for { + n, from, err := pc.ReadFrom(buf) + if err != nil { + return + } + mu.Lock() + got[name] = append(got[name], string(buf[:n])) + mu.Unlock() + _, _ = pc.WriteTo([]byte(name+":"+string(buf[:n])), from) + } + }() + return pc.LocalAddr().String() + } + members := [][2]string{{"a", sink("a")}, {"b", sink("b")}, {"c", sink("c")}} + s := startStreamLB(t, "udp", members, "", " mechanism: mirror\n") + conn, err := net.Dial("udp", s.addr) + require.NoError(t, err) + defer conn.Close() + for _, msg := range []string{"one", "two", "three"} { + _, _ = conn.Write([]byte(msg)) + _ = conn.SetReadDeadline(time.Now().Add(3 * time.Second)) + buf := make([]byte, 32) + n, err := conn.Read(buf) + require.NoError(t, err) + require.Equal(t, "a:"+msg, string(buf[:n])) + } + require.Eventually(t, func() bool { + mu.Lock() + defer mu.Unlock() + return len(got["a"]) == 3 && len(got["b"]) == 3 && len(got["c"]) == 3 + }, 5*time.Second, 50*time.Millisecond, "not every member received every datagram") + _ = conn.SetReadDeadline(time.Now().Add(200 * time.Millisecond)) + _, err = conn.Read(make([]byte, 32)) + require.Error(t, err, "a mirror's reply reached the client") +} diff --git a/integration/testdata/alb.yaml b/integration/testdata/alb.yaml index 62fc413a3..d71176128 100644 --- a/integration/testdata/alb.yaml +++ b/integration/testdata/alb.yaml @@ -93,9 +93,9 @@ backends: alb: mechanism: rr pool: - - prom1 + - name: prom1 + weight: 2 - prom2 - - prom1 alb-tsm: provider: alb alb: diff --git a/pkg/backends/alb/client.go b/pkg/backends/alb/client.go index 2781de1ca..49ec38356 100644 --- a/pkg/backends/alb/client.go +++ b/pkg/backends/alb/client.go @@ -31,6 +31,8 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/ur" "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/native" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/observe" "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" bo "github.com/trickstercache/trickster/v2/pkg/backends/options" @@ -38,6 +40,7 @@ import ( rt "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" "github.com/trickstercache/trickster/v2/pkg/cache" "github.com/trickstercache/trickster/v2/pkg/errors" + "github.com/trickstercache/trickster/v2/pkg/lb" "github.com/trickstercache/trickster/v2/pkg/observability/keys" "github.com/trickstercache/trickster/v2/pkg/observability/logging" "github.com/trickstercache/trickster/v2/pkg/observability/logging/logger" @@ -77,6 +80,67 @@ type Client struct { // dynamicNames is the current discovered member-name list, for // health/mgmt display dynamicNames atomic.Pointer[[]string] + + // healthMtx orders what the current pool's snapshots say about the ALB as a whole, and is + // held while the ALBs that pool this one react; the pool graph has no cycles + healthMtx sync.Mutex + // poolID names the current pool and lastGen its newest snapshot seen, so a superseded + // pool or a snapshot observed late cannot overwrite newer state + poolID, lastGen uint64 + // hasBackups and onBackup track failover to the pool's backup members + hasBackups, onBackup bool + // health is the ALB's own status, which follows its pool; nil unless propagate_health + health *healthcheck.Status +} + +// poolObserver reports one pool's snapshots to the client that built it +type poolObserver struct { + c *Client + id uint64 +} + +func (o poolObserver) Observe(ev lb.Event) { + if ev.Kind == lb.EventSnapshot { + o.c.observeSnapshot(o.id, ev) + } +} + +func (c *Client) observeSnapshot(id uint64, ev lb.Event) { + c.healthMtx.Lock() + defer c.healthMtx.Unlock() + if id != c.poolID || ev.Gen <= c.lastGen { + return + } + c.lastGen = ev.Gen + if onBackup := ev.Tier > 0; c.hasBackups && onBackup != c.onBackup { + c.onBackup = onBackup + c.reportFailover(onBackup) + } + if c.health == nil { + return + } + if ev.Eligible > 0 { + c.health.Set(healthcheck.StatusPassing) + return + } + c.health.Set(healthcheck.StatusFailing) +} + +func (c *Client) reportFailover(onBackup bool) { + if !onBackup { + metrics.ALBPoolOnBackup.WithLabelValues(c.Name()).Set(0) + logger.Info("alb pool returned to its primary members", logging.Pairs{keys.BackendName: c.Name()}) + return + } + metrics.ALBPoolOnBackup.WithLabelValues(c.Name()).Set(1) + logger.Warn("alb pool failed over to its backup members: no other member is available", + logging.Pairs{keys.BackendName: c.Name()}) +} + +// HealthStatus returns the status that ALBs pooling this one watch, which follows whether +// this ALB has an available member. It is nil unless propagate_health is set. +func (c *Client) HealthStatus() *healthcheck.Status { + return c.health } // Handlers returns a map of the HTTP Handlers the client has registered. @@ -117,6 +181,13 @@ func NewClient(name string, o *bo.Options, router http.Handler, return nil, err } c.handler = m + if o.ALBOptions.PropagateHealth { + c.health = healthcheck.NewStatus(name, providers.ALB, "", + healthcheck.StatusUnchecked, time.Time{}, nil) + } + if pm, ok := m.(types.PickerMechanism); ok { + observe.Track(name, pm.Balancer()) + } } return c, nil } @@ -125,6 +196,7 @@ func NewClient(name string, o *bo.Options, router http.Handler, // until all backends are processed, so the ALB's destination backend names // can be mapped to their respective clients func StartALBPools(clients backends.Backends, hcs healthcheck.StatusLookup) error { + defer forgetStatsExcept(clients) for _, c := range clients { if rc, ok := c.(*Client); ok { err := rc.ValidateAndStartPool(clients, hcs) @@ -220,7 +292,18 @@ func (c *Client) ValidateAndStartPool(clients backends.Backends, hcs healthcheck return c.validateAndStartUserRouter(clients, hcs) } targets := make(pool.Targets, 0, len(o.Pool)) + stats := make(map[string]*lb.Stats, len(o.Pool)) + tracksStats := false + if pm, ok := c.handler.(types.PickerMechanism); ok { + tracksStats = pm.Picker().Needs() != 0 + } + seen := sets.NewStringSet() for _, m := range o.Pool { + // a pool holds each member once; options loaded from config are already de-duplicated + if seen.Contains(m.Name) { + continue + } + seen.Set(m.Name) tc, ok := clients[m.Name] if !ok { return alberr.NewErrInvalidPoolMemberName(c.Name(), m.Name) @@ -231,12 +314,27 @@ func (c *Client) ValidateAndStartPool(clients backends.Backends, hcs healthcheck } } hc, ok := hcs[m.Name] + if ac, isALB := tc.(*Client); isALB && ac.health != nil { + // a load balancer that reports whether it has an available member is followed by + // that status, whatever the health checker holds for it + hc, ok = ac.health, true + } if !ok { // virtual backends (rule, alb) have no health checks; treat as passing hc = healthcheck.NewStatus(m.Name, "virtual", "", healthcheck.StatusPassing, time.Time{}, nil) } - targets = append(targets, - pool.NewWeightedTarget(tc.Router(), hc, tc, m.EffectiveWeight())) + // only a mechanism that keeps stats has any worth carrying over a reload + var kept *lb.Stats + if tracksStats { + kept = carryStats(c.Name(), m.Name) + } + t := pool.NewWeightedTarget(tc.Router(), hc, tc, m.EffectiveWeight()). + WithTier(m.Tier()).WithStats(kept) + targets = append(targets, t) + stats[m.Name] = t.Member().Stats() + } + if tracksStats { + rememberStats(c.Name(), stats) } c.poolMtx.Lock() c.staticTargets = targets @@ -274,7 +372,20 @@ func (c *Client) swapPool(targets pool.Targets) { return } oldPool := pm.Pool() - pm.SetPool(pool.New(targets, c.effectiveFloor(targets))) + hasBackups := slices.ContainsFunc(targets, func(t *pool.Target) bool { return t != nil && t.Tier() > 0 }) + c.healthMtx.Lock() + c.poolID++ + c.lastGen = 0 + if c.hasBackups && !hasBackups { + metrics.ALBPoolOnBackup.DeleteLabelValues(c.Name()) + c.onBackup = false + } else if hasBackups && !c.hasBackups { + metrics.ALBPoolOnBackup.WithLabelValues(c.Name()).Set(0) + } + c.hasBackups = hasBackups + observer := poolObserver{c: c, id: c.poolID} + c.healthMtx.Unlock() + pm.SetPool(pool.New(targets, c.effectiveFloor(targets), observer)) if oldPool != nil { oldPool.Stop() } @@ -346,6 +457,25 @@ func (c *Client) Pool() pool.Pool { return nil } +// Picker returns the balancer that commits one unit of work to one pool member, which is how +// planes other than HTTP dispatch through this ALB. It is nil for a mechanism that does not +// select one member: the fanout mechanisms and the user router. +func (c *Client) Picker() lb.Picker { + if pm, ok := c.handler.(types.PickerMechanism); ok { + return pm.Picker() + } + return nil +} + +// Spread returns how the ALB's mechanism commits one flow to several members at once, or 0 +// for a mechanism that does not. +func (c *Client) Spread() types.Spread { + if sm, ok := c.handler.(types.SpreadMechanism); ok { + return sm.Spread() + } + return 0 +} + // DynamicPoolNames returns the names of the ALB's currently-discovered pool // members, for health and management display func (c *Client) DynamicPoolNames() []string { @@ -584,12 +714,21 @@ func (c *Client) validateAndStartUserRouter(clients backends.Backends, hcs healt return nil } -// RouteResolver returns the protocol-neutral resolver implemented by a User -// Router ALB. Other ALB mechanisms do not select routes by authenticated user. +// RouteResolver returns the protocol-neutral resolver a native protocol listener commits its +// sessions by: a User Router's own, or one that balances sessions with the ALB's selection +// strategy where that strategy serves sessions. It is nil for any other mechanism. func (c *Client) RouteResolver() backends.RouteResolver { if h, ok := c.handler.(backends.RouteResolver); ok { return h } + cfg := c.Configuration() + if cfg == nil || cfg.ALBOptions == nil || + !registry.Supports(cfg.ALBOptions.MechanismName, types.PlaneNative) { + return nil + } + if pm, ok := c.handler.(types.PickerMechanism); ok { + return native.Resolver(pm.Picker(), cfg.ALBOptions) + } return nil } @@ -603,6 +742,9 @@ func (c *Client) StopPool() { if pm, ok := c.handler.(types.PoolMechanism); ok { pm.StopPool() } + if pm, ok := c.handler.(types.PickerMechanism); ok { + observe.Untrack(c.Name(), pm.Balancer()) + } } // Boilerplate Interface Functions (to EOF) diff --git a/pkg/backends/alb/client_dynamic_test.go b/pkg/backends/alb/client_dynamic_test.go index d6752443d..23783d2e1 100644 --- a/pkg/backends/alb/client_dynamic_test.go +++ b/pkg/backends/alb/client_dynamic_test.go @@ -24,11 +24,16 @@ import ( "testing" "time" + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" bo "github.com/trickstercache/trickster/v2/pkg/backends/options" "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/lb" + + "github.com/prometheus/client_golang/prometheus" ) type countingHandler struct{ hits int } @@ -167,3 +172,150 @@ func TestSetDynamicTargetsUnderLoad(t *testing.T) { t.Error("expected requests to reach pool members during the swaps") } } + +func TestClientPicker(t *testing.T) { + c := newRRALB(t, "picker-alb") + defer c.StopPool() + pk := c.Picker() + if pk == nil { + t.Fatal("a round robin ALB has a picker") + } + if _, ok := pk.Pick(lb.Flow{}); ok { + t.Error("picked before any member was installed") + } + target := pool.NewWeightedTarget(&countingHandler{}, passingStatus(), nil, 1) + if !c.SetDynamicTargets(pool.Targets{target}) { + t.Fatal("swap rejected") + } + got, ok := pk.Pick(lb.Flow{}) + if !ok || got.Member() != target.Member() { + t.Error("the picker does not follow the ALB's pool") + } + + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = &ao.Options{MechanismName: "fr"} + cl, err := NewClient("fanout-alb", o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + if cl.(*Client).Picker() != nil { + t.Error("a fanout ALB does not pick one member") + } +} + +func memberInflightSeries(t *testing.T, albName string) map[string]float64 { + t.Helper() + families, err := prometheus.DefaultGatherer.Gather() + if err != nil { + t.Fatal(err) + } + out := make(map[string]float64) + for _, f := range families { + if f.GetName() != "trickster_alb_member_inflight" { + continue + } + for _, m := range f.GetMetric() { + labels := make(map[string]string) + for _, l := range m.GetLabel() { + labels[l.GetName()] = l.GetValue() + } + if labels["alb_name"] == albName { + out[labels["member"]] = m.GetGauge().GetValue() + } + } + } + return out +} + +// an ALB whose mechanism tracks in-flight work exports it per member until its pool stops +func TestMemberInflightMetricFollowsTheALB(t *testing.T) { + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = &ao.Options{MechanismName: "lc"} + cl, err := NewClient("inflight-metric-alb", o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := cl.(*Client) + release := make(chan struct{}) + entered := make(chan struct{}) + mo := bo.New() + mo.Name = "slow-member" + member, err := backends.New("slow-member", mo, nil, http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + entered <- struct{}{} + <-release + }), nil) + if err != nil { + t.Fatal(err) + } + if !c.SetDynamicTargets(pool.Targets{pool.NewTarget(member.Router(), passingStatus(), member)}) { + t.Fatal("swap rejected") + } + var wg sync.WaitGroup + wg.Go(func() { + c.Handlers()[providers.ALB].ServeHTTP(httptest.NewRecorder(), + httptest.NewRequest(http.MethodGet, "http://example.com/", nil)) + }) + <-entered + if got := memberInflightSeries(t, "inflight-metric-alb"); got["slow-member"] != 1 { + t.Errorf("series while a request is held = %v", got) + } + close(release) + wg.Wait() + if got := memberInflightSeries(t, "inflight-metric-alb"); got["slow-member"] != 0 { + t.Errorf("series after the request = %v", got) + } + c.StopPool() + if got := memberInflightSeries(t, "inflight-metric-alb"); len(got) != 0 { + t.Errorf("series after the pool stopped = %v", got) + } + // a mechanism that tracks nothing exports nothing + rr := newRRALB(t, "untracked-metric-alb") + defer rr.StopPool() + if got := memberInflightSeries(t, "untracked-metric-alb"); len(got) != 0 { + t.Errorf("round robin exported %v", got) + } +} + +func TestClientSpread(t *testing.T) { + if got := newRRALB(t, "spread-rr").Spread(); got != 0 { + t.Errorf("a round robin ALB spreads: %d", got) + } + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = &ao.Options{MechanismName: "mirror"} + cl, err := NewClient("spread-mirror", o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + c := cl.(*Client) + if c.Spread() != types.SpreadMirror || c.Picker() != nil || c.RouteResolver() != nil { + t.Errorf("mirror ALB: spread %d, picker %v", c.Spread(), c.Picker()) + } +} + +func TestClientRouteResolver(t *testing.T) { + if newRRALB(t, "sessions-rr").RouteResolver() == nil { + t.Error("a round robin ALB does not balance sessions") + } + for _, mechanism := range []string{"fr", "lt"} { + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = &ao.Options{MechanismName: mechanism} + cl, err := NewClient("sessions-"+mechanism, o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + if cl.(*Client).RouteResolver() != nil { + t.Errorf("a %s ALB resolves routes", mechanism) + } + } + bare, err := backends.New("bare", nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + if (&Client{Backend: bare}).RouteResolver() != nil { + t.Error("an ALB with no options resolves routes") + } +} diff --git a/pkg/backends/alb/dynamic/manager.go b/pkg/backends/alb/dynamic/manager.go index c1a240d91..159f1c794 100644 --- a/pkg/backends/alb/dynamic/manager.go +++ b/pkg/backends/alb/dynamic/manager.go @@ -26,7 +26,6 @@ package dynamic import ( "context" - "errors" "fmt" "net/http" "slices" @@ -91,12 +90,6 @@ type externalRegistrar interface { RegisterExternal(name, description string, s *healthcheck.Status) } -// protocolHealthProber mirrors the unexported interface consulted by -// backends.StartHealthChecks for protocol-native probes (e.g. mysql) -type protocolHealthProber interface { - HealthCheckProbe() healthcheck.Probe -} - // memberEntry is one live discovered member type memberEntry struct { member discovery.Member @@ -414,48 +407,33 @@ func (m *Manager) instantiateMember(name string, member discovery.Member) (*memb client, nb, c, m.cfg.Tracers) e := &memberEntry{member: member, client: client} - if m.cfg.Options.HealthMode == ao.HealthModeProvider { + if m.cfg.Options.HealthMode != ao.HealthModeProvider && nb.HealthCheck != nil { + // the same registration a configured backend gets + st, err := backends.RegisterHealthCheck(m.cfg.HealthChecker, name, + m.healthDescription(nb.Provider), client) + if err != nil { + return nil, err + } + if st != nil { + if oldSt, ok := m.cfg.KnownStatuses[name]; ok && oldSt != nil { + if v := oldSt.Get(); v != healthcheck.StatusInitializing { + st.Set(v) + } + } + m.admitOnReadiness(name, st, member) + client.SetHealthCheckProbe(st.Prober()) + e.status = st + } + } + if e.status == nil && (m.cfg.Options.HealthMode == ao.HealthModeProvider || nb.HealthCheck != nil) { + // the provider's readiness is the member's health: by configuration, or because its + // origin is one that cannot be probed e.external = true e.status = healthcheck.NewStatus(name, m.healthDescription(nb.Provider), "", statusForReadyState(member.Ready), time.Time{}, nil) if er, ok := m.cfg.HealthChecker.(externalRegistrar); ok { er.RegisterExternal(name, m.healthDescription(nb.Provider), e.status) } - } else if nb.HealthCheck != nil { - // mirror backends.StartHealthChecks: overlay the provider default - // healthcheck config, then register an active probe - hco := nb.HealthCheck - nb.HealthCheck = client.DefaultHealthCheckConfig() - if nb.HealthCheck == nil { - nb.HealthCheck = hco - } else { - nb.HealthCheck.Overlay(hco) - } - var st *healthcheck.Status - if prober, ok := client.(protocolHealthProber); ok { - registrar, rok := m.cfg.HealthChecker.(healthcheck.Registrar) - if !rok { - return nil, errors.New("health checker does not support protocol probes") - } - st, err = registrar.RegisterProbe(name, - m.healthDescription(nb.Provider), nb.HealthCheck, - prober.HealthCheckProbe()) - } else { - st, err = m.cfg.HealthChecker.Register(name, - m.healthDescription(nb.Provider), - nb.HealthCheck, client.HealthCheckHTTPClient()) - } - if err != nil { - return nil, err - } - if oldSt, ok := m.cfg.KnownStatuses[name]; ok && oldSt != nil { - if v := oldSt.Get(); v != healthcheck.StatusInitializing { - st.Set(v) - } - } - m.admitOnReadiness(name, st, member) - client.SetHealthCheckProbe(st.Prober()) - e.status = st } e.target = pool.NewWeightedTarget(client.Router(), e.status, client, @@ -501,9 +479,9 @@ func (m *Manager) updateMember(name string, e *memberEntry, member discovery.Mem } if member.Weight != e.member.Weight { // targets are immutable; rebuild this member's target around the - // same client and status + // same client, status and runtime stats e.target = pool.NewWeightedTarget(e.client.Router(), e.status, - e.client, member.Weight) + e.client, member.Weight).WithStatsOf(e.target) if e.external { e.target = e.target.WithExternalHealth() } diff --git a/pkg/backends/alb/dynamic/manager_test.go b/pkg/backends/alb/dynamic/manager_test.go index bc8a9cefd..864b0ae90 100644 --- a/pkg/backends/alb/dynamic/manager_test.go +++ b/pkg/backends/alb/dynamic/manager_test.go @@ -17,8 +17,10 @@ package dynamic import ( + "net" "net/http" "net/http/httptest" + "sync/atomic" "testing" "time" @@ -344,3 +346,87 @@ func TestManagerProbeModeReadinessWithoutProbe(t *testing.T) { require.NotPanics(t, func() { m.ApplySnapshot(discovery.Snapshot{mem}) }) require.NotContains(t, hc.Statuses(), "myalb-m1") } + +// a discovered member is health checked the way a configured one is: a udp origin, which +// nothing can probe, follows the provider's readiness instead of an http probe that can only +// fail, and a tcp origin is probed by connecting to it +func TestManagerProbeModeByOriginProtocol(t *testing.T) { + m, c, hc := newTestManager(t, &ao.DiscoveryOptions{ + DiscovererName: "d", TemplateBackend: "rp-template", + }) + probedTemplate(m) + m.cfg.Template.HealthCheck.FailureThreshold = 1 + + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + + datagrams := discovery.Member{Name: "dns", Scheme: "udp", Address: "10.0.0.9:53", Ready: discovery.Ready} + pending := discovery.Member{Name: "pending", Scheme: "udp", Address: "10.0.0.10:53", Ready: discovery.NotReady} + stream := discovery.Member{Name: "db", Scheme: "tcp", Address: ln.Addr().String()} + m.ApplySnapshot(discovery.Snapshot{datagrams, pending, stream}) + + statuses := hc.Statuses() + require.Equal(t, healthcheck.StatusPassing, statuses["myalb-dns"].Get(), + "a udp member the provider reports ready is available") + require.Equal(t, healthcheck.StatusFailing, statuses["myalb-pending"].Get()) + require.Eventually(t, func() bool { return statuses["myalb-db"].Get() == healthcheck.StatusPassing }, + 5*time.Second, 10*time.Millisecond, "a tcp member is probed by connecting to it") + // no probe ever runs against the udp member, so it never goes down on one + time.Sleep(150 * time.Millisecond) + require.Equal(t, healthcheck.StatusPassing, statuses["myalb-dns"].Get()) + + names := []string{} + for _, tgt := range c.Pool().Targets() { + names = append(names, tgt.Name()) + } + require.ElementsMatch(t, []string{"myalb-dns", "myalb-db"}, names) + + // and it goes on following the provider + pending.Ready = discovery.Ready + datagrams.Ready = discovery.NotReady + m.ApplySnapshot(discovery.Snapshot{datagrams, pending, stream}) + require.Equal(t, healthcheck.StatusPassing, statuses["myalb-pending"].Get()) + require.Equal(t, healthcheck.StatusFailing, statuses["myalb-dns"].Get()) +} + +// a member that keeps its name while its origin changes from one that is probed to one that +// cannot be leaves no probe behind: the retired origin is not contacted again +func TestManagerProbeRetiredWhenAMemberBecomesUnprobeable(t *testing.T) { + m, _, hc := newTestManager(t, &ao.DiscoveryOptions{ + DiscovererName: "d", TemplateBackend: "rp-template", + }) + probedTemplate(m) + m.cfg.Template.HealthCheck.Interval = timeconv.Duration(5 * time.Millisecond) + + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + var accepted atomic.Int64 + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + accepted.Add(1) + _ = conn.Close() + } + }() + + m.ApplySnapshot(discovery.Snapshot{{Name: "db", Scheme: "tcp", Address: ln.Addr().String()}}) + require.Eventually(t, func() bool { return accepted.Load() > 2 }, 5*time.Second, 5*time.Millisecond, + "the tcp member was never probed") + probed := hc.Statuses()["myalb-db"] + + m.ApplySnapshot(discovery.Snapshot{{Name: "db", Scheme: "udp", Address: "10.0.0.9:53", Ready: discovery.Ready}}) + st := hc.Statuses()["myalb-db"] + require.NotSame(t, probed, st, "the member kept the status of the origin it left") + require.Equal(t, healthcheck.StatusPassing, st.Get()) + // a probe in flight when the origin changed may still land; none starts after it + time.Sleep(50 * time.Millisecond) + settled := accepted.Load() + time.Sleep(100 * time.Millisecond) + require.Equal(t, settled, accepted.Load(), "the retired origin is still being probed") + require.Equal(t, healthcheck.StatusPassing, hc.Statuses()["myalb-db"].Get()) +} diff --git a/pkg/backends/alb/dynamic/weights_characterization_test.go b/pkg/backends/alb/dynamic/weights_characterization_test.go new file mode 100644 index 000000000..11021ecde --- /dev/null +++ b/pkg/backends/alb/dynamic/weights_characterization_test.go @@ -0,0 +1,81 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package dynamic + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/discovery" + + "github.com/stretchr/testify/require" +) + +func weightedMember(name, addr string, weight int) discovery.Member { + m := member(name, addr) + m.Weight = weight + return m +} + +func poolTargets(t *testing.T, c *alb.Client) map[string]*pool.Target { + t.Helper() + p := c.Pool() + require.NotNil(t, p) + out := make(map[string]*pool.Target) + for _, tgt := range p.ConfiguredTargets() { + out[tgt.Name()] = tgt + } + return out +} + +// a discovered member's weight reaches its pool target unchanged, and a live weight change +// re-weights that member in place: same backend, same health status, siblings untouched +func TestDiscoveredWeightsReachThePool(t *testing.T) { + m, c, _ := newTestManager(t, &ao.DiscoveryOptions{ + DiscovererName: "d", TemplateBackend: "rp-template", + }) + m.ApplySnapshot(discovery.Snapshot{ + weightedMember("m1", "10.0.0.1:8080", 0), + weightedMember("m2", "10.0.0.2:8080", 3), + weightedMember("m3", "10.0.0.3:8080", 5), + }) + before := poolTargets(t, c) + require.Len(t, before, 3) + require.Equal(t, 1, before["myalb-m1"].Weight(), "an unset weight is 1") + require.Equal(t, 3, before["myalb-m2"].Weight()) + require.Equal(t, 5, before["myalb-m3"].Weight()) + + m.ApplySnapshot(discovery.Snapshot{ + weightedMember("m1", "10.0.0.1:8080", 0), + weightedMember("m2", "10.0.0.2:8080", 7), + weightedMember("m3", "10.0.0.3:8080", 5), + }) + after := poolTargets(t, c) + require.Len(t, after, 3) + require.Equal(t, 7, after["myalb-m2"].Weight()) + require.Same(t, before["myalb-m2"].Backend(), after["myalb-m2"].Backend(), + "a weight change must not rebuild the member's backend") + require.Same(t, before["myalb-m2"].HealthStatus(), after["myalb-m2"].HealthStatus(), + "a weight change must not reset the member's health") + require.Same(t, before["myalb-m2"].Member().Stats(), after["myalb-m2"].Member().Stats(), + "a weight change must not reset the member's runtime stats") + require.Equal(t, 7, after["myalb-m2"].Member().Weight()) + require.Same(t, before["myalb-m1"], after["myalb-m1"], "an unchanged member keeps its target") + require.Same(t, before["myalb-m3"], after["myalb-m3"], "an unchanged member keeps its target") +} diff --git a/pkg/backends/alb/failover_test.go b/pkg/backends/alb/failover_test.go new file mode 100644 index 000000000..f3d6d8906 --- /dev/null +++ b/pkg/backends/alb/failover_test.go @@ -0,0 +1,211 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package alb + +import ( + "net/http" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + + "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/require" +) + +// graph builds and starts ALBs over origin backends whose health the test drives +type graph struct { + clients backends.Backends + health healthcheck.StatusLookup +} + +func newGraph(t *testing.T, origins ...string) *graph { + t.Helper() + g := &graph{clients: backends.Backends{}, health: healthcheck.StatusLookup{}} + for _, name := range origins { + o := bo.New() + o.Name = name + b, err := backends.New(name, o, nil, http.NotFoundHandler(), nil) + require.NoError(t, err) + g.clients[name] = b + g.health[name] = healthcheck.NewStatus(name, "", "", healthcheck.StatusPassing, time.Time{}, nil) + } + return g +} + +func (g *graph) alb(t *testing.T, name string, o *ao.Options) *Client { + t.Helper() + bo := bo.New() + bo.Name = name + bo.Provider = providers.ALB + bo.ALBOptions = o + require.NoError(t, bo.ALBOptions.Initialize(name)) + cl, err := NewClient(name, bo, nil, nil, nil, nil) + require.NoError(t, err) + g.clients[name] = cl + return cl.(*Client) +} + +func (g *graph) start(t *testing.T) { + t.Helper() + require.NoError(t, StartALBPools(g.clients, g.health)) + t.Cleanup(func() { _ = StopPools(g.clients) }) +} + +func liveNames(c *Client) []string { + out := []string{} + for _, tgt := range c.Pool().Targets() { + out = append(out, tgt.Name()) + } + return out +} + +func TestBackupMembersStandBy(t *testing.T) { + g := newGraph(t, "a", "b", "standby") + c := g.alb(t, "failover-alb", &ao.Options{MechanismName: "rr", Pool: ao.PoolMemberList{ + {Name: "a"}, {Name: "standby", Backup: true}, {Name: "b"}, + }}) + g.start(t) + onBackup := func() float64 { return testutil.ToFloat64(metrics.ALBPoolOnBackup.WithLabelValues("failover-alb")) } + require.Equal(t, []string{"a", "b"}, liveNames(c)) + require.Zero(t, onBackup()) + g.health["a"].Set(healthcheck.StatusFailing) + require.Equal(t, []string{"b"}, liveNames(c)) + g.health["b"].Set(healthcheck.StatusFailing) + require.Equal(t, []string{"standby"}, liveNames(c)) + require.Equal(t, 1.0, onBackup()) + pk, ok := c.Picker().Pick(lb.Flow{}) + require.True(t, ok) + require.Equal(t, "standby", pk.Member().Name()) + pk.Done(lb.OutcomeOK) + g.health["a"].Set(healthcheck.StatusPassing) + require.Equal(t, []string{"a"}, liveNames(c)) + require.Zero(t, onBackup()) + + // a pool that no longer has a backup no longer reports on one + series := func() int { return testutil.CollectAndCount(metrics.ALBPoolOnBackup) } + before := series() + c.poolMtx.Lock() + c.staticTargets = c.staticTargets[:1] + c.swapPool(c.staticTargets) + c.poolMtx.Unlock() + require.Equal(t, before-1, series()) +} + +func TestHealthPropagatesOnlyWhenAsked(t *testing.T) { + g := newGraph(t, "a1", "a2", "b1") + inner := g.alb(t, "inner-a", &ao.Options{ + MechanismName: "rr", Pool: ao.Members("a1", "a2"), PropagateHealth: true, + }) + silent := g.alb(t, "inner-b", &ao.Options{MechanismName: "rr", Pool: ao.Members("b1")}) + outer := g.alb(t, "outer", &ao.Options{ + MechanismName: "rr", Pool: ao.Members("inner-a", "inner-b"), PropagateHealth: true, + }) + top := g.alb(t, "top", &ao.Options{MechanismName: "rr", Pool: ao.Members("outer")}) + g.start(t) + require.Nil(t, silent.HealthStatus()) + require.Equal(t, healthcheck.StatusPassing, inner.HealthStatus().Get()) + require.Equal(t, []string{"inner-a", "inner-b"}, liveNames(outer)) + + g.health["a1"].Set(healthcheck.StatusFailing) + require.Equal(t, []string{"inner-a", "inner-b"}, liveNames(outer), "one live member is enough") + g.health["a2"].Set(healthcheck.StatusFailing) + require.Equal(t, healthcheck.StatusFailing, inner.HealthStatus().Get()) + require.Equal(t, []string{"inner-b"}, liveNames(outer), "an empty pool sheds its share") + + // a load balancer that does not propagate keeps its share and fails it + g.health["b1"].Set(healthcheck.StatusFailing) + require.Equal(t, []string{"inner-b"}, liveNames(outer)) + require.Equal(t, healthcheck.StatusPassing, outer.HealthStatus().Get()) + require.Equal(t, []string{"outer"}, liveNames(top)) + + g.health["a2"].Set(healthcheck.StatusPassing) + require.Equal(t, []string{"inner-a", "inner-b"}, liveNames(outer)) +} + +func TestHealthPropagatesThroughEveryLevel(t *testing.T) { + g := newGraph(t, "leaf", "other") + g.alb(t, "low", &ao.Options{MechanismName: "rr", Pool: ao.Members("leaf"), PropagateHealth: true}) + mid := g.alb(t, "mid", &ao.Options{MechanismName: "rr", Pool: ao.Members("low"), PropagateHealth: true}) + top := g.alb(t, "top", &ao.Options{MechanismName: "rr", Pool: ao.Members("mid", "other")}) + g.start(t) + require.Equal(t, []string{"mid", "other"}, liveNames(top)) + g.health["leaf"].Set(healthcheck.StatusFailing) + require.Equal(t, healthcheck.StatusFailing, mid.HealthStatus().Get()) + require.Equal(t, []string{"other"}, liveNames(top)) + g.health["leaf"].Set(healthcheck.StatusPassing) + require.Equal(t, []string{"mid", "other"}, liveNames(top)) +} + +func TestSnapshotsOfASupersededPoolAreIgnored(t *testing.T) { + g := newGraph(t, "a") + c := g.alb(t, "stale-alb", &ao.Options{MechanismName: "rr", Pool: ao.Members("a"), PropagateHealth: true}) + g.start(t) + require.Equal(t, healthcheck.StatusPassing, c.HealthStatus().Get()) + stale := poolObserver{c: c, id: c.poolID - 1} + stale.Observe(lb.Event{Kind: lb.EventSnapshot, Gen: 99, Eligible: 0}) + require.Equal(t, healthcheck.StatusPassing, c.HealthStatus().Get()) + // nor is a snapshot of the current pool that is observed after a newer one + late := poolObserver{c: c, id: c.poolID} + late.Observe(lb.Event{Kind: lb.EventSnapshot, Gen: c.lastGen, Eligible: 0}) + late.Observe(lb.Event{Kind: lb.EventPanic}) + require.Equal(t, healthcheck.StatusPassing, c.HealthStatus().Get()) +} + +type externalRegistrar interface { + RegisterExternal(name, description string, s *healthcheck.Status) +} + +// the daemon registers every ALB with the health checker before it starts the pools; the status +// an ALB keeps for itself must win over the synthetic one, and be the one the checker reports +func TestHealthPropagatesThroughTheHealthChecker(t *testing.T) { + g := newGraph(t, "a1", "b1") + inner := g.alb(t, "inner-a", &ao.Options{MechanismName: "rr", Pool: ao.Members("a1"), PropagateHealth: true}) + g.alb(t, "inner-b", &ao.Options{MechanismName: "rr", Pool: ao.Members("b1")}) + outer := g.alb(t, "outer", &ao.Options{MechanismName: "rr", Pool: ao.Members("inner-a", "inner-b")}) + for _, c := range g.clients { + if c.Configuration().Provider == "" { + c.Configuration().Provider = providers.ReverseProxyShort + } + } + hc, err := g.clients.StartHealthChecks(nil) + require.NoError(t, err) + t.Cleanup(hc.Shutdown) + for name, st := range g.health { + hc.(externalRegistrar).RegisterExternal(name, "test", st) + } + statuses := hc.Statuses() + require.Same(t, inner.HealthStatus(), statuses["inner-a"], "the checker reports the ALB's own status") + require.NotNil(t, statuses["inner-b"], "an ALB with no status of its own is still reported") + require.NoError(t, StartALBPools(g.clients, statuses)) + t.Cleanup(func() { _ = StopPools(g.clients) }) + + require.Equal(t, []string{"inner-a", "inner-b"}, liveNames(outer)) + g.health["a1"].Set(healthcheck.StatusFailing) + require.Equal(t, healthcheck.StatusFailing, statuses["inner-a"].Get()) + require.Equal(t, []string{"inner-b"}, liveNames(outer)) + g.health["b1"].Set(healthcheck.StatusFailing) + require.Equal(t, []string{"inner-b"}, liveNames(outer), "an ALB that does not propagate keeps its share") + g.health["a1"].Set(healthcheck.StatusPassing) + require.Equal(t, []string{"inner-a", "inner-b"}, liveNames(outer)) +} diff --git a/pkg/backends/alb/hotpath_bench_test.go b/pkg/backends/alb/hotpath_bench_test.go index 1879d3b74..52a4bfb3c 100644 --- a/pkg/backends/alb/hotpath_bench_test.go +++ b/pkg/backends/alb/hotpath_bench_test.go @@ -63,8 +63,8 @@ func newDiscoveredPoolALB(t testing.TB, targets pool.Targets) *Client { return c } -// waitZeroAllocSteadyState spins until the pool's async refresh worker -// has drained and dispatch serves the cached zero-alloc fast path +// waitZeroAllocSteadyState spins until dispatch serves the pool's cached +// target view, which the first read of a new snapshot builds func waitZeroAllocSteadyState(h http.Handler, w *httptest.ResponseRecorder, r *http.Request) { deadline := time.Now().Add(2 * time.Second) for testing.AllocsPerRun(1, func() { h.ServeHTTP(w, r) }) != 0 { diff --git a/pkg/backends/alb/mech/fr/first_response.go b/pkg/backends/alb/mech/fr/first_response.go index 40dafd54a..717259a08 100644 --- a/pkg/backends/alb/mech/fr/first_response.go +++ b/pkg/backends/alb/mech/fr/first_response.go @@ -25,11 +25,11 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" rt "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" + cfgtypes "github.com/trickstercache/trickster/v2/pkg/config/types" "github.com/trickstercache/trickster/v2/pkg/proxy/handlers/trickster/failures" "github.com/trickstercache/trickster/v2/pkg/proxy/headers" "github.com/trickstercache/trickster/v2/pkg/proxy/request" "github.com/trickstercache/trickster/v2/pkg/proxy/response/capture" - "github.com/trickstercache/trickster/v2/pkg/util/sets" ) const ( @@ -40,23 +40,23 @@ const ( type handler struct { mech.PoolHolder fgr bool - fgrCodes sets.Set[int] + fgrCodes *cfgtypes.StatusTable options options.FirstGoodResponseOptions maxCaptureBytes int } func RegistryEntry() types.RegistryEntry { - return types.RegistryEntry{Name: FRName, ShortName: names.MechanismFR, New: New} + return types.RegistryEntry{Name: FRName, ShortName: names.MechanismFR, Planes: types.PlaneHTTP, New: New} } func RegistryEntryFGR() types.RegistryEntry { - return types.RegistryEntry{Name: FGRName, ShortName: names.MechanismFGR, New: NewFGR} + return types.RegistryEntry{Name: FGRName, ShortName: names.MechanismFGR, Planes: types.PlaneHTTP, New: NewFGR} } func NewFGR(o *options.Options, _ rt.Lookup) (types.Mechanism, error) { return &handler{ fgr: true, - fgrCodes: o.FgrCodesLookup, + fgrCodes: o.FGRGoodCodes, options: o.FGROptions, maxCaptureBytes: o.MaxCaptureBytes, }, nil @@ -84,19 +84,18 @@ func (h *handler) StopPool() { } // qualifies returns the winner predicate for fanout.WaitForFirst. FR (non- -// FGR) takes any captured response; FGR with no custom codes accepts any -// status < 400; FGR with custom codes accepts only configured codes. +// FGR) takes any captured response; FGR accepts the configured good codes, +// which are any status < 400 unless set otherwise. // Truncated captures are filtered out by WaitForFirst before predicate is // called. func (h *handler) qualifies(r *fanout.Result) bool { if !h.fgr { return true } - code := r.Capture.StatusCode() - if len(h.fgrCodes) > 0 { - return h.fgrCodes.Contains(code) + if h.fgrCodes == nil { + return r.Capture.StatusCode() < 400 } - return code < 400 + return h.fgrCodes.Contains(r.Capture.StatusCode()) } func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { diff --git a/pkg/backends/alb/mech/fr/first_response_hang_test.go b/pkg/backends/alb/mech/fr/first_response_hang_test.go index fdb1789e8..198998e6c 100644 --- a/pkg/backends/alb/mech/fr/first_response_hang_test.go +++ b/pkg/backends/alb/mech/fr/first_response_hang_test.go @@ -45,7 +45,6 @@ func TestFRDoesNotHangWhenAllTargetsAbort(t *testing.T) { }) p, _, _ := albpool.New(-1, []http.Handler{never, never}) defer p.Stop() - p.SetHealthy([]http.Handler{never, never}) h := &handler{} h.SetPool(p) diff --git a/pkg/backends/alb/mech/fr/first_response_test.go b/pkg/backends/alb/mech/fr/first_response_test.go index 9213c6eb0..eae7aa7ed 100644 --- a/pkg/backends/alb/mech/fr/first_response_test.go +++ b/pkg/backends/alb/mech/fr/first_response_test.go @@ -23,9 +23,9 @@ import ( "testing" "time" + cfgtypes "github.com/trickstercache/trickster/v2/pkg/config/types" tu "github.com/trickstercache/trickster/v2/pkg/testutil" "github.com/trickstercache/trickster/v2/pkg/testutil/albpool" - "github.com/trickstercache/trickster/v2/pkg/util/sets" ) func TestHandleFirstResponseNilPool(t *testing.T) { @@ -119,7 +119,7 @@ func TestFirstGoodResponse(t *testing.T) { }) t.Run("FGR custom codes", func(t *testing.T) { - codes := sets.New([]int{http.StatusAccepted}) + codes := cfgtypes.StatusCodes(http.StatusAccepted).Compile() p, _, _ := albpool.NewHealthy([]http.Handler{ albpool.StatusHandler(http.StatusOK, "not-accepted"), albpool.StatusHandler(http.StatusAccepted, "accepted"), @@ -192,7 +192,7 @@ func TestFirstGoodResponse(t *testing.T) { } func TestFGRFallbackEmits502WhenNoMemberQualifies(t *testing.T) { - codes := sets.New([]int{http.StatusOK}) + codes := cfgtypes.StatusCodes(http.StatusOK).Compile() p, _, _ := albpool.NewHealthy([]http.Handler{ albpool.StatusHandler(http.StatusInternalServerError, "body0"), albpool.StatusHandler(http.StatusInternalServerError, "body1"), @@ -235,7 +235,6 @@ func TestHandleFirstResponseContextCancel(t *testing.T) { func() { p, _, _ := albpool.New(-1, []http.Handler{slow, slow}) defer p.Stop() - p.SetHealthy([]http.Handler{slow, slow}) h := &handler{} h.SetPool(p) @@ -282,7 +281,6 @@ func TestHandleFirstResponseContextCancel_50Backends(t *testing.T) { func() { p, _, _ := albpool.New(-1, hs) defer p.Stop() - p.SetHealthy(hs) h := &handler{} h.SetPool(p) diff --git a/pkg/backends/alb/mech/fr/first_response_truncated_test.go b/pkg/backends/alb/mech/fr/first_response_truncated_test.go index 8ec80df08..d8d62cf3f 100644 --- a/pkg/backends/alb/mech/fr/first_response_truncated_test.go +++ b/pkg/backends/alb/mech/fr/first_response_truncated_test.go @@ -22,9 +22,9 @@ import ( "testing" "github.com/trickstercache/trickster/v2/pkg/appinfo" + cfgtypes "github.com/trickstercache/trickster/v2/pkg/config/types" "github.com/trickstercache/trickster/v2/pkg/proxy/headers" "github.com/trickstercache/trickster/v2/pkg/testutil/albpool" - "github.com/trickstercache/trickster/v2/pkg/util/sets" ) // TestFRDisqualifiesTruncatedWinner asserts that FR (FGR variant) does not @@ -41,11 +41,10 @@ func TestFRDisqualifiesTruncatedWinner(t *testing.T) { } p, _, _ := albpool.New(-1, hs) defer p.Stop() - p.SetHealthy(hs) h := &handler{ fgr: true, - fgrCodes: sets.New([]int{http.StatusOK}), + fgrCodes: cfgtypes.StatusCodes(http.StatusOK).Compile(), maxCaptureBytes: maxBytes, } h.SetPool(p) @@ -78,7 +77,6 @@ func TestFRTruncatedAllMembersFallback(t *testing.T) { } p, _, _ := albpool.New(-1, hs) defer p.Stop() - p.SetHealthy(hs) h := &handler{maxCaptureBytes: maxBytes} h.SetPool(p) @@ -105,11 +103,10 @@ func TestFRPrefersIntactOverTruncated(t *testing.T) { } p, _, _ := albpool.New(-1, hs) defer p.Stop() - p.SetHealthy(hs) h := &handler{ fgr: true, - fgrCodes: sets.New([]int{http.StatusOK}), + fgrCodes: cfgtypes.StatusCodes(http.StatusOK).Compile(), maxCaptureBytes: maxBytes, } h.SetPool(p) diff --git a/pkg/backends/alb/mech/nlm/newest_last_modified.go b/pkg/backends/alb/mech/nlm/newest_last_modified.go index 66d98799f..90ddff42c 100644 --- a/pkg/backends/alb/mech/nlm/newest_last_modified.go +++ b/pkg/backends/alb/mech/nlm/newest_last_modified.go @@ -41,7 +41,7 @@ type handler struct { } func RegistryEntry() types.RegistryEntry { - return types.RegistryEntry{Name: Name, ShortName: names.MechanismNLM, New: New} + return types.RegistryEntry{Name: Name, ShortName: names.MechanismNLM, Planes: types.PlaneHTTP, New: New} } func New(o *options.Options, _ rt.Lookup) (types.Mechanism, error) { diff --git a/pkg/backends/alb/mech/nlm/newest_last_modified_invariants_test.go b/pkg/backends/alb/mech/nlm/newest_last_modified_invariants_test.go index 6215fdc04..80ce32513 100644 --- a/pkg/backends/alb/mech/nlm/newest_last_modified_invariants_test.go +++ b/pkg/backends/alb/mech/nlm/newest_last_modified_invariants_test.go @@ -135,7 +135,6 @@ func testNLMAllTruncated(t *testing.T) { hs := []http.Handler{oversized(), oversized(), oversized()} p, _, _ := albpool.NewHealthy(hs) defer p.Stop() - p.SetHealthy(hs) h := &handler{maxCaptureBytes: maxBytes} h.SetPool(p) @@ -159,7 +158,6 @@ func testNLMAllFailed(t *testing.T) { p, _, _ := albpool.NewHealthy(hs) defer p.Stop() - p.SetHealthy(hs) h := &handler{} h.SetPool(p) diff --git a/pkg/backends/alb/mech/nlm/newest_last_modified_test.go b/pkg/backends/alb/mech/nlm/newest_last_modified_test.go index 6654445a2..52366ec9f 100644 --- a/pkg/backends/alb/mech/nlm/newest_last_modified_test.go +++ b/pkg/backends/alb/mech/nlm/newest_last_modified_test.go @@ -217,7 +217,6 @@ func TestHandleNewestContextCancel(t *testing.T) { func() { p, _, _ := albpool.New(-1, hs) defer p.Stop() - p.SetHealthy(hs) h := &handler{} h.SetPool(p) diff --git a/pkg/backends/alb/mech/pick/helpers_test.go b/pkg/backends/alb/mech/pick/helpers_test.go new file mode 100644 index 000000000..642d18117 --- /dev/null +++ b/pkg/backends/alb/mech/pick/helpers_test.go @@ -0,0 +1,38 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package pick + +import ( + "net/http" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" +) + +func newRR() *handler { + return New(names.MechanismRR, rr.New()).(*handler) +} + +// nextTarget selects as a request would, without dispatching +func nextTarget(h *handler) http.Handler { + pk, ok := h.balancer.Pick(lb.Flow{}) + if !ok { + return nil + } + return pk.Member().Value.(*pool.Target).Handler() +} diff --git a/pkg/backends/alb/mech/pick/key.go b/pkg/backends/alb/mech/pick/key.go new file mode 100644 index 000000000..2adcebe2b --- /dev/null +++ b/pkg/backends/alb/mech/pick/key.go @@ -0,0 +1,120 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package pick + +import ( + "net/http" + "net/netip" + "strings" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/lb" + tctx "github.com/trickstercache/trickster/v2/pkg/proxy/context" + "github.com/trickstercache/trickster/v2/pkg/proxy/headers" +) + +// keyFunc derives a request's flow key; it must not allocate +type keyFunc func(*http.Request) lb.Flow + +// newKeyFunc returns the extractor for a key source. v6Prefix applies to client_ip keys. +func newKeyFunc(ks options.KeySource, v6Prefix int) keyFunc { + name := ks.Name + switch ks.Kind { + case options.KeyHost: + return hostKey + case options.KeyHeader: + return func(r *http.Request) lb.Flow { + if v := r.Header[name]; len(v) > 0 { + return keyOf(v[0]) + } + return lb.Flow{} + } + case options.KeyCookie: + return func(r *http.Request) lb.Flow { + for _, line := range r.Header[headers.NameCookie] { + if v, ok := scanPairs(line, name, ';'); ok { + return keyOf(strings.Trim(v, `"`)) + } + } + return lb.Flow{} + } + case options.KeyQuery: + return func(r *http.Request) lb.Flow { + if r.URL == nil { + return lb.Flow{} + } + v, _ := scanPairs(r.URL.RawQuery, name, '&') + return keyOf(v) + } + } + return func(r *http.Request) lb.Flow { return clientIPKey(r, v6Prefix) } +} + +// keyOf is the flow of a key value; an empty value identifies nothing +func keyOf(v string) lb.Flow { + if v == "" { + return lb.Flow{} + } + return lb.Flow{Key: lb.HashString(v), HasKey: true} +} + +// scanPairs finds name=value among the sep-separated pairs of s without allocating. The +// value is returned as written: not unquoted, not unescaped. +func scanPairs(s, name string, sep byte) (string, bool) { + for s != "" { + var pair string + if i := strings.IndexByte(s, sep); i >= 0 { + pair, s = s[:i], s[i+1:] + } else { + pair, s = s, "" + } + pair = strings.TrimLeft(pair, " \t") + if len(pair) > len(name) && pair[len(name)] == '=' && pair[:len(name)] == name { + return strings.TrimRight(pair[len(name)+1:], " \t"), true + } + } + return "", false +} + +func hostKey(r *http.Request) lb.Flow { + host := r.Host + // drop a port, which follows the last colon unless that colon is inside an IPv6 literal + if i := strings.LastIndexByte(host, ':'); i > 0 && strings.IndexByte(host[i:], ']') < 0 { + host = host[:i] + } + if host == "" { + return lb.Flow{} + } + return lb.Flow{Key: lb.HashFold(host), HasKey: true} +} + +// clientIPKey keys on the address resolved from trusted proxies when there is one, else on +// the peer's; the port is never part of it, as an ephemeral port would end all affinity +func clientIPKey(r *http.Request, v6Prefix int) lb.Flow { + if ip := tctx.ClientIP(r.Context()); ip != "" { + if addr, err := netip.ParseAddr(ip); err == nil { + return lb.Flow{Key: lb.HashAddr(addr, v6Prefix), HasKey: true} + } + return keyOf(ip) + } + if ap, err := netip.ParseAddrPort(r.RemoteAddr); err == nil { + return lb.Flow{Key: lb.HashAddr(ap.Addr(), v6Prefix), HasKey: true} + } + if addr, err := netip.ParseAddr(r.RemoteAddr); err == nil { + return lb.Flow{Key: lb.HashAddr(addr, v6Prefix), HasKey: true} + } + return keyOf(r.RemoteAddr) +} diff --git a/pkg/backends/alb/mech/pick/key_test.go b/pkg/backends/alb/mech/pick/key_test.go new file mode 100644 index 000000000..eb8f4febf --- /dev/null +++ b/pkg/backends/alb/mech/pick/key_test.go @@ -0,0 +1,187 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package pick + +import ( + "net/http" + "net/http/httptest" + "net/url" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/lb" + tctx "github.com/trickstercache/trickster/v2/pkg/proxy/context" +) + +func keyFuncFor(t *testing.T, source string) keyFunc { + t.Helper() + ks, err := options.ParseKeySource(source) + if err != nil { + t.Fatal(err) + } + return newKeyFunc(ks, options.DefaultIPv6Prefix) +} + +func request(target string, mutate func(*http.Request)) *http.Request { + r := httptest.NewRequest(http.MethodGet, target, nil) + if mutate != nil { + mutate(r) + } + return r +} + +func TestClientIPKey(t *testing.T) { + key := keyFuncFor(t, "client_ip") + from := func(remote, resolved string) lb.Flow { + r := request("http://example.com/", func(r *http.Request) { r.RemoteAddr = remote }) + if resolved != "" { + r = r.WithContext(tctx.WithClientIP(r.Context(), resolved)) + } + return key(r) + } + a := from("192.0.2.1:50000", "") + if !a.HasKey || a != from("192.0.2.1:61234", "") { + t.Error("the client's ephemeral port changed its key") + } + if a == from("192.0.2.2:50000", "") { + t.Error("two clients share a key") + } + // the address resolved from a trusted proxy wins over the proxy's own + if from("10.0.0.9:443", "192.0.2.1") != a { + t.Error("the resolved client address was not used") + } + if from("192.0.2.1", "") != a { + t.Error("a peer address without a port keyed differently") + } + // IPv6 clients key on their /64: privacy addresses rotate the rest + v6 := from("[2001:db8:1:2:aaaa:bbbb:cccc:dddd]:443", "") + if v6 != from("[2001:db8:1:2:1111:2222:3333:4444]:443", "") { + t.Error("two addresses in one /64 keyed differently") + } + if v6 == from("[2001:db8:1:3:aaaa:bbbb:cccc:dddd]:443", "") { + t.Error("two /64s share a key") + } + // something unparsable still keys, consistently; nothing at all does not + if odd := from("@", "not-an-ip"); !odd.HasKey || odd != from("@", "not-an-ip") { + t.Error("an unparsable resolved address did not key consistently") + } + if odd := from("pipe", ""); !odd.HasKey { + t.Error("an unparsable peer address did not key") + } + if from("", "").HasKey { + t.Error("no address at all produced a key") + } +} + +func TestHostKey(t *testing.T) { + key := keyFuncFor(t, "host") + of := func(host string) lb.Flow { + return key(request("http://example.com/", func(r *http.Request) { r.Host = host })) + } + if !of("api.example.com").HasKey || of("api.example.com") != of("API.Example.COM:8443") { + t.Error("case or port changed a host's key") + } + if of("api.example.com") == of("www.example.com") { + t.Error("two hosts share a key") + } + if of("[2001:db8::1]:8080") != of("[2001:db8::1]") { + t.Error("the port of an IPv6 literal changed its key") + } + if of("").HasKey { + t.Error("an empty host produced a key") + } +} + +func TestHeaderCookieAndQueryKeys(t *testing.T) { + header := keyFuncFor(t, "header:x-tenant") + with := func(v ...string) *http.Request { + return request("http://example.com/", func(r *http.Request) { r.Header["X-Tenant"] = v }) + } + if a := header(with("acme")); !a.HasKey || a != header(with("acme", "other")) || a == header(with("globex")) { + t.Error("the header key does not follow the first header value") + } + if header(with()).HasKey || header(with("")).HasKey { + t.Error("a missing or empty header produced a key") + } + + cookie := keyFuncFor(t, "cookie:session") + jar := func(lines ...string) *http.Request { + return request("http://example.com/", func(r *http.Request) { r.Header["Cookie"] = lines }) + } + want := lb.Flow{Key: lb.HashString("abc123"), HasKey: true} + for _, lines := range [][]string{ + {"session=abc123"}, {"theme=dark; session=abc123"}, {"theme=dark;session=abc123 ; x=1"}, + {"theme=dark", "session=abc123"}, {`session="abc123"`}, {"xsession=no; session=abc123"}, + } { + if got := cookie(jar(lines...)); got != want { + t.Errorf("cookies %q keyed %+v", lines, got) + } + } + for _, lines := range [][]string{nil, {"theme=dark"}, {"session="}, {"sessionid=abc123"}, {"session"}} { + if cookie(jar(lines...)).HasKey { + t.Errorf("cookies %q produced a key", lines) + } + } + + query := keyFuncFor(t, "query:tenant") + if got := query(request("http://example.com/q?a=1&tenant=acme&b=2", nil)); got != (lb.Flow{Key: lb.HashString("acme"), HasKey: true}) { + t.Errorf("query key = %+v", got) + } + for _, target := range []string{"http://example.com/q", "http://example.com/q?tenant=", "http://example.com/q?subtenant=acme"} { + if query(request(target, nil)).HasKey { + t.Errorf("%s produced a key", target) + } + } + if query(&http.Request{}).HasKey { + t.Error("a request with no URL produced a key") + } +} + +func TestKeyExtractionDoesNotAllocate(t *testing.T) { + r := request("http://example.com/q?a=1&tenant=acme", func(r *http.Request) { + r.RemoteAddr = "[2001:db8::7]:4431" + r.Host = "API.example.com:8443" + r.Header["X-Tenant"] = []string{"acme"} + r.Header["Cookie"] = []string{"theme=dark; session=abc123"} + }) + for _, source := range []string{"client_ip", "host", "header:X-Tenant", "cookie:session", "query:tenant"} { + key := keyFuncFor(t, source) + if allocs := testing.AllocsPerRun(200, func() { _ = key(r) }); allocs != 0 { + t.Errorf("%s allocates %v per request", source, allocs) + } + } +} + +func FuzzKeyExtraction(f *testing.F) { + f.Add("10.0.0.1:1", "example.com", "a=1; session=x", "tenant=acme&x") + f.Add("[::1]:80", "[::1]:8080", ";;==;", "&&==&") + f.Add("", "", "", "") + f.Fuzz(func(t *testing.T, remote, host, cookies, rawQuery string) { + r := &http.Request{RemoteAddr: remote, Host: host, URL: &url.URL{RawQuery: rawQuery}, Header: http.Header{ + "Cookie": {cookies}, "X-Tenant": {cookies}, + }} + for _, source := range []string{"client_ip", "host", "header:X-Tenant", "cookie:session", "query:tenant"} { + ks, err := options.ParseKeySource(source) + if err != nil { + t.Fatal(err) + } + key := newKeyFunc(ks, 64) + if a, b := key(r), key(r); a != b { + t.Errorf("%s keyed one request two ways", source) + } + } + }) +} diff --git a/pkg/backends/alb/mech/pick/pick.go b/pkg/backends/alb/mech/pick/pick.go new file mode 100644 index 000000000..b970c885e --- /dev/null +++ b/pkg/backends/alb/mech/pick/pick.go @@ -0,0 +1,156 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package pick is the HTTP adapter for selection strategies: one handler that dispatches each +// request to the pool member a strategy picks, whichever strategy that is. +package pick + +import ( + "net/http" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + cfgtypes "github.com/trickstercache/trickster/v2/pkg/config/types" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/proxy/handlers/trickster/failures" +) + +// Options are what the adapter needs beyond the strategy itself. +type Options struct { + // Key is where a request's affinity key is read from, for a strategy that needs one. + Key options.KeySource + // IPv6Prefix is how many leading bits of an IPv6 address form a client_ip key. + IPv6Prefix int + // GoodCodes are the response codes that count as a good answer, for a strategy that + // needs latency; nil counts every response as good. + GoodCodes *cfgtypes.StatusTable + // Balancer carries what the balancer itself is configured with, such as passive ejection. + Balancer lb.BalancerOptions +} + +type handler struct { + mech.PoolHolder + name types.Name + balancer *lb.Balancer + // resolved once: a strategy pays on dispatch only for what it needs + tracked bool + timed bool + key keyFunc + goodCodes *cfgtypes.StatusTable +} + +// New returns the pool mechanism that serves HTTP with the provided strategy, which the +// mechanism then owns. name is the mechanism's short name. +func New(name types.Name, selector lb.Selector, opts ...Options) types.PickerMechanism { + var o Options + if len(opts) > 0 { + o = opts[0] + } + b := lb.NewBalancer(selector, o.Balancer) + h := &handler{ + name: name, balancer: b, + tracked: b.Needs() != 0, timed: b.Needs().Has(lb.NeedLatency), + } + h.goodCodes = o.GoodCodes + if b.Needs().Has(lb.NeedKey) { + h.key = newKeyFunc(o.Key, o.IPv6Prefix) + } + return h +} + +func (h *handler) Name() types.Name { + return h.name +} + +// Picker returns the balancer behind the mechanism, for planes that do not dispatch over HTTP. +func (h *handler) Picker() lb.Picker { + return h.balancer +} + +// SetPool installs the pool that requests are dispatched to; nil leaves the mechanism with none. +func (h *handler) SetPool(p pool.Pool) { + h.PoolHolder.SetPool(p) + if p == nil { + h.balancer.SetPool(nil) + return + } + h.balancer.SetPool(p.Core()) +} + +func (h *handler) StopPool() { + if p := h.Pool(); p != nil { + p.Stop() + } +} + +// Balancer returns the mechanism's balancer, whose members carry its runtime stats. +func (h *handler) Balancer() *lb.Balancer { + return h.balancer +} + +func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + var flow lb.Flow + if h.key != nil && r != nil { + flow = h.key(r) + } + pk, ok := h.balancer.Pick(flow) + if !ok { + failures.HandleBadGateway(w, r) + return + } + t, ok := pk.Member().Value.(*pool.Target) + if !ok || t.Handler() == nil { + pk.Done(lb.OutcomeFailed) + failures.HandleBadGateway(w, r) + return + } + switch { + case !h.tracked: + t.Handler().ServeHTTP(w, r) + case h.timed: + h.serveTimed(pk, t.Handler(), w, r) + default: + h.serveTracked(pk, t.Handler(), w, r) + } +} + +// serveTracked reports the pick as done even when the member's handler panics, so the +// member's in-flight count cannot leak +func (h *handler) serveTracked(pk lb.Pick, member http.Handler, w http.ResponseWriter, r *http.Request) { + outcome := lb.OutcomeFailed + defer func() { pk.Done(outcome) }() + member.ServeHTTP(w, r) + outcome = lb.OutcomeOK +} + +// serveTimed also times the member's first write and judges its answer: a response code +// outside the good set is a failure, and a client that went away is nobody's +func (h *handler) serveTimed(pk lb.Pick, member http.Handler, w http.ResponseWriter, r *http.Request) { + fw := getWriter(w, pk) + outcome := lb.OutcomeFailed + defer func() { + pk.Done(outcome) + putWriter(fw) + }() + member.ServeHTTP(fw, r) + switch { + case r != nil && r.Context().Err() != nil: + outcome = lb.OutcomeCanceled + case h.goodCodes == nil || h.goodCodes.Contains(fw.code): + outcome = lb.OutcomeOK + } +} diff --git a/pkg/backends/alb/mech/pick/pick_characterization_test.go b/pkg/backends/alb/mech/pick/pick_characterization_test.go new file mode 100644 index 000000000..182eebec2 --- /dev/null +++ b/pkg/backends/alb/mech/pick/pick_characterization_test.go @@ -0,0 +1,167 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package pick + +import ( + "net/http" + "net/http/httptest" + "runtime" + "slices" + "sync" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" +) + +// sequenceRecorder collects the member index of each dispatch, in order +type sequenceRecorder struct { + seq []int +} + +func (s *sequenceRecorder) handler(i int) http.Handler { + return http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + s.seq = append(s.seq, i) + }) +} + +func recordedTargets(rec *sequenceRecorder, weights []int) pool.Targets { + targets := make(pool.Targets, len(weights)) + for i, w := range weights { + targets[i] = pool.NewWeightedTarget(rec.handler(i), passingStatus(), nil, w) + } + return targets +} + +func sum(weights []int) int { + var total int + for _, w := range weights { + total += w + } + return total +} + +func serveN(h http.Handler, n int) { + w := httptest.NewRecorder() + r := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + for range n { + h.ServeHTTP(w, r) + } +} + +// assertEveryWindowExact fails unless every run of total consecutive selections gives each +// member exactly its weight, wherever the run starts +func assertEveryWindowExact(t *testing.T, seq, weights []int) { + t.Helper() + total := sum(weights) + counts := make([]int, len(weights)) + for i, m := range seq { + counts[m]++ + if i >= total { + counts[seq[i-total]]-- + } + if i < total-1 { + continue + } + if !slices.Equal(counts, weights) { + t.Fatalf("window ending at selection %d apportioned %v, want %v", i, counts, weights) + } + } +} + +func TestRoundRobinEveryWindowIsExact(t *testing.T) { + for name, weights := range map[string][]int{ + "uniform": {1, 1, 1, 1}, + "weighted": {1, 3, 2}, + "heavy first": {7, 1, 1}, + "single": {4}, + "uniform of 2": {2, 2}, + } { + t.Run(name, func(t *testing.T) { + rec := &sequenceRecorder{} + p := pool.New(recordedTargets(rec, weights), 0) + defer p.Stop() + h := newRR() + h.SetPool(p) + serveN(h, 7*sum(weights)+3) + assertEveryWindowExact(t, rec.seq, weights) + }) + } +} + +// the order of a rotation, from a fixed start: a heavier member's turns are spread through +// it rather than taken back to back +func TestRoundRobinSequence(t *testing.T) { + weights := []int{1, 3, 2} + rec := &sequenceRecorder{} + p := pool.New(recordedTargets(rec, weights), 0) + defer p.Stop() + h := New(names.MechanismRR, rr.NewAt(0)) + h.SetPool(p) + serveN(h, 2*sum(weights)) + want := []int{1, 2, 1, 1, 2, 0, 1, 2, 1, 1, 2, 0} + if !slices.Equal(rec.seq, want) { + t.Errorf("weighted sequence = %v, want %v", rec.seq, want) + } + + rec = &sequenceRecorder{} + u := pool.New(recordedTargets(rec, []int{1, 1, 1}), 0) + defer u.Stop() + h = New(names.MechanismRR, rr.NewAt(0)) + h.SetPool(u) + serveN(h, 6) + want = []int{1, 2, 0, 1, 2, 0} + if !slices.Equal(rec.seq, want) { + t.Errorf("uniform sequence = %v, want %v", rec.seq, want) + } +} + +// a pool swapped for one of the same membership mid-rotation must not disturb apportionment: +// the rotation belongs to the mechanism, not to the pool +func TestRoundRobinApportionmentSurvivesConcurrentSetPool(t *testing.T) { + weights := []int{2, 1, 3} + rec := &sequenceRecorder{} + targets := recordedTargets(rec, weights) + first := pool.New(targets, 0) + h := newRR() + h.SetPool(first) + + pools := []pool.Pool{first} + stop := make(chan struct{}) + var wg sync.WaitGroup + wg.Go(func() { + for range 256 { + select { + case <-stop: + return + default: + } + p := pool.New(targets, 0) + pools = append(pools, p) + h.SetPool(p) + runtime.Gosched() + } + }) + serveN(h, 2000*sum(weights)) + close(stop) + wg.Wait() + for _, p := range pools { + p.Stop() + } + assertEveryWindowExact(t, rec.seq, weights) +} diff --git a/pkg/backends/alb/mech/rr/round_robin_setpool_race_test.go b/pkg/backends/alb/mech/pick/pick_setpool_race_test.go similarity index 94% rename from pkg/backends/alb/mech/rr/round_robin_setpool_race_test.go rename to pkg/backends/alb/mech/pick/pick_setpool_race_test.go index 13610ed80..a6cb306b1 100644 --- a/pkg/backends/alb/mech/rr/round_robin_setpool_race_test.go +++ b/pkg/backends/alb/mech/pick/pick_setpool_race_test.go @@ -14,7 +14,7 @@ * limitations under the License. */ -package rr +package pick import ( "net/http" @@ -27,14 +27,14 @@ import ( "github.com/trickstercache/trickster/v2/pkg/testutil/albpool" ) -// SetPool writes h.pool while ServeHTTP reads it via h.pool.LiveTargets(). +// SetPool writes h.pool while ServeHTTP reads it via h.pool.Targets(). // Interface field assignment is two-word and not atomic, so config reload // concurrent with in-flight requests should race. func TestRoundRobinSetPoolRace(t *testing.T) { const cycles = 100 initial, _, _ := albpool.New(-1, []http.Handler{http.HandlerFunc(tu.BasicHTTPHandler)}) - h := &handler{} + h := newRR() h.SetPool(initial) pools := make([]pool.Pool, 0, cycles+1) diff --git a/pkg/backends/alb/mech/rr/round_robin_test.go b/pkg/backends/alb/mech/pick/pick_test.go similarity index 91% rename from pkg/backends/alb/mech/rr/round_robin_test.go rename to pkg/backends/alb/mech/pick/pick_test.go index b004031d9..28db1eb34 100644 --- a/pkg/backends/alb/mech/rr/round_robin_test.go +++ b/pkg/backends/alb/mech/pick/pick_test.go @@ -14,14 +14,13 @@ * limitations under the License. */ -package rr +package pick import ( "net/http" "net/http/httptest" "testing" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" "github.com/trickstercache/trickster/v2/pkg/proxy/handlers/trickster/failures" tu "github.com/trickstercache/trickster/v2/pkg/testutil" "github.com/trickstercache/trickster/v2/pkg/testutil/albpool" @@ -29,7 +28,7 @@ import ( func TestHandleRoundRobin(t *testing.T) { w := httptest.NewRecorder() - h := &handler{} + h := newRR() h.ServeHTTP(w, nil) if w.Code != http.StatusBadGateway { t.Error("expected 502 got", w.Code) @@ -64,12 +63,11 @@ func TestHandleRoundRobin(t *testing.T) { } func TestNextTarget(t *testing.T) { - p := pool.New(nil, -1) - h := &handler{} + p, _, _ := albpool.NewHealthy([]http.Handler{http.NotFoundHandler()}) + h := newRR() h.SetPool(p) - h.StopPool() - p.SetHealthy([]http.Handler{http.NotFoundHandler()}) - n := h.nextTarget(p) + defer h.StopPool() + n := nextTarget(h) if n == nil { t.Error("expected non-nil target") } @@ -84,7 +82,7 @@ func TestRoundRobinProgression(t *testing.T) { defer p.Stop() albpool.WaitHealthy(t, p, 3) - rr := &handler{} + rr := newRR() rr.SetPool(p) // Fire 6 requests and verify rotation through all 3 backends. diff --git a/pkg/backends/alb/mech/pick/pick_timed_test.go b/pkg/backends/alb/mech/pick/pick_timed_test.go new file mode 100644 index 000000000..a2bb9bc34 --- /dev/null +++ b/pkg/backends/alb/mech/pick/pick_timed_test.go @@ -0,0 +1,252 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package pick + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/hrw" + "github.com/trickstercache/trickster/v2/pkg/lb/lt" + "github.com/trickstercache/trickster/v2/pkg/testutil/albpool" +) + +func newLT(t *testing.T, member http.Handler) (*handler, *pool.Target) { + t.Helper() + p, targets, _ := albpool.NewHealthy([]http.Handler{member}) + t.Cleanup(p.Stop) + h := New(names.MechanismLT, lt.New(lt.Options{}), + Options{GoodCodes: options.DefaultLTStatusCodes().Compile()}).(*handler) + h.SetPool(p) + return h, targets[0] +} + +func get(h http.Handler, r *http.Request) *httptest.ResponseRecorder { + w := httptest.NewRecorder() + if r == nil { + r = httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + } + h.ServeHTTP(w, r) + return w +} + +// the latency sample is the time to the member's first write, not to its last +func TestTimedDispatchSamplesFirstWrite(t *testing.T) { + h, tgt := newLT(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + time.Sleep(30 * time.Millisecond) + w.WriteHeader(http.StatusAccepted) + time.Sleep(120 * time.Millisecond) + _, _ = io.WriteString(w, "body") + })) + w := get(h, nil) + if w.Code != http.StatusAccepted || w.Body.String() != "body" { + t.Fatalf("response = %d %q", w.Code, w.Body.String()) + } + st := tgt.Member().Stats() + if got := st.Latency(); got < 30*time.Millisecond || got > 110*time.Millisecond { + t.Errorf("sample = %v, want the ~30ms to the first write, not the ~150ms to the last", got) + } + if st.Inflight() != 0 || st.Failures() != 0 { + t.Errorf("after the request: %d in flight, %d failures", st.Inflight(), st.Failures()) + } +} + +func TestTimedDispatchJudgesTheAnswer(t *testing.T) { + code := http.StatusOK + h, tgt := newLT(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(code) })) + st := tgt.Member().Stats() + // a gateway failure answered in microseconds is a penalty, not a fast sample + for i, bad := range []int{http.StatusBadGateway, http.StatusServiceUnavailable, http.StatusGatewayTimeout} { + code = bad + get(h, nil) + if st.Failures() != int32(i+1) || st.Latency() < lb.DefaultLatencyPenalty { + t.Fatalf("after a %d: %d failures, latency %v", bad, st.Failures(), st.Latency()) + } + } + // other codes, client and server errors included, are the member answering + for _, good := range []int{http.StatusOK, http.StatusNotFound, http.StatusInternalServerError} { + code = good + get(h, nil) + if st.Failures() != 0 { + t.Fatalf("a %d was counted as a failure", good) + } + } + // a body with no explicit header is a 200 + body, _ := newLT(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { _, _ = io.WriteString(w, "ok") })) + if w := get(body, nil); w.Code != http.StatusOK { + t.Errorf("implicit status = %d", w.Code) + } +} + +// a client that went away says nothing about the member, whatever was half-written +func TestTimedDispatchIgnoresCanceledRequests(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + h, tgt := newLT(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + cancel() + w.WriteHeader(http.StatusBadGateway) + })) + get(h, httptest.NewRequest(http.MethodGet, "http://example.com/", nil).WithContext(ctx)) + if st := tgt.Member().Stats(); st.Failures() != 0 || st.Latency() >= lb.DefaultLatencyPenalty || st.Inflight() != 0 { + t.Errorf("a canceled request moved the stats: %d failures, %v, %d in flight", + st.Failures(), st.Latency(), st.Inflight()) + } +} + +func TestTimedDispatchSurvivesAPanic(t *testing.T) { + h, tgt := newLT(t, http.HandlerFunc(func(http.ResponseWriter, *http.Request) { panic("member blew up") })) + func() { + defer func() { _ = recover() }() + get(h, nil) + }() + if st := tgt.Member().Stats(); st.Inflight() != 0 || st.Failures() != 1 { + t.Errorf("after a panic: %d in flight, %d failures", st.Inflight(), st.Failures()) + } + // with no good codes configured, every answer is a good one + p, targets, _ := albpool.NewHealthy([]http.Handler{http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusBadGateway) + })}) + defer p.Stop() + lenient := New(names.MechanismLT, lt.New(lt.Options{})) + lenient.SetPool(p) + get(lenient, nil) + if targets[0].Member().Stats().Failures() != 0 { + t.Error("a response was judged without any good codes to judge it by") + } +} + +// the wrapper must not hide what the real writer can do +func TestFirstWriteWriterPassesThrough(t *testing.T) { + rec := httptest.NewRecorder() + fw := getWriter(rec, lb.Pick{}) + defer putWriter(fw) + if fw.Unwrap() != rec { + t.Error("Unwrap does not return the wrapped writer") + } + n, err := fw.ReadFrom(strings.NewReader("streamed")) + if err != nil || n != 8 || rec.Body.String() != "streamed" || !fw.wrote || fw.code != http.StatusOK { + t.Errorf("ReadFrom: %d %v %q", n, err, rec.Body.String()) + } + fw.Flush() + if !rec.Flushed { + t.Error("Flush did not reach the wrapped writer") + } + if err := http.NewResponseController(fw).Flush(); err != nil { + t.Errorf("the response controller cannot see through the wrapper: %v", err) + } + // an informational response is not yet the member's answer + early := getWriter(httptest.NewRecorder(), lb.Pick{}) + defer putWriter(early) + early.WriteHeader(http.StatusEarlyHints) + if early.wrote { + t.Error("an informational response was taken for the first write") + } + early.WriteHeader(http.StatusSwitchingProtocols) + if !early.wrote || early.code != http.StatusSwitchingProtocols { + t.Error("a protocol switch was not taken for the answer") + } + // a writer with no ReadFrom or Flush of its own still works + var plain plainWriter + pw := getWriter(&plain, lb.Pick{}) + defer putWriter(pw) + if n, err := pw.ReadFrom(strings.NewReader("copied")); err != nil || n != 6 || plain.body.String() != "copied" { + t.Errorf("ReadFrom over a plain writer: %d %v %q", n, err, plain.body.String()) + } + pw.Flush() + pw.WriteHeader(http.StatusTeapot) + if pw.code != http.StatusOK { + t.Error("a header after the first write replaced the recorded status") + } +} + +// readerFromWriter has the copy fast path a real connection's writer has +type readerFromWriter struct { + plainWriter + fast bool +} + +func (w *readerFromWriter) ReadFrom(r io.Reader) (int64, error) { + w.fast = true + return io.Copy(&w.body, r) +} + +func TestFirstWriteWriterKeepsTheCopyFastPath(t *testing.T) { + var under readerFromWriter + fw := getWriter(&under, lb.Pick{}) + defer putWriter(fw) + // a limited reader has no WriteTo, so the copy turns to the writer's ReadFrom + if n, err := io.Copy(fw, io.LimitReader(strings.NewReader("sendfile"), 8)); err != nil || n != 8 || !under.fast { + t.Errorf("copy: %d %v, fast path taken: %v", n, err, under.fast) + } +} + +type plainWriter struct{ body strings.Builder } + +func (*plainWriter) Header() http.Header { return http.Header{} } + +func (w *plainWriter) Write(b []byte) (int, error) { return w.body.Write(b) } + +func (*plainWriter) WriteHeader(int) {} + +// requests with one key reach one member; the mechanism reads the key it was configured for +func TestKeyedDispatchIsSticky(t *testing.T) { + hs := make([]http.Handler, 5) + for i := range hs { + hs[i] = albpool.NamedHandler(string(rune('a' + i))) + } + p, _, _ := albpool.NewHealthy(hs) + defer p.Stop() + ks, err := options.ParseKeySource("header:X-Tenant") + if err != nil { + t.Fatal(err) + } + h := New(names.MechanismHRW, hrw.New(), Options{Key: ks, IPv6Prefix: 64}) + h.SetPool(p) + owners := make(map[string]string) + reached := make(map[string]bool) + for round := range 3 { + for _, tenant := range []string{"acme", "globex", "initech", "umbrella", "hooli", "stark", "wayne", "wonka"} { + r := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + r.Header.Set("X-Tenant", tenant) + got := get(h, r).Body.String() + if round > 0 && owners[tenant] != got { + t.Fatalf("%s moved from %s to %s", tenant, owners[tenant], got) + } + owners[tenant] = got + reached[got] = true + } + } + if len(reached) < 3 { + t.Errorf("8 tenants reached only %d of 5 members", len(reached)) + } + // a request with no key is still served + if w := get(h, nil); w.Code != http.StatusOK { + t.Errorf("keyless request = %d", w.Code) + } + w := httptest.NewRecorder() + h.ServeHTTP(w, nil) + if w.Code != http.StatusOK { + t.Errorf("nil request = %d", w.Code) + } +} diff --git a/pkg/backends/alb/mech/pick/pick_tracked_test.go b/pkg/backends/alb/mech/pick/pick_tracked_test.go new file mode 100644 index 000000000..222f803ce --- /dev/null +++ b/pkg/backends/alb/mech/pick/pick_tracked_test.go @@ -0,0 +1,138 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package pick + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/testutil/albpool" +) + +// firstMember is a strategy with a need, so dispatch through it is accounted for +type firstMember struct{ members []*lb.Member } + +func (*firstMember) Name() string { return "first" } + +func (*firstMember) Needs() lb.Needs { return lb.NeedInflight } + +func (*firstMember) Prepare(s *lb.Snapshot) lb.Prepared { return &firstMember{members: s.Members} } + +func (f *firstMember) Select(lb.Flow) *lb.Member { return f.members[0] } + +var _ types.PickerMechanism = (*handler)(nil) + +func TestPickerAndName(t *testing.T) { + h := New("first", &firstMember{}).(*handler) + if h.Name() != "first" { + t.Errorf("name = %q", h.Name()) + } + if h.Balancer() == nil || h.Picker() != lb.Picker(h.Balancer()) { + t.Error("the picker is not the mechanism's balancer") + } + if h.Picker() == nil || h.Picker().Needs() != lb.NeedInflight { + t.Error("the mechanism does not expose its balancer") + } + if _, ok := h.Picker().Pick(lb.Flow{}); ok { + t.Error("picked before a pool was installed") + } + p, _, _ := albpool.NewHealthy([]http.Handler{http.NotFoundHandler()}) + defer p.Stop() + h.SetPool(p) + if h.Pool() != p { + t.Error("the pool was not held") + } + pk, ok := h.Picker().Pick(lb.Flow{}) + if !ok { + t.Fatal("no pick from an installed pool") + } + pk.Done(lb.OutcomeOK) + // removing the pool leaves nothing to pick from, and a 502 to serve + h.SetPool(nil) + if _, ok := h.Picker().Pick(lb.Flow{}); ok { + t.Error("picked after the pool was removed") + } + w := httptest.NewRecorder() + h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "http://example.com/", nil)) + if w.Code != http.StatusBadGateway { + t.Errorf("expected 502 with no pool, got %d", w.Code) + } + h.StopPool() +} + +// a strategy that tracks in-flight work sees the request while it runs and not after, even +// when the member's handler panics +func TestTrackedDispatchBalancesInflight(t *testing.T) { + var during int64 + var tgt *pool.Target + serve := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + during = tgt.Member().Stats().Inflight() + if r.URL.Path == "/panic" { + panic("member blew up") + } + w.WriteHeader(http.StatusNoContent) + }) + p, targets, _ := albpool.NewHealthy([]http.Handler{serve}) + defer p.Stop() + tgt = targets[0] + h := New("first", &firstMember{}) + h.SetPool(p) + + w := httptest.NewRecorder() + h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "http://example.com/", nil)) + if w.Code != http.StatusNoContent || during != 1 { + t.Errorf("code %d, in-flight during the request = %d", w.Code, during) + } + if got := tgt.Member().Stats().Inflight(); got != 0 { + t.Errorf("in-flight after the request = %d", got) + } + + func() { + defer func() { + if recover() == nil { + t.Error("the member's panic was swallowed") + } + }() + h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://example.com/panic", nil)) + }() + if got := tgt.Member().Stats().Inflight(); got != 0 { + t.Errorf("a panicking member leaked %d in flight", got) + } +} + +// a member whose payload is not a dispatchable target is a 502, with its pick accounted for +func TestUndispatchableMemberIsBadGateway(t *testing.T) { + stray := lb.NewMember(lb.MemberOptions{Name: "stray", Value: "not a target"}) + core, err := lb.NewPool([]*lb.Member{stray}, 0) + if err != nil { + t.Fatal(err) + } + defer core.Stop() + h := New("first", &firstMember{}).(*handler) + h.balancer.SetPool(core) + w := httptest.NewRecorder() + h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "http://example.com/", nil)) + if w.Code != http.StatusBadGateway { + t.Errorf("expected 502, got %d", w.Code) + } + if got := stray.Stats().Inflight(); got != 0 { + t.Errorf("the refused pick leaked %d in flight", got) + } +} diff --git a/pkg/backends/alb/mech/rr/round_robin_unavail_test.go b/pkg/backends/alb/mech/pick/pick_unavail_test.go similarity index 70% rename from pkg/backends/alb/mech/rr/round_robin_unavail_test.go rename to pkg/backends/alb/mech/pick/pick_unavail_test.go index 45fb8076a..f1a7f198f 100644 --- a/pkg/backends/alb/mech/rr/round_robin_unavail_test.go +++ b/pkg/backends/alb/mech/pick/pick_unavail_test.go @@ -14,7 +14,7 @@ * limitations under the License. */ -package rr +package pick import ( "net/http" @@ -26,31 +26,25 @@ import ( "github.com/trickstercache/trickster/v2/pkg/testutil/albpool" ) -// nextTarget must skip targets whose hcStatus dropped below the pool's -// healthyFloor since the snapshot was taken. Without the dispatch-time check, -// rr would route to a member the pool already considers unavailable. -func TestNextTargetSkipsStaleFailingTarget(t *testing.T) { +// rr must not route to a member whose status dropped below the pool's healthyFloor: the very +// next request after the transition already excludes it, with no wait for a refresh. +func TestNextTargetSkipsFailingTargetImmediately(t *testing.T) { var hits1, hits2 atomic.Int64 h1 := http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { hits1.Add(1) }) h2 := http.HandlerFunc(func(_ http.ResponseWriter, _ *http.Request) { hits2.Add(1) }) p, _, sts := albpool.New(1, []http.Handler{h1, h2}) + defer p.Stop() sts[0].Set(healthcheck.StatusPassing) sts[1].Set(healthcheck.StatusPassing) + if got := len(p.Targets()); got != 2 { + t.Fatalf("setup: expected 2 healthy targets, got %d", got) + } - albpool.WaitHealthy(t, p, 2) - - // Stop the pool so no auto-refresh can repair a stale snapshot. This - // pins the test to the exact race window the dispatch-time re-check - // is meant to close. - p.Stop() - - // Flip target 1 to Failing. The snapshot is now permanently stale until - // the dispatch-time re-check kicks in. sts[1].Set(healthcheck.StatusFailing) - rr := &handler{} + rr := newRR() rr.SetPool(p) const reqs = 50 for range reqs { diff --git a/pkg/backends/alb/mech/rr/round_robin_weight_test.go b/pkg/backends/alb/mech/pick/pick_weight_test.go similarity index 92% rename from pkg/backends/alb/mech/rr/round_robin_weight_test.go rename to pkg/backends/alb/mech/pick/pick_weight_test.go index 578864ea5..921d189be 100644 --- a/pkg/backends/alb/mech/rr/round_robin_weight_test.go +++ b/pkg/backends/alb/mech/pick/pick_weight_test.go @@ -14,7 +14,7 @@ * limitations under the License. */ -package rr +package pick import ( "net/http" @@ -53,7 +53,7 @@ func TestWeightedRoundRobinExactApportionment(t *testing.T) { } p := pool.New(targets, 0) defer p.Stop() - h := &handler{} + h := newRR() h.SetPool(p) const cycles = 5 @@ -85,7 +85,7 @@ func TestUniformWeightsUseRotation(t *testing.T) { } p := pool.New(targets, 0) defer p.Stop() - h := &handler{} + h := newRR() h.SetPool(p) w := httptest.NewRecorder() r := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) @@ -100,7 +100,7 @@ func TestUniformWeightsUseRotation(t *testing.T) { } func TestServeHTTPNilAndEmptyPool(t *testing.T) { - h := &handler{} + h := newRR() w := httptest.NewRecorder() r := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) // no pool installed: 502 @@ -118,7 +118,7 @@ func TestServeHTTPNilAndEmptyPool(t *testing.T) { } // StopPool is safe with and without a pool h.StopPool() - h2 := &handler{} + h2 := newRR() h2.StopPool() if h.Name() != "rr" { t.Errorf("unexpected mechanism name %s", h.Name()) @@ -138,20 +138,19 @@ func TestNextTargetZeroAlloc(t *testing.T) { targets[i] = pool.NewWeightedTarget(&countingHandler{}, passingStatus(), nil, w) } p := pool.New(targets, 0) - h := &handler{} + h := newRR() h.SetPool(p) - // the pool's async refresh worker must drain its pending flag - // before Targets() serves the cached zero-alloc fast path; wait + // the first read of a snapshot builds the cached target view; wait // for that steady state, then hold it to the bar deadline := time.Now().Add(2 * time.Second) - for testing.AllocsPerRun(1, func() { h.nextTarget(p) }) != 0 { + for testing.AllocsPerRun(1, func() { nextTarget(h) }) != 0 { if time.Now().After(deadline) { break } time.Sleep(time.Millisecond) } if allocs := testing.AllocsPerRun(1000, func() { - if h.nextTarget(p) == nil { + if nextTarget(h) == nil { t.Fatal("expected a target") } }); allocs != 0 { diff --git a/pkg/backends/alb/mech/pick/writer.go b/pkg/backends/alb/mech/pick/writer.go new file mode 100644 index 000000000..ee8ef1a7c --- /dev/null +++ b/pkg/backends/alb/mech/pick/writer.go @@ -0,0 +1,88 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package pick + +import ( + "io" + "net/http" + "sync" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// firstWriteWriter reports the first byte written to the client, which is the latency signal +// of an HTTP flow, and keeps the status code so the flow's outcome can be judged +type firstWriteWriter struct { + http.ResponseWriter + pick lb.Pick + code int + wrote bool +} + +var writerPool = sync.Pool{New: func() any { return &firstWriteWriter{} }} + +func getWriter(w http.ResponseWriter, pk lb.Pick) *firstWriteWriter { + fw := writerPool.Get().(*firstWriteWriter) + fw.ResponseWriter, fw.pick, fw.code, fw.wrote = w, pk, http.StatusOK, false + return fw +} + +func putWriter(fw *firstWriteWriter) { + fw.ResponseWriter, fw.pick = nil, lb.Pick{} + writerPool.Put(fw) +} + +func (w *firstWriteWriter) first(code int) { + if w.wrote { + return + } + w.wrote, w.code = true, code + w.pick.FirstByte() +} + +func (w *firstWriteWriter) WriteHeader(code int) { + // an informational response is not yet the member's answer + if code < 100 || code >= 200 || code == http.StatusSwitchingProtocols { + w.first(code) + } + w.ResponseWriter.WriteHeader(code) +} + +func (w *firstWriteWriter) Write(b []byte) (int, error) { + w.first(http.StatusOK) + return w.ResponseWriter.Write(b) +} + +// ReadFrom keeps the underlying writer's copy fast path, such as sendfile, reachable +func (w *firstWriteWriter) ReadFrom(r io.Reader) (int64, error) { + w.first(http.StatusOK) + if rf, ok := w.ResponseWriter.(io.ReaderFrom); ok { + return rf.ReadFrom(r) + } + return io.Copy(writerOnly{w.ResponseWriter}, r) +} + +// writerOnly hides ReadFrom so io.Copy does not recurse into it +type writerOnly struct{ io.Writer } + +func (w *firstWriteWriter) Flush() { + if f, ok := w.ResponseWriter.(http.Flusher); ok { + f.Flush() + } +} + +// Unwrap exposes the underlying writer to http.ResponseController. +func (w *firstWriteWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter } diff --git a/pkg/backends/alb/mech/registry/registry.go b/pkg/backends/alb/mech/registry/registry.go index e70b1297f..548b58d18 100644 --- a/pkg/backends/alb/mech/registry/registry.go +++ b/pkg/backends/alb/mech/registry/registry.go @@ -17,58 +17,210 @@ package registry import ( + "slices" + "time" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/errors" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/fr" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/nlm" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/rr" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/pick" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/spread" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/tsm" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/ur" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/observe" "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" rt "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/hrw" + "github.com/trickstercache/trickster/v2/pkg/lb/lc" + "github.com/trickstercache/trickster/v2/pkg/lb/lt" + "github.com/trickstercache/trickster/v2/pkg/lb/p2c" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" ) // this slice is the one and only place to aggregate all registered Mechanisms var registry = []types.RegistryEntry{ - rr.RegistryEntry(), + roundRobin(), + strategy(names.MechanismPowerOfTwoChoices, names.MechanismP2C, everyPlane, + func(*options.Options) (lb.Selector, error) { return p2c.New(), nil }), + strategy(names.MechanismHighestRandomWeight, names.MechanismHRW, everyPlane, + func(*options.Options) (lb.Selector, error) { return hrw.New(), nil }), + // a native session reports no latency for least time to rank members by + strategy(names.MechanismLeastTime, names.MechanismLT, types.PlaneHTTP|types.PlaneStream, + func(o *options.Options) (lb.Selector, error) { + if o == nil { + return lt.New(lt.Options{}), nil + } + return lt.New(lt.Options{Decay: o.LTDecay()}), nil + }), + strategy(names.MechanismLeastConnections, names.MechanismLC, everyPlane, + func(*options.Options) (lb.Selector, error) { return lc.New(), nil }), fr.RegistryEntry(), fr.RegistryEntryFGR(), nlm.RegistryEntry(), tsm.RegistryEntry(), ur.RegistryEntry(), + spread.RegistryEntryRace(), + spread.RegistryEntryMirror(), +} + +// roundRobin is a selection strategy, so it serves every plane that commits one unit of work +// to one member +func roundRobin() types.RegistryEntry { + return types.RegistryEntry{ + Name: names.MechanismRoundRobin, + ShortName: names.MechanismRR, + Planes: everyPlane, + NewSelector: func(*options.Options) (lb.Selector, error) { + return rr.New(), nil + }, + } +} + +// everyPlane is where one unit of work is committed to one member: requests, tcp, tls and udp +// flows, and the sessions of a native protocol listener +const everyPlane = types.PlaneHTTP | types.PlaneStream | types.PlaneNative + +// strategy registers a selection strategy for the planes it can serve +func strategy(name, shortName types.Name, planes types.Plane, fn types.NewSelectorFunc) types.RegistryEntry { + return types.RegistryEntry{Name: name, ShortName: shortName, Planes: planes, NewSelector: fn} } var registryByName = compileSupportedByName(registry) // compileSupportedByName indexes registry entries by both Name and ShortName. // Panics on duplicate registration: silently last-write-wins would mask a -// configuration error that only surfaces at request time. -func compileSupportedByName(entries []types.RegistryEntry) map[types.Name]types.NewMechanismFunc { - out := make(map[types.Name]types.NewMechanismFunc, len(entries)*2) - add := func(name types.Name, fn types.NewMechanismFunc) { +// configuration error that only surfaces at request time. It also panics on an +// entry that is not exactly one of a mechanism and a selection strategy. +func compileSupportedByName(entries []types.RegistryEntry) map[types.Name]types.RegistryEntry { + out := make(map[types.Name]types.RegistryEntry, len(entries)*2) + add := func(name types.Name, entry types.RegistryEntry) { if _, exists := out[name]; exists { panic("alb/mech/registry: duplicate mechanism name " + name) } - out[name] = fn + out[name] = entry } for _, entry := range entries { - add(entry.ShortName, entry.New) - add(entry.Name, entry.New) + if (entry.New == nil) == (entry.NewSelector == nil) { + panic("alb/mech/registry: mechanism " + entry.Name + " must set exactly one of New and NewSelector") + } + add(entry.ShortName, entry) + add(entry.Name, entry) } return out } +// New returns the named mechanism as an HTTP handler. A selection strategy is wrapped in +// the handler that dispatches to the member it picks. func New(name types.Name, opts *options.Options, factories rt.Lookup, ) (types.Mechanism, error) { - if f, ok := registryByName[name]; ok && f != nil { - return f(opts, factories) + entry, ok := registryByName[name] + if !ok { + return nil, errors.ErrUnsupportedMechanism + } + if entry.NewSelector == nil { + return entry.New(opts, factories) + } + s, err := entry.NewSelector(opts) + if err != nil { + return nil, err + } + return pick.New(entry.ShortName, s, pickOptions(opts)), nil +} + +// pickOptions carries to the HTTP adapter what a strategy's needs may call for +func pickOptions(o *options.Options) pick.Options { + if o == nil { + return pick.Options{} + } + return pick.Options{ + Key: o.HRW.KeySource, IPv6Prefix: o.HRW.IPv6Prefix, GoodCodes: o.LT.GoodCodes, + Balancer: balancerOptions(o), + } +} + +// balancerOptions translates what configures the balancer itself: passive ejection, which +// only a stream listener's connect failures ever feed +func balancerOptions(o *options.Options) lb.BalancerOptions { + bo := lb.BalancerOptions{Observer: observe.Balancer(o.Name)} + if o.Stream != nil && o.Stream.PassiveHealth != nil { + p := o.Stream.PassiveHealth + bo.Ejection = lb.EjectionOptions{ + Failures: p.Failures, Duration: time.Duration(p.Eject), MaxPercent: p.MaxEjectedPercent, + } + } + return bo +} + +// NewBalancer returns the named selection strategy as a balancer with no pool, for a plane +// that does not dispatch over HTTP. A mechanism that is not a strategy is unsupported. +func NewBalancer(name types.Name, opts *options.Options) (*lb.Balancer, error) { + entry, ok := registryByName[name] + if !ok || entry.NewSelector == nil { + return nil, errors.ErrUnsupportedMechanism + } + s, err := entry.NewSelector(opts) + if err != nil { + return nil, err + } + if opts == nil { + return lb.NewBalancer(s), nil } - return nil, errors.ErrUnsupportedMechanism + return lb.NewBalancer(s, balancerOptions(opts)), nil } func IsRegistered(name types.Name) bool { _, ok := registryByName[name] return ok } + +// Supports reports whether the named mechanism can serve every plane in plane. +func Supports(name types.Name, plane types.Plane) bool { + entry, ok := registryByName[name] + return ok && entry.Planes.Has(plane) +} + +// ServesProtocol reports whether the named mechanism can serve a stream listener of the given +// protocol: tcp, tls or udp. +func ServesProtocol(name types.Name, protocol string) bool { + entry, ok := registryByName[name] + if !ok || !entry.Planes.Has(types.PlaneStream) { + return false + } + return len(entry.Protocols) == 0 || slices.Contains(entry.Protocols, protocol) +} + +// ServingProtocol returns the short names of the mechanisms that can serve a stream listener +// of the given protocol, sorted. +func ServingProtocol(protocol string) []types.Name { + var out []types.Name + for _, entry := range registry { + if ServesProtocol(entry.ShortName, protocol) { + out = append(out, entry.ShortName) + } + } + slices.Sort(out) + return out +} + +// StreamProtocols returns the stream protocols the named mechanism is limited to, or nil +// when it serves them all or none. +func StreamProtocols(name types.Name) []string { + return slices.Clone(registryByName[name].Protocols) +} + +// Supporting returns the short names of the mechanisms that can serve plane, sorted. +func Supporting(plane types.Plane) []types.Name { + var out []types.Name + for _, entry := range registry { + if entry.Planes.Has(plane) { + out = append(out, entry.ShortName) + } + } + slices.Sort(out) + return out +} diff --git a/pkg/backends/alb/mech/registry/registry_planes_test.go b/pkg/backends/alb/mech/registry/registry_planes_test.go new file mode 100644 index 000000000..9aaebb541 --- /dev/null +++ b/pkg/backends/alb/mech/registry/registry_planes_test.go @@ -0,0 +1,178 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package registry + +import ( + "errors" + "testing" + + alberr "github.com/trickstercache/trickster/v2/pkg/backends/alb/errors" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + rt "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" + + "github.com/stretchr/testify/require" +) + +// what each mechanism may serve; a change here is a change to what configs validate +var wantPlanes = map[types.Name]types.Plane{ + names.MechanismRR: everyPlane, + names.MechanismP2C: everyPlane, + names.MechanismHRW: everyPlane, + // a native session reports no latency + names.MechanismLT: types.PlaneHTTP | types.PlaneStream, + names.MechanismLC: everyPlane, + names.MechanismFR: types.PlaneHTTP, + names.MechanismFGR: types.PlaneHTTP, + names.MechanismNLM: types.PlaneHTTP, + names.MechanismTSM: types.PlaneHTTP, + names.MechanismUR: types.PlaneHTTP | types.PlaneNative, + // what commits a flow to several members at once is the stream relay's alone to carry out + names.MechanismRace: types.PlaneStream, + names.MechanismMirror: types.PlaneStream, +} + +// the strategies whose weights are an exact apportionment contract +var exactWeights = map[types.Name]bool{names.MechanismRR: true} + +func TestEntriesDeclareOneConstructorAndTheirPlanes(t *testing.T) { + require.Len(t, registry, len(wantPlanes)) + for _, e := range registry { + require.NotEqual(t, e.New == nil, e.NewSelector == nil, + "%s must set exactly one of New and NewSelector", e.Name) + want, ok := wantPlanes[e.ShortName] + require.True(t, ok, "%s has no expected planes", e.ShortName) + require.Equal(t, want, e.Planes, e.ShortName) + // a mechanism serves requests, or is limited to the stream protocols it names; only a + // strategy serves both, since no other has a member to hand to a plane without a handler + require.NotEqual(t, e.Planes.Has(types.PlaneHTTP), len(e.Protocols) > 0, e.ShortName) + if e.NewSelector == nil { + require.Equal(t, len(e.Protocols) > 0, e.Planes.Has(types.PlaneStream), e.ShortName) + } + } +} + +func TestCompilePanicsWithoutExactlyOneConstructor(t *testing.T) { + mech := func(*options.Options, rt.Lookup) (types.Mechanism, error) { return nil, nil } + sel := func(*options.Options) (lb.Selector, error) { return nil, nil } + for name, e := range map[string]types.RegistryEntry{ + "neither": {Name: "long", ShortName: "short"}, + "both": {Name: "long", ShortName: "short", New: mech, NewSelector: sel}, + } { + t.Run(name, func(t *testing.T) { + require.Panics(t, func() { compileSupportedByName([]types.RegistryEntry{e}) }) + }) + } +} + +func TestSupports(t *testing.T) { + for _, name := range []types.Name{names.MechanismRR, names.MechanismRoundRobin} { + require.True(t, Supports(name, types.PlaneHTTP), name) + require.True(t, Supports(name, types.PlaneStream), name) + require.True(t, Supports(name, types.PlaneHTTP|types.PlaneStream), name) + require.True(t, Supports(name, types.PlaneNative), name) + } + require.False(t, Supports(names.MechanismLT, types.PlaneNative)) + require.False(t, Supports(names.MechanismFR, types.PlaneStream)) + require.False(t, Supports(names.MechanismTSM, types.PlaneStream)) + require.True(t, Supports(names.MechanismUR, types.PlaneNative)) + require.False(t, Supports(names.MechanismUR, types.PlaneStream)) + require.False(t, Supports("nonexistent", types.PlaneHTTP)) + require.False(t, Supports(names.MechanismRR, 0), "no plane is not a supported plane") + + require.Equal(t, []types.Name{names.MechanismHRW, names.MechanismLC, names.MechanismLT, + names.MechanismMirror, names.MechanismP2C, names.MechanismRace, names.MechanismRR}, + Supporting(types.PlaneStream)) + require.Equal(t, []types.Name{names.MechanismHRW, names.MechanismLC, names.MechanismLT, + names.MechanismP2C, names.MechanismRace, names.MechanismRR}, ServingProtocol("tcp")) + require.Equal(t, []types.Name{names.MechanismHRW, names.MechanismLC, names.MechanismLT, + names.MechanismMirror, names.MechanismP2C, names.MechanismRR}, ServingProtocol("udp")) + require.True(t, ServesProtocol(names.MechanismConnectRace, "tls")) + require.False(t, ServesProtocol(names.MechanismRace, "udp")) + require.False(t, ServesProtocol(names.MechanismFR, "tcp")) + require.False(t, ServesProtocol("nonexistent", "tcp")) + require.Equal(t, []string{"udp"}, StreamProtocols(names.MechanismUDPMirror)) + require.Empty(t, StreamProtocols(names.MechanismRR)) + require.Empty(t, StreamProtocols("nonexistent")) + require.False(t, Supports(names.MechanismRace, types.PlaneHTTP)) + require.Equal(t, []types.Name{names.MechanismHRW, names.MechanismLC, names.MechanismP2C, + names.MechanismRR, names.MechanismUR}, Supporting(types.PlaneNative)) + require.Len(t, Supporting(types.PlaneHTTP), len(registry)-2) +} + +func TestNewWrapsAStrategyForHTTP(t *testing.T) { + for _, name := range []types.Name{names.MechanismRR, names.MechanismRoundRobin} { + m, err := New(name, &options.Options{}, nil) + require.NoError(t, err) + require.Equal(t, names.MechanismRR, m.Name()) + pm, ok := m.(types.PickerMechanism) + require.True(t, ok, "a strategy is served as a picker mechanism") + require.NotNil(t, pm.Picker()) + } + // each mechanism owns its strategy: two ALBs never share a rotation + a, _ := New(names.MechanismRR, nil, nil) + b, _ := New(names.MechanismRR, nil, nil) + require.NotSame(t, a.(types.PickerMechanism).Picker(), b.(types.PickerMechanism).Picker()) +} + +func TestNewBalancer(t *testing.T) { + b, err := NewBalancer(names.MechanismRR, nil) + require.NoError(t, err) + require.NotNil(t, b) + require.Nil(t, b.Pool()) + for _, name := range []types.Name{names.MechanismFR, names.MechanismUR, "nonexistent"} { + _, err := NewBalancer(name, nil) + require.ErrorIs(t, err, alberr.ErrUnsupportedMechanism, name) + } +} + +func TestSelectorConstructorErrorsSurface(t *testing.T) { + errBad := errors.New("bad strategy options") + saved := registryByName + t.Cleanup(func() { registryByName = saved }) + registryByName = compileSupportedByName([]types.RegistryEntry{{ + Name: "broken_long", ShortName: "broken", Planes: types.PlaneHTTP, + NewSelector: func(*options.Options) (lb.Selector, error) { return nil, errBad }, + }}) + _, err := New("broken", nil, nil) + require.ErrorIs(t, err, errBad) + _, err = NewBalancer("broken", nil) + require.ErrorIs(t, err, errBad) +} + +// every registered strategy is held to the core's selector contract +func TestRegisteredSelectorsConform(t *testing.T) { + var strategies int + for _, e := range registry { + if e.NewSelector == nil { + continue + } + strategies++ + t.Run(e.ShortName, func(t *testing.T) { + lbtest.Run(t, func() lb.Selector { + s, err := e.NewSelector(&options.Options{}) + if err != nil { + t.Fatal(err) + } + return s + }, lbtest.Options{ExactWeights: exactWeights[e.ShortName]}) + }) + } + require.NotZero(t, strategies) +} diff --git a/pkg/backends/alb/mech/registry/registry_test.go b/pkg/backends/alb/mech/registry/registry_test.go index 6d590e6b0..03148a9a8 100644 --- a/pkg/backends/alb/mech/registry/registry_test.go +++ b/pkg/backends/alb/mech/registry/registry_test.go @@ -88,16 +88,3 @@ func TestNewRoundRobinNilOptions(t *testing.T) { require.NoError(t, err) require.Equal(t, names.MechanismRR, m.Name()) } - -func TestNewNilConstructor(t *testing.T) { - original := registryByName - t.Cleanup(func() { registryByName = original }) - registryByName = compileSupportedByName([]types.RegistryEntry{ - {Name: "nil_constructor", ShortName: "nil"}, - }) - for _, name := range []string{"nil_constructor", "nil"} { - m, err := New(name, nil, nil) - require.ErrorIs(t, err, alberr.ErrUnsupportedMechanism) - require.Nil(t, m) - } -} diff --git a/pkg/backends/alb/mech/rr/round_robin.go b/pkg/backends/alb/mech/rr/round_robin.go deleted file mode 100644 index 545f2f407..000000000 --- a/pkg/backends/alb/mech/rr/round_robin.go +++ /dev/null @@ -1,101 +0,0 @@ -/* - * Copyright 2018 The Trickster Authors - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -package rr - -import ( - "net/http" - "sync/atomic" - - "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" - rt "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" - "github.com/trickstercache/trickster/v2/pkg/proxy/handlers/trickster/failures" -) - -const Name types.Name = "round_robin" - -type handler struct { - mech.PoolHolder - pos atomic.Uint64 -} - -func RegistryEntry() types.RegistryEntry { - return types.RegistryEntry{Name: Name, ShortName: names.MechanismRR, New: New} -} - -func New(_ *options.Options, _ rt.Lookup) (types.Mechanism, error) { - return &handler{}, nil -} - -func (h *handler) Name() types.Name { - return names.MechanismRR -} - -func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { - p := h.Pool() - if p == nil { - failures.HandleBadGateway(w, r) - return - } - if t := h.nextTarget(p); t != nil { - t.ServeHTTP(w, r) - return - } - failures.HandleBadGateway(w, r) -} - -func (h *handler) StopPool() { - if p := h.Pool(); p != nil { - p.Stop() - } -} - -// nextTarget selects the next pool member. With uniform weights it is a -// plain modular rotation; with mixed weights, each member owns a contiguous -// weight-sized span of the [0, totalWeight) rotation, so apportionment over -// any totalWeight consecutive selections against a stable healthy set is -// exact: each member is selected exactly Weight() times. -func (h *handler) nextTarget(p pool.Pool) http.Handler { - targets := p.Targets() - n := uint64(len(targets)) - if n == 0 { - return nil - } - var total uint64 - weighted := false - for _, t := range targets { - w := uint64(t.Weight()) //nolint:gosec // Target.Weight() is always >= 1 - if w != 1 { - weighted = true - } - total += w - } - if !weighted { - return targets[h.pos.Add(1)%n].Handler() - } - k := int(h.pos.Add(1) % total) //nolint:gosec // value is < total, an int sum - for _, t := range targets { - k -= t.Weight() - if k < 0 { - return t.Handler() - } - } - return targets[n-1].Handler() -} diff --git a/pkg/backends/alb/mech/spread/spread.go b/pkg/backends/alb/mech/spread/spread.go new file mode 100644 index 000000000..70ded436b --- /dev/null +++ b/pkg/backends/alb/mech/spread/spread.go @@ -0,0 +1,85 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +// Package spread holds the mechanisms that commit one flow to several pool members at once: +// a connect race for tcp and tls, and datagram mirroring for udp. Only a stream listener can +// serve them; they hold the pool, and the stream relay does the rest. +package spread + +import ( + "net/http" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + rt "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" + "github.com/trickstercache/trickster/v2/pkg/proxy/handlers/trickster/failures" +) + +// the stream protocols each mechanism serves, as a listener spells them +const ( + protocolTCP = "tcp" + protocolTLS = "tls" + protocolUDP = "udp" +) + +type handler struct { + mech.PoolHolder + name types.Name + spread types.Spread +} + +// RegistryEntryRace returns the registry entry of the connect race mechanism. +func RegistryEntryRace() types.RegistryEntry { + return types.RegistryEntry{ + Name: names.MechanismConnectRace, ShortName: names.MechanismRace, Planes: types.PlaneStream, + Protocols: []string{protocolTCP, protocolTLS}, + New: func(*options.Options, rt.Lookup) (types.Mechanism, error) { + return &handler{name: names.MechanismRace, spread: types.SpreadRace}, nil + }, + } +} + +// RegistryEntryMirror returns the registry entry of the datagram mirror mechanism. +func RegistryEntryMirror() types.RegistryEntry { + return types.RegistryEntry{ + Name: names.MechanismUDPMirror, ShortName: names.MechanismMirror, Planes: types.PlaneStream, + Protocols: []string{protocolUDP}, + New: func(*options.Options, rt.Lookup) (types.Mechanism, error) { + return &handler{name: names.MechanismMirror, spread: types.SpreadMirror}, nil + }, + } +} + +func (h *handler) Name() types.Name { + return h.name +} + +func (h *handler) Spread() types.Spread { + return h.spread +} + +func (h *handler) StopPool() { + if p := h.Pool(); p != nil { + p.Stop() + } +} + +// ServeHTTP refuses: config validation keeps these mechanisms off every request listener +func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + failures.HandleBadGateway(w, r) +} diff --git a/pkg/backends/alb/mech/spread/spread_test.go b/pkg/backends/alb/mech/spread/spread_test.go new file mode 100644 index 000000000..86957ddc1 --- /dev/null +++ b/pkg/backends/alb/mech/spread/spread_test.go @@ -0,0 +1,63 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package spread + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + + "github.com/stretchr/testify/require" +) + +func TestSpreadMechanisms(t *testing.T) { + for _, test := range []struct { + entry types.RegistryEntry + name types.Name + spread types.Spread + protocols []string + }{ + {RegistryEntryRace(), names.MechanismRace, types.SpreadRace, []string{"tcp", "tls"}}, + {RegistryEntryMirror(), names.MechanismMirror, types.SpreadMirror, []string{"udp"}}, + } { + require.Equal(t, test.name, test.entry.ShortName) + require.Equal(t, types.PlaneStream, test.entry.Planes) + require.Equal(t, test.protocols, test.entry.Protocols) + m, err := test.entry.New(nil, nil) + require.NoError(t, err) + require.Equal(t, test.name, m.Name()) + sm, ok := m.(types.SpreadMechanism) + require.True(t, ok) + require.Equal(t, test.spread, sm.Spread()) + + // stopping before a pool is set is harmless, and a set pool is stopped + sm.StopPool() + require.Nil(t, sm.Pool()) + p := pool.New(nil, 0) + sm.SetPool(p) + require.Equal(t, p, sm.Pool()) + sm.StopPool() + + w := httptest.NewRecorder() + m.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/", nil)) + require.Equal(t, http.StatusBadGateway, w.Code) + } +} diff --git a/pkg/backends/alb/mech/tsm/registry.go b/pkg/backends/alb/mech/tsm/registry.go index 8b4af6a81..6f1249d39 100644 --- a/pkg/backends/alb/mech/tsm/registry.go +++ b/pkg/backends/alb/mech/tsm/registry.go @@ -25,7 +25,7 @@ import ( // RegistryEntry adapts ALB options to the time-series merge constructor. func RegistryEntry() types.RegistryEntry { - return types.RegistryEntry{Name: Name, ShortName: ShortName, New: newFromOptions} + return types.RegistryEntry{Name: Name, ShortName: ShortName, Planes: types.PlaneHTTP, New: newFromOptions} } func newFromOptions(o *options.Options, factories rt.Lookup) (types.Mechanism, error) { diff --git a/pkg/backends/alb/mech/types/types.go b/pkg/backends/alb/mech/types/types.go index 0191d4b7a..eae573174 100644 --- a/pkg/backends/alb/mech/types/types.go +++ b/pkg/backends/alb/mech/types/types.go @@ -22,6 +22,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" + "github.com/trickstercache/trickster/v2/pkg/lb" ) // Name is a type alias for the load balancing mechanism common name @@ -49,9 +50,62 @@ type PoolMechanism interface { Pool() pool.Pool } -// RegistryEntry defines an entry in the ALB Registry +// NewSelectorFunc returns a new selection strategy from the provided Options. It is where +// config is translated into the protocol-neutral core's own option types. +type NewSelectorFunc func(*options.Options) (lb.Selector, error) + +// Plane is a set of dispatch planes: the kinds of listener a mechanism can serve. +type Plane uint8 + +const ( + // PlaneHTTP is request dispatch through an http.Handler. + PlaneHTTP Plane = 1 << iota + // PlaneStream is tcp, tls and udp relaying. + PlaneStream + // PlaneNative is a wire-protocol session server, such as MySQL. + PlaneNative +) + +// Has reports whether every plane in want is in the set. +func (p Plane) Has(want Plane) bool { + return want != 0 && p&want == want +} + +// PickerMechanism is a pool mechanism that selects one member per unit of work, and so can +// serve planes other than HTTP through its Picker. +type PickerMechanism interface { + PoolMechanism + Picker() lb.Picker + Balancer() *lb.Balancer +} + +// Spread is how a mechanism commits one flow to several pool members at once. +type Spread uint8 + +const ( + // SpreadRace connects to several members together and keeps the first to connect. + SpreadRace Spread = iota + 1 + // SpreadMirror copies a flow's datagrams to every member and answers from the first. + SpreadMirror +) + +// SpreadMechanism is a pool mechanism that commits a flow to several members at once, which +// only a stream listener can do. +type SpreadMechanism interface { + PoolMechanism + Spread() Spread +} + +// RegistryEntry defines an entry in the ALB Registry. Exactly one of New and NewSelector is +// set: New for a mechanism that is itself an HTTP handler, NewSelector for a strategy that +// the registry wraps for whichever plane asks. type RegistryEntry struct { Name Name ShortName Name - New NewMechanismFunc + Planes Plane + // Protocols limits a mechanism that serves PlaneStream to some of tcp, tls and udp; + // empty is all of them + Protocols []string + New NewMechanismFunc + NewSelector NewSelectorFunc } diff --git a/pkg/backends/alb/mech/types/types_interfaces_test.go b/pkg/backends/alb/mech/types/types_interfaces_test.go index 5f931c582..b2663fd18 100644 --- a/pkg/backends/alb/mech/types/types_interfaces_test.go +++ b/pkg/backends/alb/mech/types/types_interfaces_test.go @@ -21,15 +21,17 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/fr" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/nlm" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/rr" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/pick" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/tsm" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/ur" uropt "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/ur/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" "github.com/trickstercache/trickster/v2/pkg/backends/prometheus" "github.com/trickstercache/trickster/v2/pkg/backends/providers" rt "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry/types" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" ) // TestPoolMechanismMembership pins which mechs implement PoolMechanism and @@ -43,11 +45,7 @@ func TestPoolMechanismMembership(t *testing.T) { wantsPool bool }{ {"rr", func(t *testing.T) types.Mechanism { - m, err := rr.New(nil, nil) - if err != nil { - t.Fatalf("rr.New: %v", err) - } - return m + return pick.New(names.MechanismRR, rr.New()) }, true}, {"fr", func(t *testing.T) types.Mechanism { m, err := fr.New(nil, nil) @@ -95,3 +93,14 @@ func TestPoolMechanismMembership(t *testing.T) { }) } } + +func TestPlaneHas(t *testing.T) { + t.Parallel() + set := types.PlaneHTTP | types.PlaneStream + if !set.Has(types.PlaneHTTP) || !set.Has(types.PlaneStream) || !set.Has(types.PlaneHTTP|types.PlaneStream) { + t.Error("a plane in the set is reported missing") + } + if set.Has(types.PlaneNative) || set.Has(types.PlaneHTTP|types.PlaneNative) || set.Has(0) { + t.Error("a plane outside the set, or no plane, is reported present") + } +} diff --git a/pkg/backends/alb/mech/ur/user_router.go b/pkg/backends/alb/mech/ur/user_router.go index 802fde16a..9f3685d3a 100644 --- a/pkg/backends/alb/mech/ur/user_router.go +++ b/pkg/backends/alb/mech/ur/user_router.go @@ -70,6 +70,7 @@ func RegistryEntry() types.RegistryEntry { return types.RegistryEntry{ Name: URName, ShortName: names.MechanismUR, + Planes: types.PlaneHTTP | types.PlaneNative, New: New, } } diff --git a/pkg/backends/alb/names/names.go b/pkg/backends/alb/names/names.go index 811ad313e..e870c3d0a 100644 --- a/pkg/backends/alb/names/names.go +++ b/pkg/backends/alb/names/names.go @@ -24,4 +24,22 @@ const ( MechanismNLM = "nlm" MechanismTSM = "tsm" MechanismUR = "ur" + MechanismP2C = "p2c" + MechanismHRW = "hrw" + MechanismLT = "lt" + MechanismLC = "lc" + + MechanismRace = "race" + MechanismMirror = "mirror" +) + +// Mechanism long name constants +const ( + MechanismRoundRobin = "round_robin" + MechanismPowerOfTwoChoices = "power_of_two_choices" + MechanismHighestRandomWeight = "highest_random_weight" + MechanismLeastTime = "least_time" + MechanismLeastConnections = "least_connections" + MechanismConnectRace = "connect_race" + MechanismUDPMirror = "udp_mirror" ) diff --git a/pkg/backends/alb/native/native.go b/pkg/backends/alb/native/native.go new file mode 100644 index 000000000..78866d5c9 --- /dev/null +++ b/pkg/backends/alb/native/native.go @@ -0,0 +1,104 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +// Package native adapts a load balancer's selection strategy to the listeners that speak a +// backend's own protocol: it commits each authenticated session to the pool member the +// strategy picks, for as long as the session lasts. +package native + +import ( + "sync" + + "github.com/trickstercache/trickster/v2/pkg/backends" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Resolver returns a route resolver that balances sessions with picker, which o configures. +// A session is one unit of work: it counts against its member until its route is released. +func Resolver(picker lb.Picker, o *ao.Options) backends.RouteResolver { + if picker == nil { + return nil + } + return &resolver{picker: picker, options: o} +} + +type resolver struct { + picker lb.Picker + options *ao.Options +} + +func (r *resolver) ResolveRoute(in backends.RouteInput) (backends.RouteDecision, bool) { + // each load balancer passed through on the way to a member reads the key its own way + flowOf := func(_ int, p lb.Picker, via *lb.Member) lb.Flow { + if !p.Needs().Has(lb.NeedKey) { + return lb.Flow{} + } + o := r.options + if via != nil { + o = optionsOf(via) + } + return key(o, in) + } + pk, ok := lb.PickLeafFunc(r.picker, flowOf) + if !ok { + return backends.RouteDecision{Outcome: backends.RouteOutcomeUnavailable}, false + } + t, ok := pk.Member().Value.(*pool.Target) + if !ok || t.Backend() == nil { + pk.Done(lb.OutcomeCanceled) + return backends.RouteDecision{Outcome: backends.RouteOutcomeUnavailable}, false + } + var once sync.Once + // the target carries no status: the pool has already held the member to the load balancer's + // healthy_floor, which a second, fixed threshold must not overrule + return backends.RouteDecision{ + Target: backends.RouteTarget{Backend: t.Backend()}, + Outcome: backends.RouteOutcomeSelected, + Release: func() { once.Do(func() { pk.Done(lb.OutcomeOK) }) }, + }, true +} + +func optionsOf(m *lb.Member) *ao.Options { + if t, ok := m.Value.(*pool.Target); ok && t.Backend() != nil { + if cfg := t.Backend().Configuration(); cfg != nil { + return cfg.ALBOptions + } + } + return nil +} + +// key is the session's affinity key as one load balancer is configured to read it: the name +// it authenticated as, or else the client's address +func key(o *ao.Options, in backends.RouteInput) lb.Flow { + prefix := ao.DefaultIPv6Prefix + if o != nil { + if o.HRW.KeySource.Kind == ao.KeyUser { + if in.Username == "" { + return lb.Flow{} + } + return lb.Flow{Key: lb.HashString(in.Username), HasKey: true} + } + if o.HRW.IPv6Prefix > 0 { + prefix = o.HRW.IPv6Prefix + } + } + if !in.Client.IsValid() { + return lb.Flow{} + } + return lb.Flow{Key: lb.HashAddr(in.Client.Unmap(), prefix), HasKey: true} +} diff --git a/pkg/backends/alb/native/native_test.go b/pkg/backends/alb/native/native_test.go new file mode 100644 index 000000000..562e6e058 --- /dev/null +++ b/pkg/backends/alb/native/native_test.go @@ -0,0 +1,241 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package native_test + +import ( + "net/http" + "net/netip" + "strconv" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/alb" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/native" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" + + "github.com/stretchr/testify/require" +) + +type fixture struct { + clients backends.Backends + health healthcheck.StatusLookup +} + +func replicas(t *testing.T, names ...string) *fixture { + t.Helper() + f := &fixture{clients: backends.Backends{}, health: healthcheck.StatusLookup{}} + for _, name := range names { + o := bo.New() + o.Name = name + b, err := backends.New(name, o, nil, http.NotFoundHandler(), nil) + require.NoError(t, err) + f.clients[name] = b + f.health[name] = healthcheck.NewStatus(name, "", "", healthcheck.StatusPassing, time.Time{}, nil) + } + return f +} + +func (f *fixture) alb(t *testing.T, name string, o *ao.Options) *alb.Client { + t.Helper() + b := bo.New() + b.Name = name + b.Provider = providers.ALB + b.ALBOptions = o + require.NoError(t, o.Initialize(name)) + cl, err := alb.NewClient(name, b, nil, nil, nil, nil) + require.NoError(t, err) + f.clients[name] = cl + return cl.(*alb.Client) +} + +func (f *fixture) start(t *testing.T) { + t.Helper() + require.NoError(t, alb.StartALBPools(f.clients, f.health)) + t.Cleanup(func() { _ = alb.StopPools(f.clients) }) +} + +func inflight(c *alb.Client) map[string]int64 { + out := map[string]int64{} + for _, tgt := range c.Pool().ConfiguredTargets() { + out[tgt.Name()] = tgt.Member().Stats().Inflight() + } + return out +} + +func TestSessionsAreBalancedAndCounted(t *testing.T) { + f := replicas(t, "r1", "r2") + c := f.alb(t, "replicas", &ao.Options{MechanismName: "lc", Pool: ao.Members("r1", "r2")}) + f.start(t) + r := c.RouteResolver() + require.NotNil(t, r) + + var held []backends.RouteDecision + for range 4 { + d, ok := r.ResolveRoute(backends.RouteInput{Username: "app", Authenticated: true}) + require.True(t, ok) + require.Equal(t, backends.RouteOutcomeSelected, d.Outcome) + require.True(t, d.Target.Available()) + held = append(held, d) + } + // least connections: open sessions are shared out evenly, and stay counted while open + require.Equal(t, map[string]int64{"r1": 2, "r2": 2}, inflight(c)) + for _, d := range held { + d.Release() + d.Release() + } + require.Equal(t, map[string]int64{"r1": 0, "r2": 0}, inflight(c), "a session is released once") + + // a replica that is down takes no sessions, and a pool with none left refuses them + f.health["r1"].Set(healthcheck.StatusFailing) + for range 3 { + d, ok := r.ResolveRoute(backends.RouteInput{}) + require.True(t, ok) + require.Equal(t, "r2", d.Target.Backend.Name()) + d.Release() + } + f.health["r2"].Set(healthcheck.StatusFailing) + d, ok := r.ResolveRoute(backends.RouteInput{}) + require.False(t, ok) + require.Equal(t, backends.RouteOutcomeUnavailable, d.Outcome) + require.Nil(t, d.Release) +} + +func TestSessionsKeepToAMemberByUserOrAddress(t *testing.T) { + names := []string{"r1", "r2", "r3", "r4", "r5"} + f := replicas(t, names...) + byUser := f.alb(t, "by-user", &ao.Options{ + MechanismName: "hrw", Pool: ao.Members(names...), HRW: ao.HRWOptions{Key: "user"}, + }) + byAddr := f.alb(t, "by-addr", &ao.Options{MechanismName: "hrw", Pool: ao.Members(names...)}) + f.start(t) + owner := func(c *alb.Client, in backends.RouteInput) string { + d, ok := c.RouteResolver().ResolveRoute(in) + require.True(t, ok) + d.Release() + return d.Target.Backend.Name() + } + here, there := netip.MustParseAddr("198.51.100.7"), netip.MustParseAddr("::ffff:203.0.113.9") + users, addrs := map[string]bool{}, map[string]bool{} + for i := range 30 { + user := "tenant" + strconv.Itoa(i) + first := owner(byUser, backends.RouteInput{Username: user, Client: here}) + users[first] = true + require.Equal(t, first, owner(byUser, backends.RouteInput{Username: user, Client: there}), user) + + addr := netip.AddrFrom4([4]byte{198, 51, 100, byte(i)}) + first = owner(byAddr, backends.RouteInput{Username: "a", Client: addr}) + addrs[first] = true + require.Equal(t, first, owner(byAddr, backends.RouteInput{Username: "b", Client: netip.AddrFrom16(addr.As16())})) + } + require.GreaterOrEqual(t, len(users), 4) + require.GreaterOrEqual(t, len(addrs), 4) + // a session with nothing to key on is served all the same + spread := map[string]bool{} + for range 60 { + spread[owner(byUser, backends.RouteInput{})] = true + spread[owner(byAddr, backends.RouteInput{})] = true + } + require.GreaterOrEqual(t, len(spread), 3) +} + +func TestNestedLoadBalancersKeyForThemselves(t *testing.T) { + f := replicas(t, "a1", "a2", "a3", "b1") + f.alb(t, "inner-a", &ao.Options{ + MechanismName: "hrw", Pool: ao.Members("a1", "a2", "a3"), HRW: ao.HRWOptions{Key: "user"}, + }) + f.alb(t, "inner-b", &ao.Options{MechanismName: "rr", Pool: ao.Members("b1")}) + outer := f.alb(t, "outer", &ao.Options{MechanismName: "rr", Pool: ao.Members("inner-a", "inner-b")}) + f.start(t) + seen := map[string]bool{} + for range 20 { + d, ok := outer.RouteResolver().ResolveRoute(backends.RouteInput{Username: "app"}) + require.True(t, ok) + seen[d.Target.Backend.Name()] = true + d.Release() + } + require.Len(t, seen, 2, "one user keeps to one member of the keyed pool: %v", seen) + require.True(t, seen["b1"]) +} + +func TestOnlySessionStrategiesResolveRoutes(t *testing.T) { + f := replicas(t, "r1") + fanout := f.alb(t, "fanout", &ao.Options{MechanismName: "fr", Pool: ao.Members("r1")}) + timed := f.alb(t, "timed", &ao.Options{MechanismName: "lt", Pool: ao.Members("r1")}) + f.start(t) + require.Nil(t, fanout.RouteResolver()) + require.Nil(t, timed.RouteResolver(), "a native session reports no latency to rank members by") + require.Nil(t, native.Resolver(nil, nil)) +} + +// a member whose payload is not a backend cannot take a session, and is not blamed for it +func TestForeignMembersAreRefused(t *testing.T) { + m := lb.NewMember(lb.MemberOptions{Name: "foreign", Value: 7}) + p, err := lb.NewPool([]*lb.Member{m}, 0) + require.NoError(t, err) + defer p.Stop() + b := lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: p, Ejection: lb.EjectionOptions{Failures: 1}}) + d, ok := native.Resolver(b, nil).ResolveRoute(backends.RouteInput{}) + require.False(t, ok) + require.Equal(t, backends.RouteOutcomeUnavailable, d.Outcome) + require.Zero(t, m.Stats().Inflight()) + require.Zero(t, m.Stats().Failures()) +} + +// the load balancer's healthy_floor alone decides which members take sessions, as it does for +// requests and connections: a route the pool admits is never refused for its status afterwards +func TestSessionsHonorTheHealthyFloor(t *testing.T) { + statuses := map[string]int32{ + "failing": healthcheck.StatusFailing, "unchecked": healthcheck.StatusUnchecked, "passing": healthcheck.StatusPassing, + } + for floor, want := range map[int][]string{ + -1: {"failing", "passing", "unchecked"}, + 0: {"passing", "unchecked"}, + 1: {"passing"}, + } { + f := replicas(t, "failing", "unchecked", "passing") + for name, status := range statuses { + f.health[name].Set(status) + // probed, so a floor of 1 is not reset for members that could never reach it + f.clients[name].Configuration().HealthCheck = &ho.Options{Interval: timeconv.Duration(time.Second)} + } + c := f.alb(t, "floor"+strconv.Itoa(floor+1), &ao.Options{ + MechanismName: "rr", HealthyFloor: floor, Pool: ao.Members("failing", "unchecked", "passing"), + }) + f.start(t) + reached := map[string]bool{} + for range 12 { + d, ok := c.RouteResolver().ResolveRoute(backends.RouteInput{}) + require.True(t, ok, "floor %d", floor) + require.True(t, d.Target.Available(), "floor %d refused %s after selecting it", floor, d.Target.Backend.Name()) + reached[d.Target.Backend.Name()] = true + d.Release() + } + got := make([]string, 0, len(reached)) + for name := range reached { + got = append(got, name) + } + require.ElementsMatch(t, want, got, "floor %d", floor) + } +} diff --git a/pkg/backends/alb/native/options_test.go b/pkg/backends/alb/native/options_test.go new file mode 100644 index 000000000..be7b23547 --- /dev/null +++ b/pkg/backends/alb/native/options_test.go @@ -0,0 +1,29 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package native + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +func TestOptionsOfAStrayMember(t *testing.T) { + if optionsOf(lb.NewMember(lb.MemberOptions{Name: "stray", Value: "not a target"})) != nil { + t.Error("a member that is no load balancer has options") + } +} diff --git a/pkg/backends/alb/observe/members.go b/pkg/backends/alb/observe/members.go new file mode 100644 index 000000000..395ca927f --- /dev/null +++ b/pkg/backends/alb/observe/members.go @@ -0,0 +1,91 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package observe + +import ( + "maps" + "sync" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/observability/keys" + + "github.com/prometheus/client_golang/prometheus" +) + +// memberInflightDesc describes the in-flight gauge. It is filled at scrape time from each +// tracked balancer's members, so dispatch pays nothing for it and a member that leaves its +// pool takes its series with it. +var memberInflightDesc = prometheus.NewDesc( + "trickster_alb_member_inflight", + "Current number of requests in flight to an ALB pool member, for mechanisms that track it.", + []string{keys.ALB_Name, keys.Member}, nil, +) + +var ( + trackedMtx sync.Mutex + tracked = make(map[string]*lb.Balancer) +) + +// Track exports the in-flight count of every member of the named ALB's balancer. It does +// nothing for a strategy that keeps no in-flight count. +func Track(albName string, b *lb.Balancer) { + if b == nil || albName == "" || !b.Needs().Has(lb.NeedInflight) { + return + } + trackedMtx.Lock() + tracked[albName] = b + trackedMtx.Unlock() +} + +// Untrack stops exporting the named ALB's members, unless another balancer has since taken +// the name, as the ALB of a reloaded config does. +func Untrack(albName string, b *lb.Balancer) { + trackedMtx.Lock() + if tracked[albName] == b { + delete(tracked, albName) + } + trackedMtx.Unlock() +} + +type memberCollector struct{} + +func (memberCollector) Describe(ch chan<- *prometheus.Desc) { + ch <- memberInflightDesc +} + +func (memberCollector) Collect(ch chan<- prometheus.Metric) { + trackedMtx.Lock() + balancers := make(map[string]*lb.Balancer, len(tracked)) + maps.Copy(balancers, tracked) + trackedMtx.Unlock() + for name, b := range balancers { + p := b.Pool() + if p == nil { + continue + } + for _, m := range p.Configured() { + if m.Name() == "" { + continue + } + ch <- prometheus.MustNewConstMetric(memberInflightDesc, prometheus.GaugeValue, + float64(m.Stats().Inflight()), name, m.Name()) + } + } +} + +func init() { + prometheus.MustRegister(memberCollector{}) +} diff --git a/pkg/backends/alb/observe/members_test.go b/pkg/backends/alb/observe/members_test.go new file mode 100644 index 000000000..62fd9061a --- /dev/null +++ b/pkg/backends/alb/observe/members_test.go @@ -0,0 +1,79 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package observe + +import ( + "strings" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lc" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" + + "github.com/prometheus/client_golang/prometheus/testutil" +) + +func TestMemberInflightIsCollectedAtScrapeTime(t *testing.T) { + members := []*lb.Member{ + lb.NewMember(lb.MemberOptions{Name: "a"}), + lb.NewMember(lb.MemberOptions{Name: "b"}), + lb.NewMember(lb.MemberOptions{}), + } + p, err := lb.NewPool(members, 0) + if err != nil { + t.Fatal(err) + } + defer p.Stop() + b := lb.NewBalancer(lc.New(), lb.BalancerOptions{Pool: p}) + Track("inflight-alb", b) + Track("", b) + Track("untracked-rr", lb.NewBalancer(rr.New(), lb.BalancerOptions{Pool: p})) + Track("no-balancer", nil) + poolless := lb.NewBalancer(lc.New()) + Track("poolless-alb", poolless) + defer Untrack("poolless-alb", poolless) + + var held []lb.Pick + for range 6 { + pk, _ := b.Pick(lb.Flow{}) + held = append(held, pk) + } + // six flows over three members, one of them unnamed and so not exported + want := ` +# HELP trickster_alb_member_inflight Current number of requests in flight to an ALB pool member, for mechanisms that track it. +# TYPE trickster_alb_member_inflight gauge +trickster_alb_member_inflight{alb_name="inflight-alb",member="a"} 2 +trickster_alb_member_inflight{alb_name="inflight-alb",member="b"} 2 +` + if err := testutil.CollectAndCompare(memberCollector{}, strings.NewReader(want)); err != nil { + t.Error(err) + } + for _, pk := range held { + pk.Done(lb.OutcomeOK) + } + + // a reloaded ALB takes over its name; stopping the old one must not drop the new series + next := lb.NewBalancer(lc.New(), lb.BalancerOptions{Pool: p}) + Track("inflight-alb", next) + Untrack("inflight-alb", b) + if got := testutil.CollectAndCount(memberCollector{}); got != 2 { + t.Errorf("series after the old balancer stopped = %d, want 2", got) + } + Untrack("inflight-alb", next) + if got := testutil.CollectAndCount(memberCollector{}); got != 0 { + t.Errorf("series after untracking = %d", got) + } +} diff --git a/pkg/backends/alb/observe/observe.go b/pkg/backends/alb/observe/observe.go new file mode 100644 index 000000000..cc0471c51 --- /dev/null +++ b/pkg/backends/alb/observe/observe.go @@ -0,0 +1,66 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package observe binds the load-balancing core's events to Trickster's logger and metrics. +package observe + +import ( + "fmt" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/observability/logging" + "github.com/trickstercache/trickster/v2/pkg/observability/logging/logger" + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" +) + +// WorkerRefresh labels a panic recovered while a pool rebuilt its snapshot. +const WorkerRefresh = "refresh" + +// balancerObserver meters one ALB's balancer events +type balancerObserver struct{ albName string } + +// Balancer returns the observer for the named ALB's balancer. +func Balancer(albName string) lb.Observer { + return balancerObserver{albName: albName} +} + +func (o balancerObserver) Observe(ev lb.Event) { + if ev.Kind != lb.EventEjected { + return + } + logger.Warn("alb pool member ejected after repeated connect failures", logging.Pairs{ + "albName": o.albName, "member": ev.Member, + }) + metrics.ALBMemberEjections.WithLabelValues(o.albName, ev.Member).Inc() +} + +type poolObserver struct{} + +// Pool returns the observer every ALB pool reports to. +func Pool() lb.Observer { + return poolObserver{} +} + +func (poolObserver) Observe(ev lb.Event) { + if ev.Kind != lb.EventPanic { + return + } + logger.Error("alb pool refresh panic", logging.Pairs{ + "worker": WorkerRefresh, + "panic": fmt.Sprintf("%v", ev.Panic), + "stack": string(ev.Stack), + }) + metrics.ALBPoolRefreshPanicRecovered.WithLabelValues(WorkerRefresh).Inc() +} diff --git a/pkg/backends/alb/observe/observe_test.go b/pkg/backends/alb/observe/observe_test.go new file mode 100644 index 000000000..f31e92e04 --- /dev/null +++ b/pkg/backends/alb/observe/observe_test.go @@ -0,0 +1,59 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package observe + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + + "github.com/prometheus/client_golang/prometheus/testutil" +) + +func TestPoolObserverMetersRecoveredPanics(t *testing.T) { + c := metrics.ALBPoolRefreshPanicRecovered.WithLabelValues(WorkerRefresh) + before := testutil.ToFloat64(c) + o := Pool() + o.Observe(lb.Event{Kind: lb.EventSnapshot, Gen: 1}) + if got := testutil.ToFloat64(c) - before; got != 0 { + t.Errorf("a snapshot event was metered as a panic: +%v", got) + } + o.Observe(lb.Event{Kind: lb.EventPanic, Panic: "boom", Stack: []byte("stack")}) + if got := testutil.ToFloat64(c) - before; got != 1 { + t.Errorf("recovered panics metered = +%v, want +1", got) + } +} + +func TestBalancerObserverMetersEjections(t *testing.T) { + c := metrics.ALBMemberEjections.WithLabelValues("observed-alb", "observed-member") + before := testutil.ToFloat64(c) + o := Balancer("observed-alb") + o.Observe(lb.Event{Kind: lb.EventSnapshot}) + o.Observe(lb.Event{Kind: lb.EventPanic, Panic: "boom"}) + if got := testutil.ToFloat64(c) - before; got != 0 { + t.Errorf("an event that is not an ejection was metered as one: +%v", got) + } + o.Observe(lb.Event{Kind: lb.EventEjected, Member: "observed-member"}) + if got := testutil.ToFloat64(c) - before; got != 1 { + t.Errorf("ejections metered = +%v, want +1", got) + } + // a member that leaves takes its series with it + metrics.DeleteBackendSeries("observed-member") + if got := testutil.ToFloat64(metrics.ALBMemberEjections.WithLabelValues("observed-alb", "observed-member")); got != 0 { + t.Errorf("the series survived its member: %v", got) + } +} diff --git a/pkg/backends/alb/options/keysource.go b/pkg/backends/alb/options/keysource.go new file mode 100644 index 000000000..aae287256 --- /dev/null +++ b/pkg/backends/alb/options/keysource.go @@ -0,0 +1,146 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package options + +import ( + "errors" + "fmt" + "net/textproto" + "strconv" + "strings" +) + +// KeyKind is where a flow's affinity key is read from. +type KeyKind uint8 + +const ( + // KeyClientIP keys on the client address, after trusted-proxy resolution; never the port. + KeyClientIP KeyKind = iota + // KeyHost keys on the request's host, without its port and without regard to case. + KeyHost + // KeyHeader keys on the first value of a request header. + KeyHeader + // KeyCookie keys on the value of a request cookie. + KeyCookie + // KeyQuery keys on the raw value of a query string parameter. + KeyQuery + // KeySNI keys on the TLS server name a client offered; tls stream listeners only. + KeySNI + // KeyProxyTLV keys on the value of a PROXY protocol v2 TLV; tcp and tls stream listeners + // that accept the PROXY protocol only. + KeyProxyTLV + // KeyUser keys on the name a session authenticated as; native protocol listeners only. + KeyUser +) + +// Key source spellings +const ( + KeySourceClientIP = "client_ip" + KeySourceHost = "host" + KeySourceSNI = "sni" + KeySourceUser = "user" + keyPrefixHeader = "header:" + keyPrefixCookie = "cookie:" + keyPrefixQuery = "query:" + keyPrefixProxyTLV = "proxy_tlv:" +) + +// ErrInvalidKeySource is returned for a key source that is not one of client_ip, host, sni, +// user, header:, cookie:, query: or proxy_tlv:. +var ErrInvalidKeySource = errors.New("invalid key source") + +// StreamListener is what decides which keys a stream listener can read. +type StreamListener struct { + // TLS is set for a tls listener, which reads the server name a client offers + TLS bool + // ProxyProtocol is set for a tcp or tls listener that accepts a PROXY protocol header + ProxyProtocol bool +} + +// OnStream reports whether a stream listener can read the key: the client address always, the +// server name when it is tls, a TLV when it accepts the PROXY protocol, nothing of a request. +func (k KeySource) OnStream(l StreamListener) bool { + switch k.Kind { + case KeyClientIP: + return true + case KeySNI: + return l.TLS + case KeyProxyTLV: + return l.ProxyProtocol + } + return false +} + +// OnHTTP reports whether an HTTP listener can read the key, which is any that is part of a +// request or is the client address. +func (k KeySource) OnHTTP() bool { + return k.Kind != KeySNI && k.Kind != KeyProxyTLV && k.Kind != KeyUser +} + +// OnNative reports whether a native protocol listener can read the key: the client address, +// or the name its session authenticated as. +func (k KeySource) OnNative() bool { + return k.Kind == KeyClientIP || k.Kind == KeyUser +} + +// KeySource is a parsed affinity key source. +type KeySource struct { + Kind KeyKind + // Name is the header (in canonical form), cookie or query parameter; empty otherwise. + Name string + // TLV is the PROXY protocol v2 TLV type of a KeyProxyTLV + TLV byte +} + +// ParseKeySource parses a key source; the empty string is client_ip. +func ParseKeySource(s string) (KeySource, error) { + s = strings.TrimSpace(s) + switch s { + case "", KeySourceClientIP: + return KeySource{Kind: KeyClientIP}, nil + case KeySourceHost: + return KeySource{Kind: KeyHost}, nil + case KeySourceSNI: + return KeySource{Kind: KeySNI}, nil + case KeySourceUser: + return KeySource{Kind: KeyUser}, nil + } + if t, ok := strings.CutPrefix(s, keyPrefixProxyTLV); ok { + // the type is a byte, written in decimal or, as the PROXY protocol does, 0x hex + if n, err := strconv.ParseUint(strings.TrimSpace(t), 0, 8); err == nil { + return KeySource{Kind: KeyProxyTLV, TLV: byte(n)}, nil + } + } + for prefix, kind := range map[string]KeyKind{ + keyPrefixHeader: KeyHeader, keyPrefixCookie: KeyCookie, keyPrefixQuery: KeyQuery, + } { + name, ok := strings.CutPrefix(s, prefix) + if !ok { + continue + } + name = strings.TrimSpace(name) + if name == "" || strings.ContainsAny(name, " \t;=&,") { + break + } + if kind == KeyHeader { + name = textproto.CanonicalMIMEHeaderKey(name) + } + return KeySource{Kind: kind, Name: name}, nil + } + return KeySource{}, fmt.Errorf("%w %q: use %s, %s, %s, %s, %s, %s, %s or %s", + ErrInvalidKeySource, s, KeySourceClientIP, KeySourceHost, KeySourceSNI, KeySourceUser, + keyPrefixHeader, keyPrefixCookie, keyPrefixQuery, keyPrefixProxyTLV) +} diff --git a/pkg/backends/alb/options/options.go b/pkg/backends/alb/options/options.go index 791d81061..e924ca7c3 100644 --- a/pkg/backends/alb/options/options.go +++ b/pkg/backends/alb/options/options.go @@ -42,6 +42,10 @@ type Options struct { // Pool provides the list of pool members (backend name + optional // weight) to be used by the load balancer Pool PoolMemberList `yaml:"pool,omitempty"` + // Name is the ALB backend's name, set by Initialize + Name string `yaml:"-"` + // PoolRepeats lists the member names Initialize found repeated in Pool and removed + PoolRepeats []PoolRepeat `yaml:"-"` // Discovery, when set, binds this ALB's pool to a named discoverer from // the top-level 'discovery' config section; discovered members are // additive to static Pool entries @@ -55,6 +59,10 @@ type Options struct { // Unknown means the first health check hasn't returned yet, or the target // backend has no health check interval configured. HealthyFloor int `yaml:"healthy_floor,omitempty"` + // PropagateHealth makes the ALB unavailable, as a member of another ALB's pool, while none + // of its own members is available, so that pool dispatches to its other members instead + // of having this one fail its share. + PropagateHealth bool `yaml:"propagate_health,omitempty"` // MaxCaptureBytes overrides the backend-level max_capture_bytes for this // ALB's fanout members. Set this when the ALB's expected response shape // differs from the backend default (e.g. a TSM fan-out of 50 small-payload @@ -77,18 +85,24 @@ type Options struct { UserRouter *ur.Options `yaml:"user_router,omitempty"` // // synthetic values - FgrCodesLookup sets.Set[int] `yaml:"-"` + // FGRGoodCodes is the compiled set of status codes that fgr accepts + FGRGoodCodes *types.StatusTable `yaml:"-"` // mechanism-specific options TSMOptions tsmoptions.Options `yaml:"tsm,omitempty"` NLMOptions NewestLastModifiedOptions `yaml:"nlm,omitempty"` - FGROptions FirstGoodResponseOptions `yaml:"fgr,omitempty"` + HRW HRWOptions `yaml:"hrw,omitempty"` + LT LTOptions `yaml:"lt,omitempty"` + // Stream holds what applies only when the ALB balances tcp, tls or udp flows + Stream *StreamOptions `yaml:"stream,omitempty"` + FGROptions FirstGoodResponseOptions `yaml:"fgr,omitempty"` } type FirstGoodResponseOptions struct { // StatusCodes provides an explicit list of status codes considered "good" when using - // the First Good Response (fgr) methodology. By default, any code < 400 is good. - StatusCodes []int `yaml:"status_codes,omitempty"` + // the First Good Response (fgr) methodology: bare codes, inclusive {start, end} ranges, + // or a mix. By default, any code < 400 is good. + StatusCodes types.StatusRanges `yaml:"status_codes,omitempty"` ConcurrencyOptions ConcurrencyOptions `yaml:",inline"` } @@ -111,8 +125,14 @@ var _ types.ConfigOptions[Options] = &Options{} const defaultTSOutputFormat = providers.Prometheus +// DefaultFGRStatusCodes returns the status codes fgr accepts when none are configured. +func DefaultFGRStatusCodes() types.StatusRanges { + return types.StatusRanges{{Start: 100, End: 399}} +} + var ( ErrUserRouterRequired = errors.New("'user_router' block is required") + ErrPropagateHealthNoPool = errors.New("'propagate_health' is not valid for mechanism 'ur', which has no pool") ErrInvalidOutputFormat = errors.New("value for 'output_format' is invalid") ErrOutputFormatOnlyForTSM = errors.New("'output_format' option is only valid for provider 'alb' and mechanism 'tsmerge'") ) @@ -132,18 +152,6 @@ func New() *Options { // Clone returns a perfect copy of the Options func (o *Options) Clone() *Options { - var fsc []int - var fscm sets.Set[int] - - if o.FGRStatusCodes != nil { - fsc = make([]int, len(o.FGRStatusCodes)) - copy(fsc, o.FGRStatusCodes) - } - - if o.FgrCodesLookup != nil { - fscm = o.FgrCodesLookup.Clone() - } - c := pointers.Clone(o) if o.UserRouter != nil { c.UserRouter = o.UserRouter.Clone() @@ -152,12 +160,33 @@ func (o *Options) Clone() *Options { c.Discovery = o.Discovery.Clone() } c.Pool = slices.Clone(o.Pool) - c.FGRStatusCodes = fsc - c.FgrCodesLookup = fscm + c.PoolRepeats = slices.Clone(o.PoolRepeats) + c.FGRStatusCodes = slices.Clone(o.FGRStatusCodes) + c.FGROptions.StatusCodes = slices.Clone(o.FGROptions.StatusCodes) + c.Stream = o.Stream.Clone() + c.LT.StatusCodes = slices.Clone(o.LT.StatusCodes) + if o.LT.GoodCodes != nil { + table := *o.LT.GoodCodes + c.LT.GoodCodes = &table + } + if o.FGRGoodCodes != nil { + table := *o.FGRGoodCodes + c.FGRGoodCodes = &table + } return c } -func (o *Options) Initialize(_ string) error { +func (o *Options) Initialize(name string) error { + if name != "" { + o.Name = name + } + pool, repeats, err := o.Pool.Dedupe(name) + if err != nil { + return err + } + if len(repeats) > 0 { + o.Pool, o.PoolRepeats = pool, repeats + } if strings.HasPrefix(o.MechanismName, names.MechanismTSM) && o.MechanismName != names.MechanismTSM { // shorten from tsmerge to tsm o.MechanismName = names.MechanismTSM @@ -166,18 +195,23 @@ func (o *Options) Initialize(_ string) error { case names.MechanismFGR: // apply deprecated top-level FGRStatusCodes to new FROptions level if len(o.FGRStatusCodes) > 0 && len(o.FGROptions.StatusCodes) == 0 { - o.FGROptions.StatusCodes = o.FGRStatusCodes + o.FGROptions.StatusCodes = types.StatusCodes(o.FGRStatusCodes...) } - if len(o.FGROptions.StatusCodes) > 0 { - o.FgrCodesLookup = sets.NewIntSet() - o.FgrCodesLookup.SetAll(o.FGROptions.StatusCodes) + codes := o.FGROptions.StatusCodes + if len(codes) == 0 { + codes = DefaultFGRStatusCodes() } + o.FGRGoodCodes = codes.Compile() case names.MechanismTSM: if o.OutputFormat == "" { o.OutputFormat = defaultTSOutputFormat } } + if err := o.initializeStrategies(); err != nil { + return err + } + if o.Discovery != nil { if err := o.Discovery.Initialize(""); err != nil { return err @@ -187,12 +221,36 @@ func (o *Options) Initialize(_ string) error { return nil } +// PoolRepeatWarning returns the deprecation warning for member names that were repeated in +// the pool, naming the weight that restores each one's former share; empty when none were. +func (o *Options) PoolRepeatWarning(albName string) string { + if len(o.PoolRepeats) == 0 { + return "" + } + var sb strings.Builder + fmt.Fprintf(&sb, "alb %q: repeating a pool member no longer increases its share;"+ + " repeats were ignored. to keep the former split, set", albName) + for i, r := range o.PoolRepeats { + if i > 0 { + sb.WriteByte(',') + } + fmt.Fprintf(&sb, " {name: %s, weight: %d}", r.Name, r.Weight) + } + return sb.String() +} + func (o *Options) Validate() (bool, error) { + if err := o.FGROptions.StatusCodes.Validate(); err != nil { + return false, fmt.Errorf("fgr.status_codes: %w", err) + } switch o.MechanismName { case names.MechanismUR: if o.UserRouter == nil { return false, ErrUserRouterRequired } + if o.PropagateHealth { + return false, ErrPropagateHealthNoPool + } case names.MechanismTSM: if o.OutputFormat != "" && !providers.IsSupportedTimeSeriesMergeProvider(o.OutputFormat) { return false, ErrInvalidOutputFormat @@ -202,6 +260,9 @@ func (o *Options) Validate() (bool, error) { return false, ErrOutputFormatOnlyForTSM } } + if err := o.validateStrategies(); err != nil { + return false, err + } if o.Discovery != nil { if _, err := o.Discovery.Validate(); err != nil { return false, err @@ -214,6 +275,10 @@ func (o *Options) ValidatePool(backendName string, allBackends sets.Set[string]) if err := o.Pool.Validate(backendName); err != nil { return err } + // discovered members are never backups, so a discovered pool may list only standbys + if o.Discovery == nil && o.Pool.AllBackups() { + return fmt.Errorf("%w (alb %q)", ErrNoPrimaryPoolMember, backendName) + } for _, m := range o.Pool { if _, ok := allBackends[m.Name]; !ok { return te.NewErrInvalidPoolMemberName(backendName, m.Name) diff --git a/pkg/backends/alb/options/options_test.go b/pkg/backends/alb/options/options_test.go index df58f8c00..b1c9914d0 100644 --- a/pkg/backends/alb/options/options_test.go +++ b/pkg/backends/alb/options/options_test.go @@ -23,6 +23,7 @@ import ( ur "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/ur/options" "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/config/types" "github.com/trickstercache/trickster/v2/pkg/util/sets" "github.com/stretchr/testify/require" @@ -76,7 +77,8 @@ func TestClone(t *testing.T) { o := New() o.Pool = Members("test") o.FGRStatusCodes = []int{200} - o.FgrCodesLookup = sets.New([]int{200}) + o.FGROptions.StatusCodes = types.StatusRanges{{Start: 200, End: 299}} + o.FGRGoodCodes = o.FGROptions.StatusCodes.Compile() require.NotNil(t, o) co := o.Clone() @@ -86,9 +88,13 @@ func TestClone(t *testing.T) { if len(co.FGRStatusCodes) != 1 || co.FGRStatusCodes[0] != 200 { t.Error("status codes mismatch") } - if len(co.FgrCodesLookup) != 1 || !co.FgrCodesLookup.Contains(200) { + if !co.FGRGoodCodes.Contains(250) || co.FGRGoodCodes == o.FGRGoodCodes { t.Error("fgr lookup mismatch") } + co.FGROptions.StatusCodes[0].End = 200 + if o.FGROptions.StatusCodes[0].End != 299 { + t.Error("the clone shares its status code ranges with the original") + } } func TestInitialize(t *testing.T) { @@ -118,7 +124,7 @@ func TestInitialize(t *testing.T) { if err != nil { t.Error("failed to set defaults") } - if o.FgrCodesLookup == nil || !o.FgrCodesLookup.Contains(200) || !o.FgrCodesLookup.Contains(201) { + if !o.FGRGoodCodes.Contains(200) || !o.FGRGoodCodes.Contains(201) || o.FGRGoodCodes.Contains(204) { t.Error("expected FGR codes lookup to be set") } @@ -160,9 +166,32 @@ func TestInitializeDeprecatedFGRAndDefaultOutputFormat(t *testing.T) { o.MechanismName = names.MechanismFGR o.FGRStatusCodes = []int{200, 204} require.NoError(t, o.Initialize("")) - require.Equal(t, []int{200, 204}, o.FGROptions.StatusCodes) - require.True(t, o.FgrCodesLookup.Contains(200)) - require.True(t, o.FgrCodesLookup.Contains(204)) + require.Equal(t, types.StatusCodes(200, 204), o.FGROptions.StatusCodes) + require.True(t, o.FGRGoodCodes.Contains(200)) + require.True(t, o.FGRGoodCodes.Contains(204)) + require.False(t, o.FGRGoodCodes.Contains(201)) + + // with nothing configured, anything below 400 is good + o = New() + o.MechanismName = names.MechanismFGR + require.NoError(t, o.Initialize("")) + require.True(t, o.FGRGoodCodes.Contains(100)) + require.True(t, o.FGRGoodCodes.Contains(399)) + require.False(t, o.FGRGoodCodes.Contains(400)) + + // ranges and bare codes mix, and a range out of bounds is refused + o = New() + require.NoError(t, yaml.Unmarshal( + []byte("mechanism: fgr\nfgr:\n status_codes: [{start: 200, end: 299}, 304]\n"), o)) + require.NoError(t, o.Initialize("")) + require.True(t, o.FGRGoodCodes.Contains(250)) + require.True(t, o.FGRGoodCodes.Contains(304)) + require.False(t, o.FGRGoodCodes.Contains(303)) + _, err := o.Validate() + require.NoError(t, err) + o.FGROptions.StatusCodes = types.StatusRanges{{Start: 200, End: 700}} + _, err = o.Validate() + require.ErrorIs(t, err, types.ErrInvalidStatusRange) o = New() o.MechanismName = names.MechanismTSM diff --git a/pkg/backends/alb/options/pool.go b/pkg/backends/alb/options/pool.go index 0d4ff1d2c..cdc616864 100644 --- a/pkg/backends/alb/options/pool.go +++ b/pkg/backends/alb/options/pool.go @@ -37,19 +37,78 @@ import ( // // Weights apply to mechanisms that select a single member per request // (round_robin); fan-out mechanisms dispatch to every member regardless of -// weight. A weight of 0 (or omitted) means 1. Weights replace the legacy -// workaround of repeating a member name to increase its share. +// weight. A weight of 0 (or omitted) means 1. A weight is the only way to +// increase a member's share: a name repeated in the list is de-duplicated. +// +// A backup member stands by: it is used only while no other member is available. type PoolMember struct { Name string `yaml:"name"` Weight int `yaml:"weight,omitempty"` + Backup bool `yaml:"backup,omitempty"` } +// BackupTier is the failover tier of a backup member; every other member is in tier 0 +const BackupTier = 1 + // PoolMemberList is the ALB pool as configured type PoolMemberList []PoolMember // ErrInvalidPoolWeight is returned when a pool entry has a negative weight var ErrInvalidPoolWeight = errors.New("pool member 'weight' cannot be negative") +// ErrNoPrimaryPoolMember is returned when every member of a pool is a backup +var ErrNoPrimaryPoolMember = errors.New("pool needs at least one member that is not a 'backup'") + +// ErrConflictingPoolWeights is returned when a pool lists one member under different weights +var ErrConflictingPoolWeights = errors.New("pool member is repeated with different 'weight' values") + +// PoolRepeat describes a member name that a pool listed more than once +type PoolRepeat struct { + Name string + // Count is how many times the name was listed + Count int + // Weight is the combined effective weight the repeated entries carried + Weight int +} + +// Dedupe returns the list with repeated member names removed, the first occurrence winning, +// and what was repeated. Repeats that set different explicit weights are an error. +func (l PoolMemberList) Dedupe(albName string) (PoolMemberList, []PoolRepeat, error) { + seen := make(map[string]int, len(l)) + var repeats []PoolRepeat + var repeatIdx map[string]int + out := l + for i, m := range l { + j, dup := seen[m.Name] + if !dup { + seen[m.Name] = i + if repeats != nil { + out = append(out, m) + } + continue + } + first := l[j] + if first.Weight > 0 && m.Weight > 0 && first.Weight != m.Weight { + return nil, nil, fmt.Errorf("%w (member %q of alb %q: %d and %d)", + ErrConflictingPoolWeights, m.Name, albName, first.Weight, m.Weight) + } + if repeats == nil { + // the first repeat: everything before it is kept as is + out = append(make(PoolMemberList, 0, len(l)-1), l[:i]...) + repeatIdx = make(map[string]int) + } + k, ok := repeatIdx[m.Name] + if !ok { + k = len(repeats) + repeatIdx[m.Name] = k + repeats = append(repeats, PoolRepeat{Name: m.Name, Count: 1, Weight: first.EffectiveWeight()}) + } + repeats[k].Count++ + repeats[k].Weight += m.EffectiveWeight() + } + return out, repeats, nil +} + // Members returns a PoolMemberList of the provided names, each with the // default weight func Members(names ...string) PoolMemberList { @@ -69,6 +128,14 @@ func (l PoolMemberList) Names() []string { return out } +// Tier returns the member's failover tier +func (m PoolMember) Tier() int { + if m.Backup { + return BackupTier + } + return 0 +} + // EffectiveWeight returns the member's weight for apportionment purposes; // an unset (0) weight is 1 func (m PoolMember) EffectiveWeight() int { @@ -97,7 +164,7 @@ func (m *PoolMember) UnmarshalYAML(value *yaml.Node) error { // MarshalYAML renders unweighted members as plain name scalars so sanitized // config output matches the common input form func (m PoolMember) MarshalYAML() (any, error) { - if m.Weight == 0 { + if m.Weight == 0 && !m.Backup { return m.Name, nil } type dumpPoolMember PoolMember @@ -114,3 +181,13 @@ func (l PoolMemberList) Validate(albName string) error { } return nil } + +// AllBackups reports whether the list has members and every one of them is a backup +func (l PoolMemberList) AllBackups() bool { + for _, m := range l { + if !m.Backup { + return false + } + } + return len(l) > 0 +} diff --git a/pkg/backends/alb/options/pool_test.go b/pkg/backends/alb/options/pool_test.go index 3ca4d4d49..631bd44ca 100644 --- a/pkg/backends/alb/options/pool_test.go +++ b/pkg/backends/alb/options/pool_test.go @@ -20,6 +20,10 @@ import ( "strings" "testing" + ur "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/ur/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/util/sets" + "github.com/stretchr/testify/require" "go.yaml.in/yaml/v3" ) @@ -71,3 +75,101 @@ func TestValidatePoolRejectsNegativeWeight(t *testing.T) { err := o.ValidatePool("alb1", nil) require.ErrorIs(t, err, ErrInvalidPoolWeight) } + +func TestPoolMemberListDedupe(t *testing.T) { + unique := PoolMemberList{{Name: "a"}, {Name: "b", Weight: 3}} + got, repeats, err := unique.Dedupe("alb1") + require.NoError(t, err) + require.Empty(t, repeats) + require.Equal(t, unique, got) + + // the first occurrence wins and keeps its position; each repeat is reported once with + // the share the entries used to carry together + got, repeats, err = PoolMemberList{ + {Name: "a"}, {Name: "b", Weight: 2}, {Name: "a"}, {Name: "c"}, {Name: "a"}, + {Name: "b", Weight: 2}, {Name: "d", Weight: 4}, {Name: "d"}, + }.Dedupe("alb1") + require.NoError(t, err) + require.Equal(t, PoolMemberList{ + {Name: "a"}, {Name: "b", Weight: 2}, {Name: "c"}, {Name: "d", Weight: 4}, + }, got) + require.Equal(t, []PoolRepeat{ + {Name: "a", Count: 3, Weight: 3}, + {Name: "b", Count: 2, Weight: 4}, + {Name: "d", Count: 2, Weight: 5}, + }, repeats) + + _, _, err = PoolMemberList{{Name: "a", Weight: 2}, {Name: "b"}, {Name: "a", Weight: 3}}.Dedupe("alb1") + require.ErrorIs(t, err, ErrConflictingPoolWeights) + require.ErrorContains(t, err, `member "a" of alb "alb1"`) + + got, repeats, err = PoolMemberList(nil).Dedupe("alb1") + require.NoError(t, err) + require.Empty(t, got) + require.Empty(t, repeats) +} + +func TestInitializeDedupesPool(t *testing.T) { + o := &Options{MechanismName: "rr", Pool: Members("a", "a", "b")} + src := o.Pool + require.NoError(t, o.Initialize("alb1")) + require.Equal(t, Members("a", "b"), o.Pool) + require.Equal(t, Members("a", "a", "b"), src, "the configured list is not edited in place") + warning := o.PoolRepeatWarning("alb1") + require.Contains(t, warning, `alb "alb1"`) + require.Contains(t, warning, "{name: a, weight: 2}") + + // a second pass finds nothing repeated and must not forget what the first one found + require.NoError(t, o.Initialize("alb1")) + require.Equal(t, warning, o.PoolRepeatWarning("alb1")) + c := o.Clone() + require.Equal(t, o.PoolRepeats, c.PoolRepeats) + c.PoolRepeats[0].Name = "changed" + require.Equal(t, "a", o.PoolRepeats[0].Name) + + two := &Options{Pool: Members("a", "a", "b", "b", "b")} + require.NoError(t, two.Initialize("alb2")) + require.Contains(t, two.PoolRepeatWarning("alb2"), "{name: a, weight: 2}, {name: b, weight: 3}") + + require.Empty(t, (&Options{Pool: Members("a", "b")}).PoolRepeatWarning("alb3")) + conflict := &Options{Pool: PoolMemberList{{Name: "a", Weight: 2}, {Name: "a", Weight: 5}}} + require.ErrorIs(t, conflict.Initialize("alb4"), ErrConflictingPoolWeights) +} + +func TestPoolMemberBackup(t *testing.T) { + var l PoolMemberList + require.NoError(t, yaml.Unmarshal([]byte(` +- primary +- name: standby + backup: true +`), &l)) + require.Equal(t, PoolMemberList{{Name: "primary"}, {Name: "standby", Backup: true}}, l) + require.Equal(t, 0, l[0].Tier()) + require.Equal(t, BackupTier, l[1].Tier()) + b, err := yaml.Marshal(l) + require.NoError(t, err) + require.Contains(t, string(b), "- primary\n") + require.Contains(t, string(b), "backup: true") + var again PoolMemberList + require.NoError(t, yaml.Unmarshal(b, &again)) + require.Equal(t, l, again) + + require.False(t, l.AllBackups()) + require.False(t, PoolMemberList{}.AllBackups()) + standbys := PoolMemberList{{Name: "a", Backup: true}, {Name: "b", Backup: true}} + require.True(t, standbys.AllBackups()) + all := sets.New([]string{"a", "b"}) + require.ErrorIs(t, (&Options{Pool: standbys}).ValidatePool("alb1", all), ErrNoPrimaryPoolMember) + // discovered members are the primaries of a pool whose configured members all stand by + require.NoError(t, (&Options{Pool: standbys, Discovery: &DiscoveryOptions{}}).ValidatePool("alb1", all)) +} + +func TestPropagateHealthNeedsAPool(t *testing.T) { + o := &Options{MechanismName: names.MechanismUR, UserRouter: &ur.Options{}, PropagateHealth: true} + _, err := o.Validate() + require.ErrorIs(t, err, ErrPropagateHealthNoPool) + o = &Options{MechanismName: names.MechanismRR, PropagateHealth: true} + require.NoError(t, o.Initialize("alb1")) + _, err = o.Validate() + require.NoError(t, err) +} diff --git a/pkg/backends/alb/options/strategies.go b/pkg/backends/alb/options/strategies.go new file mode 100644 index 000000000..ef91ca23e --- /dev/null +++ b/pkg/backends/alb/options/strategies.go @@ -0,0 +1,309 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package options + +import ( + "errors" + "fmt" + "slices" + "strings" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/config/types" + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" +) + +const ( + // DefaultIPv6Prefix is how many leading bits of an IPv6 client address form its key. + DefaultIPv6Prefix = 64 + // LTSignalFirstWrite samples latency at the first byte written to an HTTP client. + LTSignalFirstWrite = "first_write" + // LTSignalConnect samples the time to connect to the member; tcp and tls listeners. + LTSignalConnect = "connect" + // LTSignalFirstByte samples the time to the member's first byte; tcp and tls listeners. + LTSignalFirstByte = "first_byte" + // LTSignalFirstReply samples the time to the member's first datagram; udp listeners. + LTSignalFirstReply = "first_reply" + // MaxConnectRetries bounds stream.connect_retries. + MaxConnectRetries = 10 + // DefaultPassiveFailures is how many consecutive connect failures eject a member. + DefaultPassiveFailures = 3 +) + +// ltSignals are the latency signals by the listener protocols that can sample them; the +// empty signal is each protocol's first +var ltSignals = map[string][]string{ + "http": {LTSignalFirstWrite}, + "tcp": {LTSignalConnect, LTSignalFirstByte}, + "tls": {LTSignalConnect, LTSignalFirstByte}, + "udp": {LTSignalFirstReply}, +} + +// LTSignalFor returns the latency signal in effect on a listener of the given protocol, or +// an error when the configured one cannot be sampled there. +func (o *Options) LTSignalFor(protocol string) (string, error) { + valid, ok := ltSignals[protocol] + if !ok { + valid = ltSignals["http"] + } + if o.LT.Signal == "" { + return valid[0], nil + } + if !slices.Contains(valid, o.LT.Signal) { + return "", fmt.Errorf("%w: %q on a %s listener (use %s)", ErrInvalidLTSignal, o.LT.Signal, + protocol, strings.Join(valid, " or ")) + } + return o.LT.Signal, nil +} + +// StreamOptions configure how an ALB balances tcp, tls and udp flows. +type StreamOptions struct { + // ConnectRetries is how many other pool members a tcp or tls connection may be offered + // when it cannot connect to the one it was given, all within the listener's connect + // timeout. The default, 0, refuses the connection instead. + ConnectRetries int `yaml:"connect_retries,omitempty"` + // PassiveHealth takes a member out of the pool when connections keep failing to reach it, + // without waiting for a health check. Off unless set. + PassiveHealth *PassiveHealthOptions `yaml:"passive_health,omitempty"` + // RaceWidth is how many pool members the race mechanism connects to at once, from 2 to + // 8. The default is every member, up to 4. + RaceWidth int `yaml:"race_width,omitempty"` +} + +// The bounds on how many members a race connects to at once, and on how many a mirror +// copies each flow to, the one that answers it included +const ( + DefaultRaceWidth = 4 + MinRaceWidth = 2 + MaxRaceWidth = 8 + MaxMirrorMembers = 8 +) + +// PassiveHealthOptions configure passive ejection. +type PassiveHealthOptions struct { + // Failures is how many consecutive failed connects eject a member; the default is 3. + Failures int `yaml:"failures,omitempty"` + // Eject is how long an ejected member stays out; the default is 30s. + Eject timeconv.Duration `yaml:"eject,omitempty"` + // MaxEjectedPercent is the most of the pool that may be ejected at once; the default is + // 50. The last live member is never ejected. + MaxEjectedPercent int `yaml:"max_ejected_percent,omitempty"` +} + +var ( + // ErrInvalidConnectRetries is returned for a stream.connect_retries outside 0-10. + ErrInvalidConnectRetries = errors.New("'stream.connect_retries' must be between 0 and 10") + // ErrInvalidRaceWidth is returned for a stream.race_width that is set outside 2-8. + ErrInvalidRaceWidth = errors.New("'stream.race_width' must be between 2 and 8") + // ErrRaceWidthOnlyForRace is returned when stream.race_width is set for another mechanism. + ErrRaceWidthOnlyForRace = errors.New("'stream.race_width' is only valid for mechanism 'race'") + // ErrTooManyMirrorMembers is returned for a mirror pool of more than MaxMirrorMembers. + ErrTooManyMirrorMembers = errors.New("mechanism 'mirror' supports a pool of at most 8 members") + // ErrStreamOptionsNeedOneMember is returned when a mechanism that commits a flow to + // several members is given the options of one that commits it to a single member. + ErrStreamOptionsNeedOneMember = errors.New("'stream.connect_retries' and 'stream.passive_health' " + + "are not valid for mechanisms 'race' and 'mirror'") + // ErrInvalidPassiveHealth is returned for a negative or out-of-range passive_health value. + ErrInvalidPassiveHealth = errors.New("'stream.passive_health' values cannot be negative, " + + "and 'max_ejected_percent' cannot exceed 100") +) + +func (o *StreamOptions) validate() error { + if o == nil { + return nil + } + if o.ConnectRetries < 0 || o.ConnectRetries > MaxConnectRetries { + return ErrInvalidConnectRetries + } + if p := o.PassiveHealth; p != nil && + (p.Failures < 0 || p.Eject < 0 || p.MaxEjectedPercent < 0 || p.MaxEjectedPercent > 100) { + return ErrInvalidPassiveHealth + } + if o.RaceWidth != 0 && (o.RaceWidth < MinRaceWidth || o.RaceWidth > MaxRaceWidth) { + return ErrInvalidRaceWidth + } + return nil +} + +// validateFor holds the block to what the mechanism it configures can use +func (o *StreamOptions) validateFor(mechanism string) error { + if o == nil { + return nil + } + race := mechanism == names.MechanismRace || mechanism == names.MechanismConnectRace + mirror := mechanism == names.MechanismMirror || mechanism == names.MechanismUDPMirror + if o.RaceWidth != 0 && !race { + return ErrRaceWidthOnlyForRace + } + if (race || mirror) && (o.ConnectRetries != 0 || o.PassiveHealth != nil) { + return ErrStreamOptionsNeedOneMember + } + return nil +} + +// Clone returns a deep copy of the options. +func (o *StreamOptions) Clone() *StreamOptions { + if o == nil { + return nil + } + c := *o + if o.PassiveHealth != nil { + p := *o.PassiveHealth + c.PassiveHealth = &p + } + return &c +} + +var ( + // ErrHRWOnlyForHRW is returned when the hrw block is set for another mechanism. + ErrHRWOnlyForHRW = errors.New("'hrw' options are only valid for mechanism 'hrw'") + // ErrLTOnlyForLT is returned when the lt block is set for another mechanism. + ErrLTOnlyForLT = errors.New("'lt' options are only valid for mechanism 'lt'") + // ErrInvalidIPv6Prefix is returned for an hrw.ipv6_prefix outside 1-128. + ErrInvalidIPv6Prefix = errors.New("'hrw.ipv6_prefix' must be between 1 and 128") + // ErrInvalidLTSignal is returned for an lt.signal the listener's protocol cannot sample. + ErrInvalidLTSignal = errors.New("value for 'lt.signal' is invalid") + // ErrInvalidLTDecay is returned for a negative lt.decay. + ErrInvalidLTDecay = errors.New("'lt.decay' cannot be negative") +) + +// HRWOptions configures the highest random weight mechanism. +type HRWOptions struct { + // Key is what a client's affinity follows: client_ip (the default), host, + // header:, cookie: or query:. + Key string `yaml:"key,omitempty"` + // IPv6Prefix is how many leading bits of an IPv6 client address form a client_ip key. + // The default, 64, keeps a client that rotates its privacy address on one member. + IPv6Prefix int `yaml:"ipv6_prefix,omitempty"` + // KeySource is Key, parsed + KeySource KeySource `yaml:"-"` +} + +// LTOptions configures the least time mechanism. +type LTOptions struct { + // StatusCodes are the response codes that count as a good answer, as bare codes or + // inclusive {start, end} ranges; any other records a latency penalty instead of a sample. + // The default is every code but 502, 503 and 504. + StatusCodes types.StatusRanges `yaml:"status_codes,omitempty"` + // Decay is the time constant of a member's latency average; the default is 10s. + Decay timeconv.Duration `yaml:"decay,omitempty"` + // Signal is what is timed. The default is the listener protocol's own: first_write on + // http, connect on tcp and tls (or first_byte), first_reply on udp. + Signal string `yaml:"signal,omitempty"` + // GoodCodes is StatusCodes, compiled + GoodCodes *types.StatusTable `yaml:"-"` +} + +// DefaultLTStatusCodes returns the response codes lt counts as a good answer when none are +// configured: all but the gateway failures. +func DefaultLTStatusCodes() types.StatusRanges { + return types.StatusRanges{{Start: 100, End: 501}, {Start: 505, End: 599}} +} + +func (o HRWOptions) isZero() bool { + return o.Key == "" && o.IPv6Prefix == 0 +} + +func (o LTOptions) isZero() bool { + return len(o.StatusCodes) == 0 && o.Decay == 0 && o.Signal == "" +} + +// initializeStrategies fills the defaults of the block that belongs to the mechanism +func (o *Options) initializeStrategies() error { + switch o.MechanismName { + case names.MechanismHRW, names.MechanismHighestRandomWeight: + ks, err := ParseKeySource(o.HRW.Key) + if err != nil { + return fmt.Errorf("hrw.key: %w", err) + } + o.HRW.KeySource = ks + if o.HRW.IPv6Prefix == 0 { + o.HRW.IPv6Prefix = DefaultIPv6Prefix + } + case names.MechanismLT, names.MechanismLeastTime: + codes := o.LT.StatusCodes + if len(codes) == 0 { + codes = DefaultLTStatusCodes() + } + o.LT.GoodCodes = codes.Compile() + } + if o.Stream != nil && o.Stream.PassiveHealth != nil && o.Stream.PassiveHealth.Failures == 0 { + o.Stream.PassiveHealth.Failures = DefaultPassiveFailures + } + return nil +} + +func (o *Options) validateStrategies() error { + if err := o.Stream.validate(); err != nil { + return err + } + if err := o.Stream.validateFor(o.MechanismName); err != nil { + return err + } + switch o.MechanismName { + case names.MechanismMirror, names.MechanismUDPMirror: + if len(o.Pool) > MaxMirrorMembers { + return ErrTooManyMirrorMembers + } + if !o.HRW.isZero() { + return ErrHRWOnlyForHRW + } + if !o.LT.isZero() { + return ErrLTOnlyForLT + } + case names.MechanismHRW, names.MechanismHighestRandomWeight: + if !o.LT.isZero() { + return ErrLTOnlyForLT + } + if _, err := ParseKeySource(o.HRW.Key); err != nil { + return fmt.Errorf("hrw.key: %w", err) + } + if o.HRW.IPv6Prefix < 0 || o.HRW.IPv6Prefix > 128 { + return ErrInvalidIPv6Prefix + } + case names.MechanismLT, names.MechanismLeastTime: + if !o.HRW.isZero() { + return ErrHRWOnlyForHRW + } + if err := o.LT.StatusCodes.Validate(); err != nil { + return fmt.Errorf("lt.status_codes: %w", err) + } + if o.LT.Decay < 0 { + return ErrInvalidLTDecay + } + known := false + for _, signals := range ltSignals { + known = known || slices.Contains(signals, o.LT.Signal) + } + if o.LT.Signal != "" && !known { + return fmt.Errorf("%w: %q", ErrInvalidLTSignal, o.LT.Signal) + } + default: + if !o.HRW.isZero() { + return ErrHRWOnlyForHRW + } + if !o.LT.isZero() { + return ErrLTOnlyForLT + } + } + return nil +} + +// LTDecay returns the configured lt.decay, or zero for the default. +func (o *Options) LTDecay() time.Duration { + return time.Duration(o.LT.Decay) +} diff --git a/pkg/backends/alb/options/strategies_test.go b/pkg/backends/alb/options/strategies_test.go new file mode 100644 index 000000000..0ca8c3881 --- /dev/null +++ b/pkg/backends/alb/options/strategies_test.go @@ -0,0 +1,258 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package options + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + "github.com/trickstercache/trickster/v2/pkg/config/types" + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" + + "github.com/stretchr/testify/require" + "go.yaml.in/yaml/v3" +) + +func TestParseKeySource(t *testing.T) { + for in, want := range map[string]KeySource{ + "": {Kind: KeyClientIP}, + "client_ip": {Kind: KeyClientIP}, + " host ": {Kind: KeyHost}, + "header:x-tenant": {Kind: KeyHeader, Name: "X-Tenant"}, + "header: X-Tenant ": {Kind: KeyHeader, Name: "X-Tenant"}, + "cookie:session": {Kind: KeyCookie, Name: "session"}, + "query:tenant": {Kind: KeyQuery, Name: "tenant"}, + "sni": {Kind: KeySNI}, + "user": {Kind: KeyUser}, + "proxy_tlv:0xEA": {Kind: KeyProxyTLV, TLV: 0xEA}, + "proxy_tlv: 5": {Kind: KeyProxyTLV, TLV: 5}, + } { + got, err := ParseKeySource(in) + require.NoError(t, err, in) + require.Equal(t, want, got, in) + } + for _, in := range []string{"snI:x", "header:", "cookie: ", "query:a=b", "header:two words", "cookie:a;b", "ip", "header", + "proxy_tlv:", "proxy_tlv:256", "proxy_tlv:-1", "proxy_tlv:authority", + } { + _, err := ParseKeySource(in) + require.ErrorIs(t, err, ErrInvalidKeySource, in) + } +} + +func load(t *testing.T, doc string) *Options { + t.Helper() + o := New() + require.NoError(t, yaml.Unmarshal([]byte(doc), o)) + return o +} + +func TestHRWOptions(t *testing.T) { + o := load(t, "mechanism: hrw\n") + require.NoError(t, o.Initialize("alb1")) + require.Equal(t, KeySource{Kind: KeyClientIP}, o.HRW.KeySource) + require.Equal(t, DefaultIPv6Prefix, o.HRW.IPv6Prefix) + _, err := o.Validate() + require.NoError(t, err) + + o = load(t, "mechanism: highest_random_weight\nhrw:\n key: header:x-tenant\n ipv6_prefix: 56\n") + require.NoError(t, o.Initialize("alb1")) + require.Equal(t, KeySource{Kind: KeyHeader, Name: "X-Tenant"}, o.HRW.KeySource) + require.Equal(t, 56, o.HRW.IPv6Prefix) + _, err = o.Validate() + require.NoError(t, err) + + require.ErrorIs(t, load(t, "mechanism: hrw\nhrw:\n key: port\n").Initialize("alb1"), ErrInvalidKeySource) + bad := load(t, "mechanism: hrw\nhrw:\n ipv6_prefix: 129\n") + require.NoError(t, bad.Initialize("alb1")) + _, err = bad.Validate() + require.ErrorIs(t, err, ErrInvalidIPv6Prefix) + unparsed := &Options{MechanismName: names.MechanismHRW, HRW: HRWOptions{Key: "nonsense"}} + _, err = unparsed.Validate() + require.ErrorIs(t, err, ErrInvalidKeySource) +} + +func TestLTOptions(t *testing.T) { + o := load(t, "mechanism: lt\n") + require.NoError(t, o.Initialize("alb1")) + require.Empty(t, o.LT.Signal, "an unset signal is resolved per listener protocol, not at load") + require.Zero(t, o.LTDecay()) + for _, code := range []int{200, 404, 500, 501, 505} { + require.True(t, o.LT.GoodCodes.Contains(code), code) + } + for _, code := range []int{502, 503, 504} { + require.False(t, o.LT.GoodCodes.Contains(code), code) + } + _, err := o.Validate() + require.NoError(t, err) + + o = load(t, "mechanism: least_time\nlt:\n status_codes: [{start: 200, end: 499}]\n decay: 30s\n signal: first_write\n") + require.NoError(t, o.Initialize("alb1")) + require.Equal(t, 30*time.Second, o.LTDecay()) + require.True(t, o.LT.GoodCodes.Contains(404)) + require.False(t, o.LT.GoodCodes.Contains(500)) + _, err = o.Validate() + require.NoError(t, err) + + c := o.Clone() + require.Equal(t, o.LT.StatusCodes, c.LT.StatusCodes) + require.NotSame(t, o.LT.GoodCodes, c.LT.GoodCodes) + c.LT.StatusCodes[0].End = 299 + require.Equal(t, 499, o.LT.StatusCodes[0].End, "the clone shares its ranges with the original") + + for doc, want := range map[string]error{ + "mechanism: lt\nlt:\n signal: last_byte\n": ErrInvalidLTSignal, + "mechanism: lt\nlt:\n decay: -5s\n": ErrInvalidLTDecay, + "mechanism: lt\nlt:\n status_codes: [{start: 500, end: 200}]\n": types.ErrInvalidStatusRange, + } { + bad := load(t, doc) + require.NoError(t, bad.Initialize("alb1")) + _, err := bad.Validate() + require.ErrorIs(t, err, want, doc) + } +} + +// a strategy's block is refused on any other mechanism, as output_format is outside tsm +func TestStrategyBlocksBelongToTheirMechanism(t *testing.T) { + for doc, want := range map[string]error{ + "mechanism: rr\nhrw:\n key: host\n": ErrHRWOnlyForHRW, + "mechanism: rr\nlt:\n decay: 5s\n": ErrLTOnlyForLT, + "mechanism: lt\nhrw:\n ipv6_prefix: 48\n": ErrHRWOnlyForHRW, + "mechanism: hrw\nlt:\n signal: first_write": ErrLTOnlyForLT, + "mechanism: p2c\nlt:\n status_codes: [200]": ErrLTOnlyForLT, + } { + o := load(t, doc) + require.NoError(t, o.Initialize("alb1")) + _, err := o.Validate() + require.ErrorIs(t, err, want, doc) + } + for _, mech := range []string{names.MechanismP2C, names.MechanismLC, names.MechanismRR} { + o := load(t, "mechanism: "+mech+"\n") + require.NoError(t, o.Initialize("alb1")) + _, err := o.Validate() + require.NoError(t, err, mech) + } +} + +func TestKeySourcePlanes(t *testing.T) { + for in, want := range map[string][5]bool{ + // readable on: a tcp or udp listener, a tls listener, one that accepts the PROXY + // protocol, an http listener, a native protocol listener + "client_ip": {true, true, true, true, true}, + "sni": {false, true, false, false, false}, + "proxy_tlv:0xEA": {false, false, true, false, false}, + "user": {false, false, false, false, true}, + "host": {false, false, false, true, false}, + "header:X-Tenant": {false, false, false, true, false}, + "cookie:session": {false, false, false, true, false}, + "query:tenant": {false, false, false, true, false}, + } { + ks, err := ParseKeySource(in) + require.NoError(t, err) + require.Equal(t, want, [5]bool{ + ks.OnStream(StreamListener{}), ks.OnStream(StreamListener{TLS: true}), + ks.OnStream(StreamListener{ProxyProtocol: true}), ks.OnHTTP(), ks.OnNative(), + }, in) + } +} + +func TestLTSignalFor(t *testing.T) { + o := &Options{} + for protocol, want := range map[string]string{ + "http": LTSignalFirstWrite, "tcp": LTSignalConnect, "tls": LTSignalConnect, + "udp": LTSignalFirstReply, "mysql": LTSignalFirstWrite, + } { + got, err := o.LTSignalFor(protocol) + require.NoError(t, err) + require.Equal(t, want, got, protocol) + } + o.LT.Signal = LTSignalFirstByte + got, err := o.LTSignalFor("tls") + require.NoError(t, err) + require.Equal(t, LTSignalFirstByte, got) + for _, protocol := range []string{"http", "udp"} { + _, err := o.LTSignalFor(protocol) + require.ErrorIs(t, err, ErrInvalidLTSignal, protocol) + } +} + +func TestStreamOptions(t *testing.T) { + o := load(t, "mechanism: rr\nstream:\n connect_retries: 2\n passive_health: {}\n") + require.NoError(t, o.Initialize("alb1")) + require.Equal(t, 2, o.Stream.ConnectRetries) + require.Equal(t, DefaultPassiveFailures, o.Stream.PassiveHealth.Failures, "an empty block turns it on") + _, err := o.Validate() + require.NoError(t, err) + + o = load(t, "mechanism: p2c\nstream:\n passive_health: {failures: 5, eject: 10s, max_ejected_percent: 25}\n") + require.NoError(t, o.Initialize("alb1")) + require.Equal(t, PassiveHealthOptions{Failures: 5, Eject: timeconv.Duration(10 * time.Second), MaxEjectedPercent: 25}, + *o.Stream.PassiveHealth) + c := o.Clone() + c.Stream.PassiveHealth.Failures = 9 + c.Stream.ConnectRetries = 4 + require.Equal(t, 5, o.Stream.PassiveHealth.Failures, "the clone shares its stream options") + require.Zero(t, o.Stream.ConnectRetries) + require.Nil(t, (*StreamOptions)(nil).Clone()) + + for doc, want := range map[string]error{ + "stream:\n connect_retries: -1\n": ErrInvalidConnectRetries, + "stream:\n connect_retries: 11\n": ErrInvalidConnectRetries, + "stream:\n passive_health: {failures: -1}\n": ErrInvalidPassiveHealth, + "stream:\n passive_health: {eject: -1s}\n": ErrInvalidPassiveHealth, + "stream:\n passive_health: {max_ejected_percent: 101}\n": ErrInvalidPassiveHealth, + "stream:\n passive_health: {max_ejected_percent: -5}\n": ErrInvalidPassiveHealth, + "lt:\n signal: last_byte\n": ErrInvalidLTSignal, + } { + mech := "mechanism: rr\n" + if want == ErrInvalidLTSignal { + mech = "mechanism: lt\n" + } + bad := load(t, mech+doc) + require.NoError(t, bad.Initialize("alb1")) + _, err := bad.Validate() + require.ErrorIs(t, err, want, doc) + } +} + +func TestSpreadMechanismOptions(t *testing.T) { + for name, test := range map[string]struct { + doc string + want error + }{ + "race": {"mechanism: race\n", nil}, + "race width": {"mechanism: connect_race\nstream:\n race_width: 3\n", nil}, + "race too narrow": {"mechanism: race\nstream:\n race_width: 1\n", ErrInvalidRaceWidth}, + "race too wide": {"mechanism: race\nstream:\n race_width: 9\n", ErrInvalidRaceWidth}, + "width without race": {"mechanism: rr\nstream:\n race_width: 2\n", ErrRaceWidthOnlyForRace}, + "race retries": {"mechanism: race\nstream:\n connect_retries: 1\n", ErrStreamOptionsNeedOneMember}, + "mirror ejection": {"mechanism: mirror\nstream:\n passive_health:\n failures: 2\n", ErrStreamOptionsNeedOneMember}, + "mirror": {"mechanism: udp_mirror\npool: [a, b, c, d, e, f, g, h]\n", nil}, + "mirror too wide": {"mechanism: mirror\npool: [a, b, c, d, e, f, g, h, i]\n", ErrTooManyMirrorMembers}, + "mirror hrw block": {"mechanism: mirror\nhrw:\n key: host\n", ErrHRWOnlyForHRW}, + "mirror lt block": {"mechanism: mirror\nlt:\n decay: 5s\n", ErrLTOnlyForLT}, + } { + o := load(t, test.doc) + require.NoError(t, o.Initialize("alb1"), name) + _, err := o.Validate() + if test.want == nil { + require.NoError(t, err, name) + continue + } + require.ErrorIs(t, err, test.want, name) + } + require.NoError(t, (*StreamOptions)(nil).validateFor("race")) +} diff --git a/pkg/backends/alb/pool/health.go b/pkg/backends/alb/pool/health.go deleted file mode 100644 index f2db80f62..000000000 --- a/pkg/backends/alb/pool/health.go +++ /dev/null @@ -1,89 +0,0 @@ -/* - * Copyright 2018 The Trickster Authors - * - * Licensed under the Apache License, Version 2.0 (the "License"); - * you may not use this file except in compliance with the License. - * You may obtain a copy of the License at - * - * http://www.apache.org/licenses/LICENSE-2.0 - * - * Unless required by applicable law or agreed to in writing, software - * distributed under the License is distributed on an "AS IS" BASIS, - * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. - * See the License for the specific language governing permissions and - * limitations under the License. - */ - -package pool - -import ( - "github.com/trickstercache/trickster/v2/pkg/observability/logging" - "github.com/trickstercache/trickster/v2/pkg/observability/logging/logger" - "github.com/trickstercache/trickster/v2/pkg/observability/metrics" - "github.com/trickstercache/trickster/v2/pkg/util/safego" -) - -// runWithRecover runs fn under defer-recover. A panic inside a refresh worker -// would otherwise kill the goroutine and leave the healthy-targets snapshot -// stale. The per-call re-filter in Targets() still produces correct dispatch -// even with a dead worker, so this is defense-in-depth for operator -// observability (healthy-count gauges, status-driven dashboards) rather than -// for correctness. Known panic scenarios this guards against: a Target whose -// hcStatus is mutated to nil concurrent with RefreshHealthy, and any future -// subscriber callback added to the status-update path that could panic. -func (p *pool) runWithRecover(worker string, fn func()) { - safego.Run(func(r any, stack []byte) { - logger.Error("alb pool refresh worker panic", logging.Pairs{ - "worker": worker, - "panic": r, - "stack": string(stack), - }) - metrics.ALBPoolRefreshPanicRecovered.WithLabelValues(worker).Inc() - }, fn) -} - -// listenStatusUpdates bridges target health-status notifications into refresh -// scheduling. It marks the pool list dirty for every received update and -// coalesces worker wakeups via scheduleRefresh so bursty changes cannot strand -// a stale healthy-target list. -func (p *pool) listenStatusUpdates() { - defer p.workers.Done() - for { - stop := false - p.runWithRecover("listenStatusUpdates", func() { - select { - case <-p.done: - stop = true - return - case <-p.statusCh: - p.scheduleRefresh() - } - }) - if stop { - return - } - } -} - -func (p *pool) checkHealth() { - defer p.workers.Done() - for { - stop := false - p.runWithRecover("checkHealth", func() { - select { - case <-p.done: - logger.Debug("stopping ALB pool", nil) - stop = true - return - case <-p.ch: // msg arrives whenever the healthy list must be rebuilt - // this coalesces bursts of updates into a single refresh - for p.refreshPending.Swap(false) { - p.RefreshHealthy() - } - } - }) - if stop { - return - } - } -} diff --git a/pkg/backends/alb/pool/health_test.go b/pkg/backends/alb/pool/health_test.go index b557ffff6..416c05125 100644 --- a/pkg/backends/alb/pool/health_test.go +++ b/pkg/backends/alb/pool/health_test.go @@ -24,31 +24,6 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" ) -func TestCheckHealth(t *testing.T) { - synctest.Test(t, func(t *testing.T) { - tgt := &Target{ - hcStatus: &healthcheck.Status{}, - } - - tgt.hcStatus.Set(healthcheck.StatusPassing) - - p := &pool{ch: make(chan bool, 1), done: make(chan struct{}), targets: []*Target{tgt}, healthyFloor: -1} - p.workers.Add(1) - go p.checkHealth() - defer p.Stop() - p.scheduleRefresh() - synctest.Wait() - - h := p.healthyHandlers.Load() - if h == nil { - t.Fatal("expected non-nil healthy list") - } - if got := len(*h); got != 1 { - t.Errorf("expected %d got %d", 1, got) - } - }) -} - func TestBurstUpdatesEvictFailingTarget(t *testing.T) { synctest.Test(t, func(t *testing.T) { st1 := &healthcheck.Status{} diff --git a/pkg/backends/alb/pool/main_test.go b/pkg/backends/alb/pool/main_test.go index 31dbfa340..cd8b3b802 100644 --- a/pkg/backends/alb/pool/main_test.go +++ b/pkg/backends/alb/pool/main_test.go @@ -22,10 +22,8 @@ import ( "go.uber.org/goleak" ) -// goleak guards against pool/healthcheck goroutines outliving the tests that -// created them. Pool spawns long-running listenStatusUpdates / checkHealth -// goroutines; a regression that fails to stop them would leak until process -// exit. -race won't catch leaks, only this will. +// goleak holds the pool to running no goroutines of its own: health transitions reach it +// synchronously, so a test that leaves one behind has reintroduced a worker. func TestMain(m *testing.M) { goleak.VerifyTestMain(m) } diff --git a/pkg/backends/alb/pool/pool.go b/pkg/backends/alb/pool/pool.go index ebbc11bee..1c2255bb1 100644 --- a/pkg/backends/alb/pool/pool.go +++ b/pkg/backends/alb/pool/pool.go @@ -13,24 +13,23 @@ * See the License for the specific language governing permissions and * limitations under the License. */ - // Package pool provides an application load balancer pool package pool import ( - "net/http" - "sync" "sync/atomic" - "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/observe" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/observability/logging" + "github.com/trickstercache/trickster/v2/pkg/observability/logging/logger" ) // Pool defines the interface for a load balancer pool type Pool interface { - // Targets returns the current set of dispatchable targets, re-filtered - // against each target's atomic hcStatus. This closes the race window - // between a status flip and the asynchronous healthy-list refresh, so it - // is the correct method for request dispatch. + // Targets returns the current set of dispatchable targets: the members whose health + // status meets the pool's floor. A health transition is reflected by the time the + // status change returns. The slice is shared; callers must not modify it. Targets() Targets // ConfiguredLen returns the number of pool members as configured, regardless // of current health. Mechanisms compare this against len(Targets()) to @@ -39,77 +38,82 @@ type Pool interface { // ConfiguredTargets returns the configured target topology in pool order. // The returned slice is a shallow copy and is safe for callers to retain. ConfiguredTargets() Targets - // SetHealthy seeds the pool's healthy set from a handler list. Intended - // for tests and bootstrap paths that don't drive status updates through - // healthcheck subscribers. - SetHealthy([]http.Handler) - // Stop stops the pool and its health checker goroutines. + // Core returns the protocol-neutral pool that selection strategies pick from. + Core() *lb.Pool + // Stop ends the pool's health subscriptions; its dispatchable set is then frozen. Stop() - // RefreshHealthy forces a refresh of the pool's healthy handlers list. + // RefreshHealthy forces a rebuild of the pool's dispatchable set. RefreshHealthy() } -// pool implements Pool +// pool implements Pool over the protocol-neutral core type pool struct { - targets Targets - healthyTargets atomic.Pointer[Targets] - liveTargets atomic.Pointer[Targets] - healthyHandlers atomic.Pointer[[]http.Handler] - refreshPending atomic.Bool // sticky dirty flag indicating healthyTargets must be rebuilt - healthyFloor int - done chan struct{} - statusCh chan bool // receives raw health status change notifications from targets - ch chan bool - mtx sync.Mutex - stopOnce sync.Once - workers sync.WaitGroup + targets Targets + core *lb.Pool + view atomic.Pointer[targetsView] } -// scheduleRefresh marks the healthy list as dirty and coalesces wakeups for -// the refresh worker. The pending flag preserves refresh intent even when -// bursty status changes saturate the channel. -func (p *pool) scheduleRefresh() { - p.refreshPending.Store(true) - select { - case p.ch <- true: - default: - } +// targetsView is the Targets form of one core snapshot, built once per snapshot +type targetsView struct { + snap *lb.Snapshot + targets Targets } -func (p *pool) RefreshHealthy() { - p.mtx.Lock() - defer p.mtx.Unlock() - hh := make([]http.Handler, len(p.targets)) - ht := make(Targets, len(p.targets)) +// observers hands each event to every observer in turn +type observers []lb.Observer + +func (o observers) Observe(ev lb.Event) { + for _, obs := range o { + obs.Observe(ev) + } +} - var k int - for _, t := range p.targets { +// New returns a new Pool. The observers are told of its events, the first snapshot included, +// which is published before New returns. +func New(targets Targets, healthyFloor int, extra ...lb.Observer) Pool { + p := &pool{targets: targets} + members := make([]*lb.Member, 0, len(targets)) + names := make(map[string]struct{}, len(targets)) + for _, t := range targets { + // a target with no health status can never be dispatched to if t == nil || t.hcStatus == nil { continue } - if int(t.hcStatus.Get()) >= p.healthyFloor { - hh[k] = t.handler - ht[k] = t - k++ + if t.name != "" { + if _, dup := names[t.name]; dup { + logger.Warn("alb pool member listed more than once; keeping the first", + logging.Pairs{"member": t.name}) + continue + } + names[t.name] = struct{}{} + } + if t.member == nil { + // a target assembled without a constructor + t.bind(nil) } + members = append(members, t.member) + } + observer := observe.Pool() + if len(extra) > 0 { + observer = append(observers{observer}, extra...) } - hh = hh[:k] - ht = ht[:k] - p.healthyHandlers.Store(&hh) - p.healthyTargets.Store(&ht) - lt := ht - p.liveTargets.Store(<) + core, err := lb.NewPool(members, healthyFloor, lb.PoolOptions{Observer: observer}) + if err != nil { + // unreachable: nil and repeated members were filtered above + logger.Error("alb pool could not be built", logging.Pairs{"error": err.Error()}) + core, _ = lb.NewPool(nil, healthyFloor) + } + p.core = core + p.Targets() + return p } -// snapshot returns the eventually-consistent healthy-targets snapshot. Snapshots -// can lag behind atomic status flips; only the refresh worker and internal tests -// should read this directly. Dispatch callers must use Targets(). -func (p *pool) snapshot() Targets { - t := p.healthyTargets.Load() - if t != nil { - return *t - } - return nil +func (p *pool) Core() *lb.Pool { + return p.core +} + +func (p *pool) RefreshHealthy() { + p.core.Refresh() } func (p *pool) ConfiguredLen() int { @@ -121,55 +125,26 @@ func (p *pool) ConfiguredTargets() Targets { } func (p *pool) Targets() Targets { - if lt := p.liveTargets.Load(); lt != nil && !p.refreshPending.Load() { - cached := *lt - allLive := true - for _, t := range cached { - if t == nil || t.hcStatus == nil || int(t.hcStatus.Get()) < p.healthyFloor { - allLive = false - break - } - } - if allLive { - return cached - } + snap := p.core.Snapshot() + if v := p.view.Load(); v != nil && v.snap == snap { + return v.targets } - hl := p.snapshot() - live := make(Targets, 0, len(hl)) - for _, t := range hl { - if t == nil || t.hcStatus == nil || int(t.hcStatus.Get()) < p.healthyFloor { - continue - } - live = append(live, t) - } - return live + return p.buildView(snap) } -func (p *pool) SetHealthy(h []http.Handler) { - p.healthyHandlers.Store(&h) - // Materialize parallel Targets each backed by a synthetic Passing status - // so dispatch-time re-checks against HealthyFloor won't reject them. - t := make(Targets, len(h)) - for i, hh := range h { - st := &healthcheck.Status{} - st.Set(healthcheck.StatusPassing) - t[i] = NewTarget(hh, st, nil) +// buildView runs once per snapshot, off the steady-state path. Racing builders store equal +// views; one that stores a superseded view is corrected by the next call. +func (p *pool) buildView(snap *lb.Snapshot) Targets { + targets := make(Targets, 0, len(snap.Members)) + for _, m := range snap.Members { + if t, ok := m.Value.(*Target); ok { + targets = append(targets, t) + } } - p.healthyTargets.Store(&t) - lt := t - p.liveTargets.Store(<) + p.view.Store(&targetsView{snap: snap, targets: targets}) + return targets } func (p *pool) Stop() { - p.stopOnce.Do(func() { - close(p.done) - for _, t := range p.targets { - if t != nil && t.hcStatus != nil { - t.hcStatus.UnregisterSubscriber(p.statusCh) - } - } - // Wait for refresh goroutines so SetHealthy after Stop cannot - // be overwritten by a late RefreshHealthy. - p.workers.Wait() - }) + p.core.Stop() } diff --git a/pkg/backends/alb/pool/pool_livetargets_bench_test.go b/pkg/backends/alb/pool/pool_livetargets_bench_test.go index f5bfbe5b8..7848e6b7e 100644 --- a/pkg/backends/alb/pool/pool_livetargets_bench_test.go +++ b/pkg/backends/alb/pool/pool_livetargets_bench_test.go @@ -31,8 +31,8 @@ func BenchmarkLiveTargets(b *testing.B) { st.Set(healthcheck.StatusPassing) targets[i] = NewTarget(http.NotFoundHandler(), st, nil) } - p := &pool{targets: targets, healthyFloor: 1} - p.RefreshHealthy() + p := New(targets, 1) + defer p.Stop() if got := len(p.Targets()); got != n { b.Fatalf("setup: expected %d healthy targets, got %d", n, got) } diff --git a/pkg/backends/alb/pool/pool_race_test.go b/pkg/backends/alb/pool/pool_race_test.go index d087c0118..428d1b899 100644 --- a/pkg/backends/alb/pool/pool_race_test.go +++ b/pkg/backends/alb/pool/pool_race_test.go @@ -18,31 +18,25 @@ package pool import ( "net/http" + "runtime" + "slices" "sync" "sync/atomic" "testing" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" - "github.com/trickstercache/trickster/v2/pkg/observability/metrics" - - "github.com/prometheus/client_golang/prometheus" - dto "github.com/prometheus/client_model/go" ) -// TestPoolRace consolidates the prior pool_stop_race, pool_stop_refresh_race, -// pool_refresh_panic, and pool_targets_nil_hcstatus tests. Each subtest names -// its axis so -race or panic output still points at the failing scenario. -// Subtests are run sequentially (no t.Parallel) because they share Prometheus -// metric vectors and the pool internal state machine. +// Subtests are named by axis so -race or panic output points at the failing scenario. func TestPoolRace(t *testing.T) { - t.Run("stop_double_close", testPoolStopConcurrentDoubleClose) - t.Run("stop_then_set_healthy_no_refresh_overwrite", testPoolStopThenSetHealthy) - t.Run("refresh_worker_survives_panic", testPoolRefreshWorkerSurvivesPanic) - t.Run("targets_cached_path_nil_hcstatus", testPoolTargetsCachedPathNilHCStatus) + t.Run("stop_concurrent", testPoolStopConcurrent) + t.Run("stop_during_transitions", testPoolStopDuringTransitions) + t.Run("transition_storm_converges", testPoolTransitionStormConverges) + t.Run("transition_racing_construction", testPoolTransitionRacingConstruction) + t.Run("target_without_constructor", testPoolTargetWithoutConstructor) } -// concurrent Stop must not panic on a closed channel. -func testPoolStopConcurrentDoubleClose(t *testing.T) { +func testPoolStopConcurrent(t *testing.T) { const iterations = 200 for i := range iterations { s := &healthcheck.Status{} @@ -71,113 +65,114 @@ func testPoolStopConcurrentDoubleClose(t *testing.T) { done.Wait() if panics.Load() > 0 { - t.Fatalf("iteration %d: Stop panicked on concurrent invocation (close of closed channel)", i) + t.Fatalf("iteration %d: Stop panicked on concurrent invocation", i) } } } -// scheduleRefresh queued by New() must not fire after Stop, or it would -// overwrite a subsequent SetHealthy and break Stop-then-SetHealthy callers. -func testPoolStopThenSetHealthy(t *testing.T) { - const iterations = 500 - for i := range iterations { - s := &healthcheck.Status{} - tgt := NewTarget(http.NotFoundHandler(), s, nil) - p := New(Targets{tgt}, 1) +// Stop racing in-flight transitions must neither deadlock nor panic, and nothing may be +// published once it has returned +func testPoolStopDuringTransitions(t *testing.T) { + for range 200 { + st := &healthcheck.Status{} + st.Set(healthcheck.StatusPassing) + p := New(Targets{NewTarget(http.NotFoundHandler(), st, nil)}, 1) + var wg sync.WaitGroup + wg.Go(func() { + for range 50 { + st.Set(healthcheck.StatusFailing) + st.Set(healthcheck.StatusPassing) + } + }) + runtime.Gosched() p.Stop() - h := []http.Handler{http.NotFoundHandler(), http.NotFoundHandler()} - p.SetHealthy(h) - if got := len(p.Targets()); got != 2 { - t.Fatalf("iteration %d: Targets: expected 2 got %d "+ - "(RefreshHealthy ran after Stop returned)", i, got) - } - if got := len(p.(*pool).snapshot()); got != 2 { - t.Fatalf("iteration %d: snapshot: expected 2 got %d", i, got) + frozen := p.Targets() + wg.Wait() + st.Set(healthcheck.StatusFailing) + st.Set(healthcheck.StatusPassing) + st.Set(healthcheck.StatusFailing) + if got := p.Targets(); !slices.Equal(got, frozen) { + t.Fatalf("a stopped pool republished: %d then %d targets", len(frozen), len(got)) } } } -// runWithRecover must absorb panics, increment the recover counter, and -// permit subsequent iterations to execute normally. -func testPoolRefreshWorkerSurvivesPanic(t *testing.T) { - p := &pool{} - - before := counterValue(t, metrics.ALBPoolRefreshPanicRecovered, "checkHealth") - - func() { - defer func() { - if r := recover(); r != nil { - t.Fatalf("panic escaped runWithRecover: %v", r) +// concurrent transitions on every member must leave the dispatchable set matching the final +// statuses, with no wait +func testPoolTransitionStormConverges(t *testing.T) { + const n = 16 + statuses := make([]*healthcheck.Status, n) + targets := make(Targets, n) + for i := range n { + statuses[i] = &healthcheck.Status{} + targets[i] = NewTarget(http.NotFoundHandler(), statuses[i], nil) + } + p := New(targets, 1) + defer p.Stop() + var wg sync.WaitGroup + for i := range n { + wg.Go(func() { + for range 200 { + statuses[i].Set(healthcheck.StatusPassing) + statuses[i].Set(healthcheck.StatusFailing) + } + if i%2 == 0 { + statuses[i].Set(healthcheck.StatusPassing) } - }() - p.runWithRecover("checkHealth", func() { - panic("simulated refresh panic") }) - }() - - after := counterValue(t, metrics.ALBPoolRefreshPanicRecovered, "checkHealth") - if got := after - before; got != 1 { - t.Fatalf("expected ALBPoolRefreshPanicRecovered{worker=checkHealth} +1, got +%v", got) } - - ran := false - p.runWithRecover("checkHealth", func() { - ran = true + // readers never see a nil or duplicated member while the storm runs + wg.Go(func() { + for range 2000 { + seen := make(map[*Target]bool, n) + for _, tgt := range p.Targets() { + if tgt == nil || seen[tgt] { + t.Error("Targets() returned a nil or repeated target") + return + } + seen[tgt] = true + } + } }) - if !ran { - t.Fatal("iteration after recovered panic did not execute") + wg.Wait() + var want Targets + for i := 0; i < n; i += 2 { + want = append(want, targets[i]) + } + if got := p.Targets(); !slices.Equal(got, want) { + t.Fatalf("after the storm: %d live targets, want the %d passing ones in pool order", len(got), len(want)) } +} - lbefore := counterValue(t, metrics.ALBPoolRefreshPanicRecovered, "listenStatusUpdates") - p.runWithRecover("listenStatusUpdates", func() { - panic("simulated status panic") - }) - lafter := counterValue(t, metrics.ALBPoolRefreshPanicRecovered, "listenStatusUpdates") - if got := lafter - lbefore; got != 1 { - t.Fatalf("expected ALBPoolRefreshPanicRecovered{worker=listenStatusUpdates} +1, got +%v", got) +// a transition that lands while the pool is being built is not lost +func testPoolTransitionRacingConstruction(t *testing.T) { + for range 500 { + st := &healthcheck.Status{} + tgt := NewTarget(http.NotFoundHandler(), st, nil) + var wg sync.WaitGroup + wg.Go(func() { st.Set(healthcheck.StatusPassing) }) + p := New(Targets{tgt}, 1) + wg.Wait() + if got := len(p.Targets()); got != 1 { + p.Stop() + t.Fatalf("a transition racing construction was lost: %d live targets", got) + } + p.Stop() } } -// Targets() cached path must tolerate a target whose hcStatus is nil, matching -// the snapshot path. Defense-in-depth for future callers that inject targets -// bypassing RefreshHealthy / SetHealthy. -func testPoolTargetsCachedPathNilHCStatus(t *testing.T) { +// a target assembled without a constructor has no member yet and may lack a status +func testPoolTargetWithoutConstructor(t *testing.T) { st := &healthcheck.Status{} st.Set(healthcheck.StatusPassing) - good := NewTarget(http.NotFoundHandler(), st, nil) + good := &Target{handler: http.NotFoundHandler(), hcStatus: st} bad := &Target{handler: http.NotFoundHandler()} - - p := &pool{ - targets: Targets{good}, - done: make(chan struct{}), - statusCh: make(chan bool, 1), - ch: make(chan bool, 1), - healthyFloor: 1, - } - - cached := Targets{good, bad} - p.liveTargets.Store(&cached) - - defer func() { - if r := recover(); r != nil { - t.Fatalf("Targets() panicked on nil hcStatus in cached path: %v", r) - } - }() - _ = p.Targets() -} - -func counterValue(t *testing.T, vec *prometheus.CounterVec, labels ...string) float64 { - t.Helper() - c, err := vec.GetMetricWithLabelValues(labels...) - if err != nil { - t.Fatalf("GetMetricWithLabelValues: %v", err) - } - var m dto.Metric - if err := c.Write(&m); err != nil { - t.Fatalf("write metric: %v", err) + p := New(Targets{good, bad, nil}, 1) + defer p.Stop() + if got := p.Targets(); len(got) != 1 || got[0] != good { + t.Fatalf("expected only the target with a status, got %d", len(got)) } - if m.Counter == nil || m.Counter.Value == nil { - return 0 + if p.ConfiguredLen() != 3 { + t.Errorf("configured length = %d", p.ConfiguredLen()) } - return *m.Counter.Value } diff --git a/pkg/backends/alb/pool/pool_stale_cache_repro_test.go b/pkg/backends/alb/pool/pool_stale_cache_repro_test.go index 9132a5959..ac4f711d4 100644 --- a/pkg/backends/alb/pool/pool_stale_cache_repro_test.go +++ b/pkg/backends/alb/pool/pool_stale_cache_repro_test.go @@ -28,11 +28,9 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" ) -// H1: simulate a dead refresh worker. Construct a pool but never start the -// refresh goroutines, mimicking what happens if listenStatusUpdates or -// checkHealth panicked and exited. Seed a populated cache and flip a target -// Failing. If Targets() still returns the Failing one, repro. -func TestRepro_H1_DeadRefreshWorker(t *testing.T) { +// H1: no background worker stands between a status flip and dispatch, so there is none to +// die: a Failing target is gone from Targets() by the time Set returns. +func TestRepro_H1_NoRefreshWorker(t *testing.T) { const n = 3 targets := make(Targets, n) statuses := make([]*healthcheck.Status, n) @@ -42,21 +40,8 @@ func TestRepro_H1_DeadRefreshWorker(t *testing.T) { statuses[i] = st targets[i] = NewTarget(http.NotFoundHandler(), st, nil) } - p := &pool{ - targets: targets, - done: make(chan struct{}), - statusCh: make(chan bool, 1), - ch: make(chan bool, 1), - healthyFloor: 1, - } - all := append(Targets(nil), targets...) - p.healthyTargets.Store(&all) - p.liveTargets.Store(&all) - hh := make([]http.Handler, n) - for i, tt := range targets { - hh[i] = tt.handler - } - p.healthyHandlers.Store(&hh) + p := New(targets, 1) + defer p.Stop() statuses[0].Set(healthcheck.StatusFailing) @@ -111,8 +96,8 @@ func TestRepro_H1b_RealPoolFlapStorm(t *testing.T) { }) } -// H2: saturate statusCh while refreshPending is being toggled, then flip -// target 0 Failing and assert it propagates. +// H2: flap every target concurrently, then flip target 0 Failing and assert +// it propagates. func TestRepro_H2_ChannelDrop(t *testing.T) { synctest.Test(t, func(t *testing.T) { const n = 10 diff --git a/pkg/backends/alb/pool/pool_test.go b/pkg/backends/alb/pool/pool_test.go index d20ed2167..98be8cb72 100644 --- a/pkg/backends/alb/pool/pool_test.go +++ b/pkg/backends/alb/pool/pool_test.go @@ -34,46 +34,54 @@ func TestNewTarget(t *testing.T) { func TestNewPool(t *testing.T) { s := &healthcheck.Status{} tgt := NewTarget(http.NotFoundHandler(), s, nil) - if tgt.hcStatus != s { - t.Error("unexpected mismatch") - } - p := New(Targets{tgt}, 1) if p == nil { - t.Error("expected non-nil") + t.Fatal("expected non-nil") } - - p2 := p.(*pool) - if got := len(p2.snapshot()); got != 0 { - t.Error("expected 0 healthy target", got) + defer p.Stop() + if got := len(p.Targets()); got != 0 { + t.Error("expected 0 healthy targets", got) } - - p.Stop() - - ht := Targets{tgt} - p2.healthyTargets.Store(&ht) - lt := ht - p2.liveTargets.Store(<) - - if got := len(p2.snapshot()); got != 1 { - t.Error("expected 1 healthy target", got) + if p.ConfiguredLen() != 1 || len(p.ConfiguredTargets()) != 1 { + t.Error("expected 1 configured target") + } + // a transition across the floor is dispatchable by the time Set returns + s.Set(healthcheck.StatusPassing) + if got := p.Targets(); len(got) != 1 || got[0] != tgt { + t.Error("expected the passing target", got) } } -func TestSetHealthyUpdatesTargets(t *testing.T) { +func TestTargetMemberAndAddr(t *testing.T) { s := &healthcheck.Status{} - tgt := NewTarget(http.NotFoundHandler(), s, nil) - p := New(Targets{tgt}, 1) - p.Stop() - - h := []http.Handler{http.NotFoundHandler(), http.NotFoundHandler()} - p.SetHealthy(h) - - if got := len(p.Targets()); got != 2 { - t.Errorf("Targets: expected 2 got %d", got) + tgt := NewWeightedTarget(http.NotFoundHandler(), s, nil, 3) + m := tgt.Member() + if m == nil || m.Value != tgt || m.Weight() != 3 || m.Health() != s { + t.Fatalf("member does not describe its target: %+v", m) } - if got := len(p.(*pool).snapshot()); got != 2 { - t.Errorf("snapshot: expected 2 got %d", got) + if tgt.Addr() != "" { + t.Errorf("a target without a backend has no address, got %q", tgt.Addr()) + } + // a replacement target keeps its predecessor's runtime stats; one built fresh does not + next := NewWeightedTarget(http.NotFoundHandler(), s, nil, 5).WithStatsOf(tgt) + if next.Member().Stats() != m.Stats() || next.Member().Weight() != 5 || next.Member().Value != next { + t.Error("stats were not carried to the replacement target") + } + kept := NewTarget(http.NotFoundHandler(), s, nil).WithStats(m.Stats()) + if kept.Member().Stats() != m.Stats() || kept.Member().Value != kept { + t.Error("stats kept from an earlier target were not adopted") + } + if NewTarget(http.NotFoundHandler(), s, nil).WithStats(nil).Member().Stats() == nil { + t.Error("a target with nothing to adopt lost its own stats") + } + if NewTarget(http.NotFoundHandler(), s, nil).WithStatsOf(nil).Member().Stats() == m.Stats() { + t.Error("unrelated targets share stats") + } + // a target without a health status is never dispatchable, and must not panic the pool + p := New(Targets{NewTarget(http.NotFoundHandler(), nil, nil), tgt}, -1) + defer p.Stop() + if got := p.Targets(); len(got) != 1 || got[0] != tgt { + t.Errorf("expected only the target with a status, got %d", len(got)) } } @@ -84,3 +92,45 @@ func TestStopIdempotent(t *testing.T) { p.Stop() p.Stop() // must not panic } + +// the dispatchable set is rebuilt only when a member crosses the floor, and reading it in +// steady state allocates nothing +func TestTargetsRebuiltOnlyOnFloorCrossing(t *testing.T) { + st1, st2 := &healthcheck.Status{}, &healthcheck.Status{} + st1.Set(healthcheck.StatusPassing) + st2.Set(healthcheck.StatusPassing) + p := New(Targets{NewTarget(http.NotFoundHandler(), st1, nil), + NewTarget(http.NotFoundHandler(), st2, nil)}, 0) + defer p.Stop() + before := p.Targets() + // Passing, Unchecked: all at or above a floor of 0 + st1.Set(healthcheck.StatusUnchecked) + st1.Set(healthcheck.StatusPassing) + if after := p.Targets(); len(after) != 2 || &after[0] != &before[0] { + t.Error("a transition that did not cross the floor rebuilt the dispatchable set") + } + if allocs := testing.AllocsPerRun(1000, func() { _ = p.Targets() }); allocs != 0 { + t.Errorf("Targets allocates %v in steady state", allocs) + } + st1.Set(healthcheck.StatusFailing) + if after := p.Targets(); len(after) != 1 { + t.Errorf("expected 1 live target, got %d", len(after)) + } + // the view is rebuilt once per snapshot, then served as is + if a, b := p.Targets(), p.Targets(); &a[0] != &b[0] { + t.Error("Targets rebuilt its view without a new snapshot") + } +} + +// a member listed twice is kept once: the core refuses duplicate names +func TestNewPoolKeepsFirstOfARepeatedName(t *testing.T) { + st := &healthcheck.Status{} + st.Set(healthcheck.StatusPassing) + first := &Target{handler: http.NotFoundHandler(), hcStatus: st, name: "a", weight: 1} + again := &Target{handler: http.NotFoundHandler(), hcStatus: st, name: "a", weight: 1} + p := New(Targets{first, again}, 0) + defer p.Stop() + if got := p.Targets(); len(got) != 1 || got[0] != first { + t.Errorf("expected only the first of a repeated name, got %d", len(got)) + } +} diff --git a/pkg/backends/alb/pool/pool_unavail_test.go b/pkg/backends/alb/pool/pool_unavail_test.go index 824196cad..3e8431fcd 100644 --- a/pkg/backends/alb/pool/pool_unavail_test.go +++ b/pkg/backends/alb/pool/pool_unavail_test.go @@ -19,43 +19,39 @@ package pool import ( "net/http" "testing" - "testing/synctest" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" ) -// Targets must drop a target whose status flipped below the floor after the -// cached snapshot was last refreshed. The internal snapshot keeps the stale -// view intact; Targets re-checks against the current atomic status to close -// the race window. -func TestLiveTargetsDropsStaleFailingTarget(t *testing.T) { - synctest.Test(t, func(t *testing.T) { - st1 := &healthcheck.Status{} - st2 := &healthcheck.Status{} - t1 := NewTarget(http.NotFoundHandler(), st1, nil) - t2 := NewTarget(http.NotFoundHandler(), st2, nil) - - p := New(Targets{t1, t2}, 1) - defer p.Stop() - st1.Set(healthcheck.StatusPassing) - st2.Set(healthcheck.StatusPassing) - synctest.Wait() - if got := len(p.Targets()); got != 2 { - t.Fatalf("setup: expected 2 healthy targets, got %d", got) - } - - // Pin the snapshot stale by stopping the pool's refresh goroutines, then - // flip t2 to Failing. - p.Stop() - st2.Set(healthcheck.StatusFailing) - - if got := len(p.(*pool).snapshot()); got != 2 { - t.Fatalf("snapshot: expected 2 (stale), got %d", got) - } - - live := p.Targets() - if len(live) != 1 || live[0] != t1 { - t.Fatalf("Targets: expected only t1, got %#v", live) - } - }) +// Targets drops a member the moment its status falls below the floor, with no wait for any +// background worker; a stopped pool no longer follows its members. +func TestTargetsDropsFailingTargetImmediately(t *testing.T) { + st1 := &healthcheck.Status{} + st2 := &healthcheck.Status{} + t1 := NewTarget(http.NotFoundHandler(), st1, nil) + t2 := NewTarget(http.NotFoundHandler(), st2, nil) + + p := New(Targets{t1, t2}, 1) + defer p.Stop() + st1.Set(healthcheck.StatusPassing) + st2.Set(healthcheck.StatusPassing) + if got := len(p.Targets()); got != 2 { + t.Fatalf("setup: expected 2 healthy targets, got %d", got) + } + + st2.Set(healthcheck.StatusFailing) + if live := p.Targets(); len(live) != 1 || live[0] != t1 { + t.Fatalf("Targets: expected only t1, got %#v", live) + } + + p.Stop() + st2.Set(healthcheck.StatusPassing) + st1.Set(healthcheck.StatusFailing) + if live := p.Targets(); len(live) != 1 || live[0] != t1 { + t.Fatalf("a stopped pool republished: %#v", live) + } + p.RefreshHealthy() + if live := p.Targets(); len(live) != 1 || live[0] != t1 { + t.Fatalf("a stopped pool refreshed: %#v", live) + } } diff --git a/pkg/backends/alb/pool/target.go b/pkg/backends/alb/pool/target.go index 322d53dec..bc67a59da 100644 --- a/pkg/backends/alb/pool/target.go +++ b/pkg/backends/alb/pool/target.go @@ -21,6 +21,8 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/proxy/hostnames" ) // Target defines an alb pool target @@ -31,38 +33,15 @@ type Target struct { name string group string weight int + tier int probed bool + dialable bool + addr string + member *lb.Member } type Targets []*Target -// New returns a new Pool -func New(targets Targets, healthyFloor int) Pool { - p := &pool{ - targets: targets, - done: make(chan struct{}), - statusCh: make(chan bool, 1), - ch: make(chan bool, 1), - healthyFloor: healthyFloor, - } - p.scheduleRefresh() - - for _, t := range targets { - if t == nil || t.hcStatus == nil { - continue - } - t.hcStatus.RegisterSubscriber(p.statusCh) - } - // populate the healthy snapshot synchronously so a pool installed by a - // runtime membership swap is dispatchable the moment SetPool returns, - // rather than 502ing until the async refresh worker's first pass - p.RefreshHealthy() - p.workers.Add(2) - go p.listenStatusUpdates() - go p.checkHealth() - return p -} - // NewTarget returns a new Target with the default weight of 1 func NewTarget(handler http.Handler, hcStatus *healthcheck.Status, backend backends.Backend, @@ -84,6 +63,10 @@ func NewWeightedTarget(handler http.Handler, hcStatus *healthcheck.Status, } if backend != nil { t.name, t.group = backendIdentity(backend) + if cfg := backend.Configuration(); cfg != nil { + t.addr = cfg.Host + t.dialable = t.addr != "" && !hostnames.Reserved(t.addr) + } if cfg := backend.Configuration(); cfg != nil && backends.HasOrigin(cfg.Provider) { // members with an origin are probed only when an active health @@ -95,9 +78,81 @@ func NewWeightedTarget(handler http.Handler, hcStatus *healthcheck.Status, if t.group == "" { t.group = t.name } + t.bind(nil) return t } +// bind builds the target's core member, which points back at the target +func (t *Target) bind(stats *lb.Stats) { + o := lb.MemberOptions{Name: t.name, Group: t.group, Weight: t.weight, Tier: t.tier, Stats: stats, Value: t} + if t.hcStatus != nil { + // a nil *Status must not become a non-nil Health + o.Health = t.hcStatus + } + t.member = lb.NewMember(o) +} + +// WithStats gives the target runtime stats kept from an earlier target of the same member, +// such as across a config reload; nil leaves its own. It returns the target. +func (t *Target) WithStats(stats *lb.Stats) *Target { + if stats != nil { + t.bind(stats) + } + return t +} + +// WithTier sets the target's failover tier: a pool dispatches to a tier above 0 only while no +// member of a lower tier is available. It returns the target. +func (t *Target) WithTier(tier int) *Target { + if tier = max(tier, 0); tier != t.tier { + t.tier = tier + t.bind(t.member.Stats()) + } + return t +} + +// Tier returns the target's failover tier, 0 unless it stands by for other members. +func (t *Target) Tier() int { + return t.tier +} + +// WithStatsOf carries prev's runtime stats over to the target that replaces it, so a member +// rebuilt in place, such as for a weight change, is not reset. It returns the target. +func (t *Target) WithStatsOf(prev *Target) *Target { + if prev != nil && prev.member != nil { + t.bind(prev.member.Stats()) + } + return t +} + +// Member returns the target's protocol-neutral pool member, whose Value is the target. +func (t *Target) Member() *lb.Member { + return t.member +} + +// Picker returns the balancer of a target whose backend is itself a load balancer, which is +// how a pool of pools is followed to a member that can be dialed. It is nil for any other +// target, and for a load balancer whose mechanism does not select one member. +func (t *Target) Picker() lb.Picker { + if pp, ok := t.backend.(lb.PickerProvider); ok { + return pp.Picker() + } + return nil +} + +// Dialable reports whether the target has an origin address that can be connected to. One +// without an address, or under the reserved .invalid domain, holds its share of a stream +// pool's flows and refuses them. It is decided once, when the target is built. +func (t *Target) Dialable() bool { + return t.dialable +} + +// Addr returns the host:port of the target's origin, captured when the target was built; +// empty when its backend has none. +func (t *Target) Addr() string { + return t.addr +} + // WithExternalHealth marks the target's health status as externally driven // (e.g., by discovery-provider readiness), so it counts as probed for // healthy-floor purposes even without an active health check interval. diff --git a/pkg/backends/alb/pool/target_test.go b/pkg/backends/alb/pool/target_test.go new file mode 100644 index 000000000..22e215f21 --- /dev/null +++ b/pkg/backends/alb/pool/target_test.go @@ -0,0 +1,141 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package pool + +import ( + "net/http" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" +) + +func testBackend(t *testing.T, name string, o *bo.Options) backends.Backend { + t.Helper() + if err := o.Initialize(name); err != nil { + t.Fatal(err) + } + b, err := backends.New(name, o, nil, http.NotFoundHandler(), nil) + if err != nil { + t.Fatal(err) + } + return b +} + +// balanced is a backend that is itself a load balancer +type balanced struct { + backends.Backend + picker lb.Picker +} + +func (b balanced) Picker() lb.Picker { return b.picker } + +func TestTargetDescribesItsBackend(t *testing.T) { + o := bo.New() + o.OriginURL = "tcp://10.0.0.1:5432" + b := testBackend(t, "pg1", o) + st := &healthcheck.Status{} + h := http.NotFoundHandler() + tgt := NewWeightedTarget(h, st, b, 0) + if tgt.Name() != "pg1" || tgt.ReplicaGroup() != "pg1" || tgt.Weight() != 1 { + t.Errorf("identity = %q %q %d", tgt.Name(), tgt.ReplicaGroup(), tgt.Weight()) + } + if tgt.Addr() != "10.0.0.1:5432" { + t.Errorf("addr = %q", tgt.Addr()) + } + if tgt.Backend() != b || tgt.HealthStatus() != st || tgt.Handler() == nil { + t.Error("the target lost its backend, status or handler") + } + m := tgt.Member() + if m.Name() != "pg1" || m.Group() != "pg1" || m.Weight() != 1 || m.Value != tgt { + t.Errorf("member = %+v", m) + } + if !tgt.Dialable() || tgt.Picker() != nil { + t.Error("an origin with an address is dialable, and is not itself balanced") + } + // a member under the reserved .invalid domain, or with no address, can never be dialed + gone := bo.New() + gone.OriginURL = "tcp://unresolved.kgw.invalid:1" + if NewTarget(h, st, testBackend(t, "gone", gone)).Dialable() { + t.Error("a member under .invalid reports as dialable") + } + if NewTarget(h, st, testBackend(t, "hostless", bo.New())).Dialable() || NewTarget(h, st, nil).Dialable() { + t.Error("a member with no address reports as dialable") + } + // a member that is itself balanced offers its picker, or nil when it has none to offer + picker := lb.NewBalancer(rr.New()) + if got := NewTarget(h, st, balanced{Backend: b, picker: picker}).Picker(); got != lb.Picker(picker) { + t.Errorf("nested picker = %v", got) + } + if NewTarget(h, st, balanced{Backend: b}).Picker() != nil { + t.Error("a balanced member with no picker offered one") + } + // with no health check interval the status can never leave Unchecked + if tgt.Probed() { + t.Error("a member without a health check interval reports as probed") + } + if !tgt.WithExternalHealth().Probed() { + t.Error("an externally driven member reports as unprobed") + } + + po := bo.New() + po.OriginURL = "http://10.0.0.2:9090" + po.HealthCheck = ho.New() + po.HealthCheck.Interval = timeconv.Duration(time.Second) + if !NewTarget(h, st, testBackend(t, "probed", po)).Probed() { + t.Error("a member with a health check interval reports as unprobed") + } +} + +type snapshotCounter struct{ snapshots int } + +func (c *snapshotCounter) Observe(ev lb.Event) { + if ev.Kind == lb.EventSnapshot { + c.snapshots++ + } +} + +func TestTargetTier(t *testing.T) { + primary := NewTarget(nil, healthcheck.NewStatus("p", "", "", healthcheck.StatusPassing, time.Time{}, nil), nil) + standby := NewTarget(nil, healthcheck.NewStatus("s", "", "", healthcheck.StatusPassing, time.Time{}, nil), nil) + stats := standby.Member().Stats() + if standby.WithTier(0) != standby || standby.Member().Stats() != stats { + t.Error("an unchanged tier rebuilt the member") + } + if standby.WithTier(1).Tier() != 1 || standby.Member().Tier() != 1 || standby.Member().Stats() != stats { + t.Errorf("tier = %d, member tier = %d", standby.Tier(), standby.Member().Tier()) + } + if standby.WithTier(-4).Tier() != 0 { + t.Error("a negative tier is tier 0") + } + standby.WithTier(1) + counter := &snapshotCounter{} + p := New(Targets{standby, primary}, 1, counter) + defer p.Stop() + if got := p.Targets(); len(got) != 1 || got[0] != primary || counter.snapshots != 1 { + t.Fatalf("targets = %v after %d snapshots", got, counter.snapshots) + } + primary.HealthStatus().Set(healthcheck.StatusFailing) + if got := p.Targets(); len(got) != 1 || got[0] != standby || counter.snapshots != 2 { + t.Fatalf("targets = %v after %d snapshots", got, counter.snapshots) + } +} diff --git a/pkg/backends/alb/pool_dedupe_test.go b/pkg/backends/alb/pool_dedupe_test.go new file mode 100644 index 000000000..fcd0150b7 --- /dev/null +++ b/pkg/backends/alb/pool_dedupe_test.go @@ -0,0 +1,114 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package alb + +import ( + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + + "github.com/stretchr/testify/require" + "go.yaml.in/yaml/v3" +) + +type atomicCounter struct{ hits atomic.Int64 } + +func (c *atomicCounter) ServeHTTP(w http.ResponseWriter, _ *http.Request) { + c.hits.Add(1) + w.WriteHeader(http.StatusOK) +} + +// startALBFromYAML loads an alb block as the config loader does, then starts its pool over +// counting members named by the pool +func startALBFromYAML(t *testing.T, initialize bool, albYAML string) (*Client, map[string]*atomicCounter) { + t.Helper() + a := &ao.Options{} + require.NoError(t, yaml.Unmarshal([]byte(albYAML), a)) + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = a + names := a.Pool.Names() + if initialize { + require.NoError(t, o.Initialize("dedupe")) + } + cl, err := NewClient("dedupe", o, nil, nil, nil, nil) + require.NoError(t, err) + c := cl.(*Client) + t.Cleanup(c.StopPool) + clients := backends.Backends{"dedupe": cl} + hits := make(map[string]*atomicCounter) + for _, name := range names { + if _, ok := hits[name]; ok { + continue + } + hits[name] = &atomicCounter{} + mo := bo.New() + mo.Name = name + b, err := backends.New(name, mo, nil, hits[name], nil) + require.NoError(t, err) + clients[name] = b + } + require.NoError(t, c.ValidateAndStartPool(clients, nil)) + return c, hits +} + +func serve(c *Client, n int) { + h := c.Handlers()[providers.ALB] + for range n { + h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "http://example.com/", nil)) + } +} + +func TestRepeatedPoolMembersAddNoShare(t *testing.T) { + c, hits := startALBFromYAML(t, true, "mechanism: rr\npool: [a, a, b]\n") + require.Equal(t, 2, c.Pool().ConfiguredLen()) + require.Contains(t, c.Configuration().ALBOptions.PoolRepeatWarning("dedupe"), "{name: a, weight: 2}") + serve(c, 12) + require.EqualValues(t, 6, hits["a"].hits.Load()) + require.EqualValues(t, 6, hits["b"].hits.Load()) +} + +func TestWeightIsTheOnlyWayToWeight(t *testing.T) { + c, hits := startALBFromYAML(t, true, "mechanism: rr\npool: [{name: a, weight: 3}, b]\n") + require.Empty(t, c.Configuration().ALBOptions.PoolRepeatWarning("dedupe")) + serve(c, 12) + require.EqualValues(t, 9, hits["a"].hits.Load()) + require.EqualValues(t, 3, hits["b"].hits.Load()) +} + +func TestFanoutDispatchesToARepeatedMemberOnce(t *testing.T) { + c, hits := startALBFromYAML(t, true, "mechanism: nlm\npool: [a, b, a]\n") + require.Equal(t, 2, c.Pool().ConfiguredLen()) + serve(c, 1) + require.EqualValues(t, 1, hits["a"].hits.Load()) + require.EqualValues(t, 1, hits["b"].hits.Load()) +} + +// options that never went through the loader still yield a pool of unique members +func TestStartPoolSkipsRepeatsInUninitializedOptions(t *testing.T) { + c, hits := startALBFromYAML(t, false, "mechanism: rr\npool: [a, b, a, a]\n") + require.Equal(t, 2, c.Pool().ConfiguredLen()) + serve(c, 10) + require.EqualValues(t, 5, hits["a"].hits.Load()) + require.EqualValues(t, 5, hits["b"].hits.Load()) +} diff --git a/pkg/backends/alb/statscarry.go b/pkg/backends/alb/statscarry.go new file mode 100644 index 000000000..65f9987db --- /dev/null +++ b/pkg/backends/alb/statscarry.go @@ -0,0 +1,61 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package alb + +import ( + "sync" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// carriedStats holds each ALB's static members' runtime stats by name, so that a config +// reload, which rebuilds every ALB, does not send a latency-ranking mechanism back to knowing +// nothing about its members. Discovered members keep theirs through their own manager. +var carriedStats = struct { + mtx sync.Mutex + byALB map[string]map[string]*lb.Stats +}{byALB: make(map[string]map[string]*lb.Stats)} + +// carryStats returns the stats a member of the named ALB had before, if any +func carryStats(albName, member string) *lb.Stats { + carriedStats.mtx.Lock() + defer carriedStats.mtx.Unlock() + return carriedStats.byALB[albName][member] +} + +// rememberStats replaces what is carried for the named ALB with its current members' stats +func rememberStats(albName string, members map[string]*lb.Stats) { + carriedStats.mtx.Lock() + defer carriedStats.mtx.Unlock() + if len(members) == 0 { + delete(carriedStats.byALB, albName) + return + } + carriedStats.byALB[albName] = members +} + +// forgetStatsExcept drops what is carried for ALBs that the running config no longer has +func forgetStatsExcept(clients backends.Backends) { + carriedStats.mtx.Lock() + defer carriedStats.mtx.Unlock() + for name := range carriedStats.byALB { + if _, ok := clients[name].(*Client); !ok { + delete(carriedStats.byALB, name) + } + } +} diff --git a/pkg/backends/alb/statscarry_test.go b/pkg/backends/alb/statscarry_test.go new file mode 100644 index 000000000..14eee8dc8 --- /dev/null +++ b/pkg/backends/alb/statscarry_test.go @@ -0,0 +1,105 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package alb + +import ( + "net/http" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/lb" + + "github.com/stretchr/testify/require" +) + +// generation builds one config generation: an ALB over members a and b, started +func generation(t *testing.T, albName, mechanism string, members ...string) (*Client, backends.Backends) { + t.Helper() + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = &ao.Options{MechanismName: mechanism, Pool: ao.Members(members...)} + require.NoError(t, o.Initialize(albName)) + cl, err := NewClient(albName, o, nil, nil, nil, nil) + require.NoError(t, err) + clients := backends.Backends{albName: cl} + for _, name := range members { + mo := bo.New() + mo.Name = name + b, err := backends.New(name, mo, nil, http.NotFoundHandler(), nil) + require.NoError(t, err) + clients[name] = b + } + require.NoError(t, StartALBPools(clients, nil)) + return cl.(*Client), clients +} + +func memberStats(c *Client) map[string]*lb.Stats { + out := make(map[string]*lb.Stats) + for _, tgt := range c.Pool().ConfiguredTargets() { + out[tgt.Name()] = tgt.Member().Stats() + } + return out +} + +// a reload rebuilds every ALB; a mechanism that ranks members by what it has learned of them +// must not start over each time +func TestStatsSurviveAReload(t *testing.T) { + first, _ := generation(t, "carry-alb", "lt", "a", "b") + before := memberStats(first) + first.StopPool() + + second, _ := generation(t, "carry-alb", "lt", "b", "c") + defer second.StopPool() + after := memberStats(second) + require.Same(t, before["b"], after["b"], "a member kept across the reload keeps its stats") + require.NotNil(t, after["c"]) + require.NotSame(t, before["a"], after["c"]) + + // a member dropped in one generation comes back fresh in a later one + second.StopPool() + third, _ := generation(t, "carry-alb", "lt", "a", "b") + defer third.StopPool() + require.NotSame(t, before["a"], memberStats(third)["a"]) + require.Same(t, before["b"], memberStats(third)["b"]) +} + +func TestStatsAreNotCarriedWhereTheyAreNotKept(t *testing.T) { + first, _ := generation(t, "untracked-carry-alb", "rr", "a") + before := memberStats(first) + first.StopPool() + second, _ := generation(t, "untracked-carry-alb", "rr", "a") + defer second.StopPool() + require.NotSame(t, before["a"], memberStats(second)["a"], "round robin keeps no stats to carry") + require.Nil(t, carryStats("untracked-carry-alb", "a")) +} + +// an ALB that leaves the config takes what was carried for it +func TestCarriedStatsAreForgottenWithTheirALB(t *testing.T) { + gone, _ := generation(t, "departing-alb", "lc", "a") + gone.StopPool() + require.NotNil(t, carryStats("departing-alb", "a")) + staying, _ := generation(t, "staying-alb", "lc", "a") + defer staying.StopPool() + require.Nil(t, carryStats("departing-alb", "a")) + require.NotNil(t, carryStats("staying-alb", "a")) + + rememberStats("staying-alb", nil) + require.Nil(t, carryStats("staying-alb", "a")) +} diff --git a/pkg/backends/alb/stream/bench_test.go b/pkg/backends/alb/stream/bench_test.go new file mode 100644 index 000000000..c1bc5d7da --- /dev/null +++ b/pkg/backends/alb/stream/bench_test.go @@ -0,0 +1,84 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package stream + +import ( + "strconv" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" +) + +const benchPoolSize = 8 + +// benchPool returns a started pool of n origin members; every other member is given weight +// when it is above 1 +func benchPool(b *testing.B, name string, n, weight int) backends.Backend { + b.Helper() + members := make([]spec, n) + for i := range members { + w := 1 + if weight > 1 && i%2 == 0 { + w = weight + } + members[i] = up(origin(b, name+"-"+strconv.Itoa(i), "10.0.0.1:"+strconv.Itoa(1000+i)), w) + } + return newALB(b, name, "rr", members...) +} + +// benchmarkPick measures what a connection costs the relay: the pick, and the route's reports +func benchmarkPick(b *testing.B, u l4.Upstream) { + b.Helper() + b.ReportAllocs() + for b.Loop() { + r, ok := u.Pick(l4.Flow{}) + if !ok { + b.Fatal("refused") + } + r.Dialed(0, nil) + r.Closed(nil) + } +} + +func BenchmarkPoolUpstreamPickUniform(b *testing.B) { + benchmarkPick(b, FromBackend(benchPool(b, "uniform", benchPoolSize, 1))) +} + +func BenchmarkPoolUpstreamPickWeighted(b *testing.B) { + benchmarkPick(b, FromBackend(benchPool(b, "weighted", benchPoolSize, 3))) +} + +func BenchmarkPoolUpstreamPickNested(b *testing.B) { + // a weighted outer pool over uniform inner pools, as a weighted rule over endpoints compiles to + outer := make([]spec, 2) + for i := range outer { + outer[i] = up(benchPool(b, "inner"+strconv.Itoa(i), benchPoolSize, 1), 1+2*i) + } + benchmarkPick(b, FromBackend(newALB(b, "outer", "rr", outer...))) +} + +func BenchmarkPoolUpstreamPickUniformParallel(b *testing.B) { + u := FromBackend(benchPool(b, "uniform", benchPoolSize, 1)) + b.ReportAllocs() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + if r, ok := u.Pick(l4.Flow{}); ok { + r.Closed(nil) + } + } + }) +} diff --git a/pkg/backends/alb/stream/helpers_test.go b/pkg/backends/alb/stream/helpers_test.go new file mode 100644 index 000000000..df6311663 --- /dev/null +++ b/pkg/backends/alb/stream/helpers_test.go @@ -0,0 +1,142 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package stream + +import ( + "maps" + "net/http" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/alb" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" +) + +// refusedKey counts the flows an upstream refused +const refusedKey = "" + +// spec is one pool member as a test describes it +type spec struct { + backend backends.Backend + weight int + status *healthcheck.Status +} + +func up(b backends.Backend, weight int) spec { + return spec{backend: b, weight: weight, + status: healthcheck.NewStatus(b.Name(), "", "", healthcheck.StatusPassing, time.Time{}, nil)} +} + +func down(b backends.Backend, weight int) spec { + return spec{backend: b, weight: weight, + status: healthcheck.NewStatus(b.Name(), "", "", healthcheck.StatusFailing, time.Time{}, nil)} +} + +// origin is a backend dialed at addr; an empty addr leaves it with nothing to dial +func origin(t testing.TB, name, addr string) backends.Backend { + t.Helper() + o := bo.New() + if addr != "" { + o.OriginURL = "tcp://" + addr + } + if err := o.Initialize(name); err != nil { + t.Fatal(err) + } + b, err := backends.New(name, o, nil, http.NotFoundHandler(), nil) + if err != nil { + t.Fatal(err) + } + return b +} + +// newALB builds and starts a real load balancer backend over the members, which may be +// load balancers themselves +func newALB(t testing.TB, name, mechanism string, members ...spec) backends.Backend { + t.Helper() + return newALBWith(t, name, mechanism, nil, members...) +} + +// newALBWith is newALB with the load balancer's options adjusted before it is built +func newALBWith(t testing.TB, name, mechanism string, adjust func(*ao.Options), members ...spec) backends.Backend { + t.Helper() + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = &ao.Options{MechanismName: mechanism, HealthyFloor: int(healthcheck.StatusUnchecked)} + if adjust != nil { + adjust(o.ALBOptions) + } + clients := backends.Backends{} + statuses := healthcheck.StatusLookup{} + for _, m := range members { + o.ALBOptions.Pool = append(o.ALBOptions.Pool, ao.PoolMember{Name: m.backend.Name(), Weight: m.weight}) + clients[m.backend.Name()] = m.backend + statuses[m.backend.Name()] = m.status + } + if err := o.Initialize(name); err != nil { + t.Fatal(err) + } + cl, err := alb.NewClient(name, o, nil, nil, nil, nil) + if err != nil { + t.Fatal(err) + } + clients[name] = cl + c := cl.(*alb.Client) + if err := c.ValidateAndStartPool(clients, statuses); err != nil { + t.Fatal(err) + } + t.Cleanup(c.StopPool) + return cl +} + +func pickN(u l4.Upstream, n int) []string { + seq := make([]string, n) + for i := range seq { + if r, ok := u.Pick(l4.Flow{}); ok { + seq[i] = r.Addr() + r.Dialed(time.Millisecond, nil) + r.Closed(nil) + } + } + return seq +} + +func tally(seq []string) map[string]int { + counts := make(map[string]int) + for _, addr := range seq { + counts[addr]++ + } + return counts +} + +// assertEveryWindowExact fails unless every run of total consecutive selections matches want, +// wherever the run starts: apportionment is pinned, the rotation's phase is not +func assertEveryWindowExact(t *testing.T, seq []string, want map[string]int) { + t.Helper() + var total int + for _, n := range want { + total += n + } + for start := 0; start+total <= len(seq); start++ { + if got := tally(seq[start : start+total]); !maps.Equal(got, want) { + t.Fatalf("selections %d-%d apportioned %v, want %v", start, start+total-1, got, want) + } + } +} diff --git a/pkg/backends/alb/stream/relay_test.go b/pkg/backends/alb/stream/relay_test.go new file mode 100644 index 000000000..d17c8b403 --- /dev/null +++ b/pkg/backends/alb/stream/relay_test.go @@ -0,0 +1,177 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package stream + +import ( + "bufio" + "net" + "strconv" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" +) + +// echoTCP answers each line with prefix + line +func echoTCP(t *testing.T, prefix string) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = ln.Close() }) + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go func() { + defer conn.Close() + r := bufio.NewReader(conn) + for { + line, err := r.ReadString('\n') + if err != nil { + return + } + _, _ = conn.Write([]byte(prefix + strings.TrimSpace(line) + "\n")) + } + }() + } + }() + return ln.Addr().String() +} + +func echoUDP(t *testing.T, prefix string) string { + t.Helper() + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = pc.Close() }) + go func() { + buf := make([]byte, 1500) + for { + n, from, err := pc.ReadFrom(buf) + if err != nil { + return + } + _, _ = pc.WriteTo([]byte(prefix+string(buf[:n])), from) + } + }() + return pc.LocalAddr().String() +} + +func tableOf(t *testing.T, u l4.Upstream) *l4.Table { + t.Helper() + tbl := l4.NewTable() + if err := tbl.Add("", u); err != nil { + t.Fatal(err) + } + return tbl +} + +// each connection commits to one member: live members answer in turn, and the share of one +// that cannot be dialed is refused rather than handed to a sibling +func TestTCPRelayBalancesAcrossThePool(t *testing.T) { + deadLn, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + dead := deadLn.Addr().String() + _ = deadLn.Close() + pool := newALB(t, "alb", "rr", + up(origin(t, "dead", dead), 1), up(origin(t, "a", echoTCP(t, "a:")), 1), up(origin(t, "b", echoTCP(t, "b:")), 1)) + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + srv := l4.NewServer("test", l4.ProtocolTCP, &l4.Config{Table: tableOf(t, FromBackend(pool))}) + go func() { _ = srv.Serve(ln) }() + t.Cleanup(func() { _ = srv.Close() }) + + seen := make(map[string]int) + var refused int + for range 6 { + conn, err := net.DialTimeout("tcp", ln.Addr().String(), time.Second) + if err != nil { + t.Fatal(err) + } + _ = conn.SetDeadline(time.Now().Add(3 * time.Second)) + _, _ = conn.Write([]byte("x\n")) + reply, err := bufio.NewReader(conn).ReadString('\n') + _ = conn.Close() + if err != nil { + refused++ + continue + } + seen[strings.TrimSpace(reply)]++ + } + if seen["a:x"] != 2 || seen["b:x"] != 2 || refused != 2 { + t.Errorf("replies = %v, refused %d; want each member its share", seen, refused) + } +} + +// a session commits to one member for life; a member under the reserved .invalid domain +// holds its share of sessions and refuses them +func TestUDPRelayBalancesAcrossThePool(t *testing.T) { + pool := newALB(t, "alb", "rr", + up(origin(t, "a", echoUDP(t, "a:")), 1), up(origin(t, "b", echoUDP(t, "b:")), 1), + up(origin(t, "gone", "unresolved.kgw.invalid:1"), 2)) + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + srv := l4.NewPacketServer("test", &l4.Config{Table: tableOf(t, FromBackend(pool))}) + go func() { _ = srv.Serve(pc) }() + t.Cleanup(func() { _ = srv.Close() }) + + seen := make(map[string]int) + var refused int + for i := range 8 { + conn, err := net.Dial("udp", pc.LocalAddr().String()) + if err != nil { + t.Fatal(err) + } + var first string + for j := range 2 { + _, _ = conn.Write([]byte(strconv.Itoa(j))) + _ = conn.SetReadDeadline(time.Now().Add(250 * time.Millisecond)) + buf := make([]byte, 16) + n, err := conn.Read(buf) + if err != nil { + break + } + member := string(buf[:2]) + if j == 0 { + first = member + } else if member != first { + t.Errorf("session %d moved from %s to %s", i, first, member) + } + _ = n + } + _ = conn.Close() + if first == "" { + refused++ + continue + } + seen[first]++ + } + if seen["a:"] != 2 || seen["b:"] != 2 || refused != 4 { + t.Errorf("sessions = %v, refused %d; want 2, 2 and the refusing member's 4", seen, refused) + } +} diff --git a/pkg/backends/alb/stream/spread.go b/pkg/backends/alb/stream/spread.go new file mode 100644 index 000000000..8a8c9530e --- /dev/null +++ b/pkg/backends/alb/stream/spread.go @@ -0,0 +1,156 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package stream + +import ( + "errors" + "sync/atomic" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" +) + +// spreader is a load balancer whose mechanism commits one flow to several members at once +type spreader interface { + Spread() types.Spread + Pool() pool.Pool +} + +// spreadUpstream commits a flow to several of a pool's available members: a connect race for +// tcp and tls, a copy of every datagram for udp. Members that cannot be dialed are passed over. +type spreadUpstream struct { + lb spreader + kind types.Spread + width int + next atomic.Uint64 + series seriesCache +} + +func newSpread(sp spreader, cfg *bo.Options) *spreadUpstream { + u := &spreadUpstream{lb: sp, kind: sp.Spread(), width: ao.DefaultRaceWidth} + if cfg != nil && cfg.ALBOptions != nil && cfg.ALBOptions.Stream != nil && cfg.ALBOptions.Stream.RaceWidth > 0 { + u.width = cfg.ALBOptions.Stream.RaceWidth + } + return u +} + +// dialable returns the pool's available members that have an address to connect to. The +// pool is read on every flow, since discovery and reloads replace it. +func (u *spreadUpstream) dialable() pool.Targets { + p := u.lb.Pool() + if p == nil { + return nil + } + live := p.Targets() + for i, t := range live { + if t.Dialable() { + continue + } + // rare: copy the shared slice only when something must be left out of it + out := append(make(pool.Targets, 0, len(live)-1), live[:i]...) + for _, rest := range live[i+1:] { + if rest.Dialable() { + out = append(out, rest) + } + } + return out + } + return live +} + +func (u *spreadUpstream) routeTo(t *pool.Target, f l4.Flow) *spreadRoute { + return &spreadRoute{addr: t.Addr(), series: u.series.seriesFor(t.Member(), f, t.Name())} +} + +// Pick returns the flow's one answering route: the first available member. A race is asked +// for its routes through Race instead. +func (u *spreadUpstream) Pick(f l4.Flow) (l4.Route, bool) { + live := u.dialable() + if len(live) == 0 { + return nil, false + } + return u.routeTo(live[0], f), true +} + +// Race returns up to the configured width of members to connect to at once, starting one +// member further on with each flow so a pool wider than the race shares its connects. +func (u *spreadUpstream) Race(f l4.Flow) []l4.Route { + live := u.dialable() + if len(live) == 0 { + return nil + } + if u.kind != types.SpreadRace { + return []l4.Route{u.routeTo(live[0], f)} + } + n := min(len(live), u.width) + start := int(u.next.Add(1) % uint64(len(live))) // #nosec G115 -- bounded by the pool's length + routes := make([]l4.Route, n) + for i := range routes { + routes[i] = u.routeTo(live[(start+i)%len(live)], f) + } + return routes +} + +// Mirror returns the members beside the first, which answers the flow, to copy it to. +func (u *spreadUpstream) Mirror(f l4.Flow, _ l4.Route) []l4.Route { + live := u.dialable() + if u.kind != types.SpreadMirror || len(live) < 2 { + return nil + } + routes := make([]l4.Route, len(live)-1) + for i, t := range live[1:] { + routes[i] = u.routeTo(t, f) + } + return routes +} + +// spreadRoute is one member's part in a flow that was committed to several +type spreadRoute struct { + addr string + series *memberSeries +} + +func (r *spreadRoute) Addr() string { return r.addr } + +// Final is true: the flow's other members were tried along with this one +func (r *spreadRoute) Final() bool { return true } + +func (r *spreadRoute) Dialed(d time.Duration, err error) { + switch { + case err == nil: + r.series.connect.Observe(d.Seconds()) + r.series.active.Inc() + case errors.Is(err, l4.ErrAbandoned): + default: + r.series.dialFailed.Inc() + } +} + +func (r *spreadRoute) FirstByte() {} + +func (r *spreadRoute) Closed(err error) { + r.series.active.Dec() + if err != nil { + r.series.unreachable.Inc() + return + } + r.series.proxied.Inc() +} diff --git a/pkg/backends/alb/stream/spread_test.go b/pkg/backends/alb/stream/spread_test.go new file mode 100644 index 000000000..2309415c4 --- /dev/null +++ b/pkg/backends/alb/stream/spread_test.go @@ -0,0 +1,160 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package stream + +import ( + "slices" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" + + "github.com/prometheus/client_golang/prometheus/testutil" +) + +func addrsOf(routes []l4.Route) []string { + out := make([]string, len(routes)) + for i, r := range routes { + out[i] = r.Addr() + } + return out +} + +func TestRaceConnectsToSeveralMembers(t *testing.T) { + m := origins(t, 5) + unroutable := origin(t, "no-address", "") + sick := down(m[4], 1) + u := FromBackend(newALB(t, "race-alb", "race", up(m[0], 1), up(unroutable, 1), up(m[1], 1), up(m[2], 1), up(m[3], 1), sick)) + racer, ok := u.(l4.Racer) + if !ok { + t.Fatal("a race mechanism does not race") + } + f := clientFlow(l4.ProtocolTCP, "198.51.100.1:1000", "") + // the default width of 4 covers the four members that are up and can be dialed, and each + // flow starts one member further on + first, second := addrsOf(racer.Race(f)), addrsOf(racer.Race(f)) + if len(first) != ao.DefaultRaceWidth || slices.Contains(first, "10.0.0.5:9000") { + t.Fatalf("race = %v", first) + } + if first[1] != second[0] { + t.Errorf("consecutive races start at %s and %s", first[0], second[0]) + } + if r, ok := u.Pick(f); !ok || r.Addr() != "10.0.0.1:9000" || !r.Final() { + t.Errorf("pick = %v, %v", r, ok) + } + if _, mirrors := u.(l4.Mirrorer); mirrors && len(u.(l4.Mirrorer).Mirror(f, nil)) != 0 { + t.Error("a race mirrors") + } + + narrow := func(o *ao.Options) { o.Stream = &ao.StreamOptions{RaceWidth: 2} } + u = FromBackend(newALBWith(t, "narrow-alb", "connect_race", narrow, up(m[0], 1), up(m[1], 1), up(m[2], 1))) + if got := u.(l4.Racer).Race(f); len(got) != 2 { + t.Errorf("a race of width 2 connects to %d members", len(got)) + } + + empty := FromBackend(newALB(t, "empty-alb", "race", down(m[0], 1), up(unroutable, 1))) + if got := empty.(l4.Racer).Race(f); len(got) != 0 { + t.Errorf("a pool with nothing to dial races %v", addrsOf(got)) + } + if _, ok := empty.Pick(f); ok { + t.Error("a pool with nothing to dial picked a member") + } +} + +func TestMirrorCopiesToEveryOtherMember(t *testing.T) { + m := origins(t, 4) + primary := up(m[0], 1) + u := FromBackend(newALB(t, "mirror-alb", "mirror", primary, up(m[1], 1), down(m[2], 1), up(m[3], 1))) + f := clientFlow(l4.ProtocolUDP, "198.51.100.1:1000", "") + r, ok := u.Pick(f) + if !ok || r.Addr() != "10.0.0.1:9000" { + t.Fatalf("pick = %v, %v", r, ok) + } + if got := addrsOf(u.(l4.Mirrorer).Mirror(f, r)); !slices.Equal(got, []string{"10.0.0.2:9000", "10.0.0.4:9000"}) { + t.Errorf("mirrors = %v", got) + } + // a mirror has one route to offer a connection, so nothing to race + if got := u.(l4.Racer).Race(f); len(got) != 1 || got[0].Addr() != "10.0.0.1:9000" { + t.Errorf("race = %v", addrsOf(got)) + } + // the next member answers while the first is down + primary.status.Set(healthcheck.StatusFailing) + if r, ok := u.Pick(f); !ok || r.Addr() != "10.0.0.2:9000" { + t.Errorf("pick with the first member down = %v, %v", r, ok) + } + alone := FromBackend(newALB(t, "alone-alb", "udp_mirror", up(m[0], 1))) + if got := alone.(l4.Mirrorer).Mirror(f, nil); got != nil { + t.Errorf("a pool of one mirrors to %v", addrsOf(got)) + } +} + +func TestSpreadRoutesCountTheirMember(t *testing.T) { + m := origins(t, 2) + u := FromBackend(newALB(t, "counted-alb", "race", up(m[0], 1), up(m[1], 1))) + f := clientFlow(l4.ProtocolTCP, "198.51.100.1:1000", "") + f.Listener = "spread-counts" + count := func(member, result string) float64 { + return testutil.ToFloat64(metrics.ProxyStreamMemberConnections.WithLabelValues(f.Listener, f.Protocol, member, result)) + } + active := func(member string) float64 { + return testutil.ToFloat64(metrics.ProxyStreamMemberActiveConnections.WithLabelValues(f.Listener, f.Protocol, member)) + } + routes := u.(l4.Racer).Race(f) + won, lost := routes[0], routes[1] + won.Dialed(time.Millisecond, nil) + won.FirstByte() + lost.Dialed(0, l4.ErrAbandoned) + if active("m1") != 1 || active("m0") != 0 { + t.Errorf("active = m0 %v, m1 %v", active("m0"), active("m1")) + } + won.Closed(nil) + if active("m1") != 0 || count("m1", ResultProxied) != 1 || count("m0", ResultDialFailed) != 0 { + t.Error("a relayed connection was not counted for the member that won it alone") + } + failed := u.(l4.Racer).Race(f)[0] + failed.Dialed(time.Millisecond, errRefused) + if count("m0", ResultDialFailed) != 1 { + t.Error("a failed connect was not counted") + } + unreachable := u.(l4.Racer).Race(f)[0] + unreachable.Dialed(time.Millisecond, nil) + unreachable.Closed(errRefused) + if count("m1", ResultUnreachable) != 1 { + t.Error("an unreachable member was not counted") + } +} + +type poolless struct{} + +func (poolless) Spread() types.Spread { return types.SpreadRace } + +func (poolless) Pool() pool.Pool { return nil } + +func TestSpreadBeforeItsPoolStarts(t *testing.T) { + u := newSpread(poolless{}, nil) + if _, ok := u.Pick(l4.Flow{}); ok || len(u.Race(l4.Flow{})) != 0 || u.Mirror(l4.Flow{}, nil) != nil { + t.Error("a load balancer with no pool committed a flow") + } + if ao.MaxMirrorMembers != l4.MaxUDPMirrors+1 { + t.Error("the pool a mirror may have is not what the relay will copy a flow to") + } +} diff --git a/pkg/backends/alb/stream/strategies_test.go b/pkg/backends/alb/stream/strategies_test.go new file mode 100644 index 000000000..fff89a0c7 --- /dev/null +++ b/pkg/backends/alb/stream/strategies_test.go @@ -0,0 +1,660 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package stream + +import ( + "errors" + "net/netip" + "strconv" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/alb" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" + + "github.com/prometheus/client_golang/prometheus/testutil" +) + +var errRefused = errors.New("connection refused") + +func origins(t testing.TB, n int) []backends.Backend { + t.Helper() + out := make([]backends.Backend, n) + for i := range out { + out[i] = origin(t, "m"+strconv.Itoa(i), "10.0.0."+strconv.Itoa(i+1)+":9000") + } + return out +} + +func clientFlow(protocol, client, serverName string) l4.Flow { + return l4.Flow{Listener: "test", Protocol: protocol, Client: netip.MustParseAddrPort(client), ServerName: serverName} +} + +func statsOf(b backends.Backend) map[string]*lb.Stats { + out := make(map[string]*lb.Stats) + for _, t := range b.(*alb.Client).Pool().ConfiguredTargets() { + out[t.Addr()] = t.Member().Stats() + } + return out +} + +// every strategy balances connections as it balances requests: only live members, all of +// them reachable, nothing left in flight, and a refusal once none is live +func TestEveryStrategyServesAStreamListener(t *testing.T) { + for _, mech := range []string{"rr", "p2c", "lc", "lt", "hrw"} { + t.Run(mech, func(t *testing.T) { + m := origins(t, 4) + members := []spec{up(m[0], 1), up(m[1], 2), up(m[2], 1), down(m[3], 5)} + pool := newALB(t, "alb", mech, members...) + u := FromBackend(pool) + seen := make(map[string]int) + var held []l4.Route + for i := range 400 { + r, ok := u.Pick(clientFlow(l4.ProtocolTCP, "198.51.100."+strconv.Itoa(i%250)+":"+strconv.Itoa(1024+i), "")) + if !ok { + t.Fatal("refused although members are live") + } + seen[r.Addr()]++ + r.Dialed(time.Millisecond, nil) + held = append(held, r) + if len(held) > 16 { + held[0].Closed(nil) + held = held[1:] + } + } + for _, r := range held { + r.Closed(nil) + } + if seen["10.0.0.4:9000"] != 0 { + t.Errorf("%d connections went to a member that is down", seen["10.0.0.4:9000"]) + } + for _, addr := range []string{"10.0.0.1:9000", "10.0.0.2:9000", "10.0.0.3:9000"} { + if seen[addr] == 0 { + t.Errorf("%s took none of 400 connections: %v", addr, seen) + } + } + for addr, st := range statsOf(pool) { + if st.Inflight() != 0 { + t.Errorf("%s holds %d in flight after every connection closed", addr, st.Inflight()) + } + } + for _, member := range members[:3] { + member.status.Set(-1) + } + if _, ok := u.Pick(clientFlow(l4.ProtocolTCP, "198.51.100.1:5000", "")); ok { + t.Error("dialed although no member is live") + } + }) + } +} + +// a client keeps its member across connections and across its ephemeral ports, and an IPv6 +// client across the privacy addresses of its /64 +func TestHRWKeysOnTheClientAddress(t *testing.T) { + m := origins(t, 5) + u := FromBackend(newALB(t, "alb", "hrw", up(m[0], 1), up(m[1], 1), up(m[2], 1), up(m[3], 1), up(m[4], 1))) + owner := func(f l4.Flow) string { + r, ok := u.Pick(f) + if !ok { + t.Fatal("refused") + } + r.Dialed(0, l4.ErrAbandoned) + return r.Addr() + } + reached := make(map[string]bool) + for i := range 40 { + ip := "203.0.113." + strconv.Itoa(i+1) + first := owner(clientFlow(l4.ProtocolTCP, ip+":40000", "")) + reached[first] = true + for _, port := range []string{"40001", "51234"} { + if got := owner(clientFlow(l4.ProtocolUDP, ip+":"+port, "")); got != first { + t.Fatalf("%s moved from %s to %s with its port", ip, first, got) + } + } + } + if len(reached) < 4 { + t.Errorf("40 clients reached only %d of 5 members", len(reached)) + } + a := owner(clientFlow(l4.ProtocolTCP, "[2001:db8:1:2:aaaa:bbbb:cccc:dddd]:443", "")) + if b := owner(clientFlow(l4.ProtocolTCP, "[2001:db8:1:2:1111:2222:3333:4444]:443", "")); a != b { + t.Errorf("two addresses of one /64 reached %s and %s", a, b) + } + // a flow with no usable client address has no affinity, and is still served + if _, ok := u.Pick(l4.Flow{Protocol: l4.ProtocolTCP}); !ok { + t.Error("a flow with no client address was refused") + } +} + +// keyed on the server name, every client of one name shares a member, whoever they are +func TestHRWKeysOnTheServerName(t *testing.T) { + m := origins(t, 5) + sni := func(o *ao.Options) { o.HRW.Key = "sni" } + u := FromBackend(newALBWith(t, "alb", "hrw", sni, up(m[0], 1), up(m[1], 1), up(m[2], 1), up(m[3], 1), up(m[4], 1))) + owner := func(client, name string) string { + r, ok := u.Pick(clientFlow(l4.ProtocolTLS, client, name)) + if !ok { + t.Fatal("refused") + } + r.Dialed(0, l4.ErrAbandoned) + return r.Addr() + } + reached := make(map[string]bool) + for i := range 30 { + name := "tenant" + strconv.Itoa(i) + ".example.com" + first := owner("198.51.100.1:1000", name) + reached[first] = true + if got := owner("203.0.113.77:2000", "Tenant"+strconv.Itoa(i)+".Example.COM"); got != first { + t.Fatalf("%s reached %s from one client and %s from another", name, first, got) + } + } + if len(reached) < 4 { + t.Errorf("30 server names reached only %d of 5 members", len(reached)) + } + // a client that offers no name has nothing to be kept by + spread := make(map[string]bool) + for range 60 { + spread[owner("198.51.100.1:1000", "")] = true + } + if len(spread) < 3 { + t.Errorf("60 nameless flows reached only %d members", len(spread)) + } +} + +// each load balancer a flow passes through reads the key its own way: round robin across the +// inner pools, as a weighted rule compiles to, and affinity within the one chosen +func TestNestedLoadBalancersKeyForThemselves(t *testing.T) { + m := origins(t, 6) + inner := func(name string, members ...spec) backends.Backend { return newALB(t, name, "hrw", members...) } + outer := newALB(t, "outer", "rr", + up(inner("inner-a", up(m[0], 1), up(m[1], 1), up(m[2], 1)), 1), + up(inner("inner-b", up(m[3], 1), up(m[4], 1), up(m[5], 1)), 1)) + u := FromBackend(outer) + first := make(map[bool]string) + for i := range 12 { + r, ok := u.Pick(clientFlow(l4.ProtocolTCP, "203.0.113.9:"+strconv.Itoa(30000+i), "")) + if !ok { + t.Fatal("refused") + } + r.Dialed(0, l4.ErrAbandoned) + inA := r.Addr() == "10.0.0.1:9000" || r.Addr() == "10.0.0.2:9000" || r.Addr() == "10.0.0.3:9000" + if prev, seen := first[inA]; seen && prev != r.Addr() { + t.Fatalf("one client reached %s and %s within one inner pool", prev, r.Addr()) + } + first[inA] = r.Addr() + } + if len(first) != 2 { + t.Errorf("round robin across the inner pools reached %d of them", len(first)) + } +} + +// what is timed depends on the protocol and on lt.signal: the connect, or the member's +// first byte, and never both +func TestLatencySignals(t *testing.T) { + for _, test := range []struct { + name, protocol, signal string + wantConnect bool + }{ + {"tcp default is the connect", l4.ProtocolTCP, "", true}, + {"tls default is the connect", l4.ProtocolTLS, "", true}, + {"tcp first_byte", l4.ProtocolTCP, ao.LTSignalFirstByte, false}, + {"udp default is the first reply", l4.ProtocolUDP, "", false}, + } { + t.Run(test.name, func(t *testing.T) { + // a name of its own: an ALB that keeps stats inherits them from its predecessor + pool := newALBWith(t, "signal-"+test.protocol+test.signal, "lt", func(o *ao.Options) { o.LT.Signal = test.signal }, + up(origin(t, "a", "10.0.0.1:9000"), 1)) + st := statsOf(pool)["10.0.0.1:9000"] + r, _ := FromBackend(pool).Pick(clientFlow(test.protocol, "198.51.100.1:5000", "")) + r.Dialed(40*time.Millisecond, nil) + if got := st.Latency() == 40*time.Millisecond; got != test.wantConnect { + t.Fatalf("after the connect the average is %v", st.Latency()) + } + time.Sleep(15 * time.Millisecond) + r.FirstByte() + if test.wantConnect { + if st.Latency() != 40*time.Millisecond { + t.Errorf("the first byte was sampled as well as the connect: %v", st.Latency()) + } + } else if st.Latency() < 15*time.Millisecond || st.Latency() >= 40*time.Millisecond { + t.Errorf("first-byte sample = %v, want the time since the pick", st.Latency()) + } + r.Closed(nil) + }) + } +} + +// a failed connect is retried on another member, as many times as configured, and never onto +// the member that just failed +func TestConnectRetries(t *testing.T) { + m := origins(t, 3) + retries := func(n int) func(*ao.Options) { + return func(o *ao.Options) { o.Stream = &ao.StreamOptions{ConnectRetries: n} } + } + u := FromBackend(newALBWith(t, "alb", "rr", retries(2), up(m[0], 1), up(m[1], 1), up(m[2], 1))) + retrier, ok := u.(l4.Retrier) + if !ok { + t.Fatal("a pooled upstream cannot retry") + } + flow := clientFlow(l4.ProtocolTCP, "198.51.100.1:5000", "") + first, _ := u.Pick(flow) + first.Dialed(time.Millisecond, errRefused) + second, ok := retrier.Retry(flow, first) + if !ok || second.Addr() == first.Addr() { + t.Fatalf("first retry = %v, %v", second, ok) + } + second.Dialed(time.Millisecond, errRefused) + third, ok := retrier.Retry(flow, second) + if !ok || third.Addr() == second.Addr() { + t.Fatalf("second retry = %v, %v", third, ok) + } + third.Dialed(time.Millisecond, errRefused) + if _, ok := retrier.Retry(flow, third); ok { + t.Error("a third retry was offered with connect_retries: 2") + } + // the default offers none, and a route that is not the adapter's is not retried + none := FromBackend(newALB(t, "none", "rr", up(m[0], 1), up(m[1], 1))).(l4.Retrier) + r, _ := none.(l4.Upstream).Pick(flow) + r.Dialed(time.Millisecond, errRefused) + if _, ok := none.Retry(flow, r); ok { + t.Error("a retry was offered with no connect_retries configured") + } + static, _ := l4.Static("10.9.9.9:1").Pick(flow) + if _, ok := retrier.Retry(flow, static); ok { + t.Error("retried a route the adapter never issued") + } + // a retry that lands on a member that must refuse its share ends there + gone := FromBackend(newALBWith(t, "gone", "rr", retries(3), up(m[0], 1), + up(origin(t, "gone", "unresolved.kgw.invalid:1"), 1))) + for range 4 { + r, ok := gone.Pick(flow) + if !ok { + continue + } + r.Dialed(time.Millisecond, errRefused) + if next, ok := gone.(l4.Retrier).Retry(flow, r); ok { + t.Fatalf("retried onto %s, past a member that refuses its share", next.Addr()) + } + } +} + +// connects that keep failing eject the member, and what reaches the others afterwards is +// all of the traffic +func TestPassiveEjectionThroughTheAdapter(t *testing.T) { + m := origins(t, 3) + passive := func(o *ao.Options) { + o.Stream = &ao.StreamOptions{PassiveHealth: &ao.PassiveHealthOptions{Failures: 2, Eject: timeconv.Duration(time.Hour)}} + } + pool := newALBWith(t, "ejecting-alb", "rr", passive, up(m[0], 1), up(m[1], 1), up(m[2], 1)) + u := FromBackend(pool) + before := testutil.ToFloat64(metrics.ALBMemberEjections.WithLabelValues("ejecting-alb", "m1")) + flow := clientFlow(l4.ProtocolTCP, "198.51.100.1:5000", "") + for range 12 { + r, ok := u.Pick(flow) + if !ok { + t.Fatal("refused") + } + if r.Addr() == "10.0.0.2:9000" { + r.Dialed(time.Millisecond, errRefused) + continue + } + r.Dialed(time.Millisecond, nil) + r.Closed(nil) + } + if got := testutil.ToFloat64(metrics.ALBMemberEjections.WithLabelValues("ejecting-alb", "m1")) - before; got != 1 { + t.Fatalf("ejections metered = %v", got) + } + for range 20 { + r, _ := u.Pick(flow) + if r.Addr() == "10.0.0.2:9000" { + t.Fatal("an ejected member was dialed") + } + r.Dialed(time.Millisecond, nil) + r.Closed(nil) + } + // a udp member that answers with a port-unreachable counts the same way + udp := newALBWith(t, "udp-alb", "rr", passive, up(m[0], 1), up(m[1], 1)) + uu := FromBackend(udp) + for range 8 { + r, _ := uu.Pick(clientFlow(l4.ProtocolUDP, "198.51.100.1:5000", "")) + r.Dialed(0, nil) + if r.Addr() == "10.0.0.1:9000" { + r.Closed(errRefused) + continue + } + r.Closed(nil) + } + if !statsOf(udp)["10.0.0.1:9000"].Ejected(time.Now()) { + t.Error("a udp member that kept refusing datagrams was not ejected") + } +} + +func TestMemberMetrics(t *testing.T) { + pool := newALB(t, "metered", "rr", up(origin(t, "metered-member", "10.0.0.1:9000"), 1)) + u := FromBackend(pool) + flow := l4.Flow{Listener: "metered-listener", Protocol: l4.ProtocolTCP} + labels := []string{"metered-listener", l4.ProtocolTCP, "metered-member"} + active := metrics.ProxyStreamMemberActiveConnections.WithLabelValues(labels...) + count := func(result string) float64 { + return testutil.ToFloat64(metrics.ProxyStreamMemberConnections.WithLabelValues(append(labels, result)...)) + } + r, _ := u.Pick(flow) + r.Dialed(5*time.Millisecond, nil) + if testutil.ToFloat64(active) != 1 { + t.Errorf("active = %v", testutil.ToFloat64(active)) + } + r.Closed(nil) + r, _ = u.Pick(flow) + r.Dialed(time.Millisecond, errRefused) + r, _ = u.Pick(flow) + r.Dialed(0, nil) + r.Closed(errRefused) + r, _ = u.Pick(flow) + r.Dialed(0, l4.ErrAbandoned) + if testutil.ToFloat64(active) != 0 { + t.Errorf("active after every route ended = %v", testutil.ToFloat64(active)) + } + for result, want := range map[string]float64{ResultProxied: 1, ResultDialFailed: 1, ResultUnreachable: 1} { + if got := count(result); got != want { + t.Errorf("%s = %v, want %v", result, got, want) + } + } + if n := testutil.CollectAndCount(metrics.ProxyStreamMemberConnectDuration); n == 0 { + t.Error("no connect duration was observed") + } + // a discovered member that leaves takes its series with it + metrics.DeleteBackendSeries("metered-member") + if got := count(ResultProxied); got != 0 { + t.Errorf("series survived the member: %v", got) + } +} + +// a level with no load balancer options of its own takes each protocol's defaults +func TestDefaultsWithoutOptions(t *testing.T) { + if optionsOf(lb.NewMember(lb.MemberOptions{Name: "stray", Value: "not a target"})) != nil { + t.Error("options were found for a member that is not a pool target") + } + if !timesConnect(nil, l4.ProtocolTCP) || !timesConnect(nil, l4.ProtocolTLS) || timesConnect(nil, l4.ProtocolUDP) { + t.Error("the default signal is the connect on tcp and tls, and the first reply on udp") + } + a := key(nil, clientFlow(l4.ProtocolTCP, "[2001:db8:1:2::1]:1", "")) + b := key(nil, clientFlow(l4.ProtocolTCP, "[2001:db8:1:2::2]:2", "")) + if !a.HasKey || a != b { + t.Error("without options a client is not keyed on its /64") + } +} + +// series are resolved once per member, and the cache cannot grow without bound as discovered +// members come and go +func TestMemberSeriesAreCached(t *testing.T) { + u := &upstream{} + f := l4.Flow{Listener: "cached-listener", Protocol: l4.ProtocolTCP} + m := lb.NewMember(lb.MemberOptions{Name: "cached-member"}) + first := u.series.seriesFor(m, f, "cached-member") + if u.series.seriesFor(m, f, "cached-member") != first { + t.Error("a member's series were resolved twice") + } + if allocs := testing.AllocsPerRun(100, func() { _ = u.series.seriesFor(m, f, "cached-member") }); allocs != 0 { + t.Errorf("a cached lookup allocates %v", allocs) + } + for range maxCachedSeries + 10 { + u.series.seriesFor(lb.NewMember(lb.MemberOptions{Name: "cached-member"}), f, "cached-member") + } + if got := u.series.seriesOf.Load(); got > maxCachedSeries { + t.Errorf("the cache holds %d members' series", got) + } + if u.series.seriesFor(m, f, "cached-member") == nil { + t.Error("a member dropped from the cache could not be resolved again") + } +} + +type tlvs map[byte]string + +func (h tlvs) ProxyTLV(typ byte) ([]byte, bool) { + v, ok := h[typ] + return []byte(v), ok +} + +// keyed on a PROXY protocol TLV, every connection that carries one value shares a member +func TestHRWKeysOnAProxyTLV(t *testing.T) { + m := origins(t, 5) + const endpoint = 0xEA + tlv := func(o *ao.Options) { o.HRW.Key = "proxy_tlv:0xEA" } + u := FromBackend(newALBWith(t, "alb", "hrw", tlv, up(m[0], 1), up(m[1], 1), up(m[2], 1), up(m[3], 1), up(m[4], 1))) + owner := func(client string, header l4.ProxyHeader) string { + f := clientFlow(l4.ProtocolTCP, client, "") + f.Proxy = header + r, ok := u.Pick(f) + if !ok { + t.Fatal("refused") + } + r.Dialed(0, l4.ErrAbandoned) + return r.Addr() + } + reached := make(map[string]bool) + for i := range 30 { + header := tlvs{endpoint: "vpce-" + strconv.Itoa(i), 0x02: "ignored"} + first := owner("198.51.100.1:1000", header) + reached[first] = true + if got := owner("203.0.113.77:2000", header); got != first { + t.Fatalf("endpoint %d reached %s from one client and %s from another", i, first, got) + } + } + if len(reached) < 4 { + t.Errorf("30 endpoints reached only %d of 5 members", len(reached)) + } + // a connection with no header, without the TLV, or with an empty one has nothing to be kept by + for name, header := range map[string]l4.ProxyHeader{"none": nil, "other": tlvs{0x02: "x"}, "empty": tlvs{endpoint: ""}} { + spread := make(map[string]bool) + for range 60 { + spread[owner("198.51.100.1:1000", header)] = true + } + if len(spread) < 3 { + t.Errorf("%s: 60 unkeyed flows reached only %d members", name, len(spread)) + } + } +} + +// exhaust offers a flow every route its upstream will give it, failing each, and returns the +// addresses in the order they were tried +func exhaust(t *testing.T, u l4.Upstream, flow l4.Flow) []string { + t.Helper() + r, ok := u.Pick(flow) + var tried []string + for ok { + tried = append(tried, r.Addr()) + r.Dialed(time.Millisecond, errRefused) + r, ok = u.(l4.Retrier).Retry(flow, r) + } + return tried +} + +func distinct(addrs []string) bool { + seen := make(map[string]bool, len(addrs)) + for _, a := range addrs { + if seen[a] { + return false + } + seen[a] = true + } + return true +} + +// a retry never returns to a member the flow has already failed to reach, whatever the strategy: +// with enough retries every member is tried once, and then the flow is refused +func TestConnectRetriesVisitEveryMemberOnce(t *testing.T) { + m := origins(t, 4) + retries := func(o *ao.Options) { o.Stream = &ao.StreamOptions{ConnectRetries: 10} } + for _, mech := range []string{"rr", "p2c", "lc", "lt", "hrw"} { + b := newALBWith(t, "distinct-"+mech, mech, retries, up(m[0], 1), up(m[1], 3), up(m[2], 1), up(m[3], 1)) + u := FromBackend(b) + for i := range 20 { + flow := clientFlow(l4.ProtocolTCP, "198.51.100."+strconv.Itoa(i)+":5000", "") + if tried := exhaust(t, u, flow); len(tried) != 4 || !distinct(tried) { + t.Fatalf("%s: tried %v", mech, tried) + } + } + for addr, st := range statsOf(b) { + if st.Inflight() != 0 { + t.Errorf("%s: %s left %d in flight", mech, addr, st.Inflight()) + } + } + } +} + +// across nested pools too: a pool with no member left to try is passed over for its siblings +func TestConnectRetriesCrossNestedPools(t *testing.T) { + m := origins(t, 5) + retries := func(o *ao.Options) { o.Stream = &ao.StreamOptions{ConnectRetries: 10} } + left := newALB(t, "left", "hrw", up(m[0], 1), up(m[1], 1)) + right := newALB(t, "right", "rr", up(m[2], 1), up(m[3], 1), up(m[4], 1)) + u := FromBackend(newALBWith(t, "outer", "rr", retries, up(left, 1), up(right, 1))) + for i := range 10 { + flow := clientFlow(l4.ProtocolTCP, "198.51.100."+strconv.Itoa(i)+":5000", "") + if tried := exhaust(t, u, flow); len(tried) != 5 || !distinct(tried) { + t.Fatalf("tried %v", tried) + } + } +} + +// ejection follows the order of connects, not of connections ending: a long-lived connection +// between two failed dials breaks the run when it connects, and says nothing when it closes +func TestPassiveEjectionFollowsConnectOrder(t *testing.T) { + passive := func(o *ao.Options) { + o.Stream = &ao.StreamOptions{PassiveHealth: &ao.PassiveHealthOptions{Failures: 2, Eject: timeconv.Duration(time.Hour)}} + } + // a standby keeps the member ejectable: the last live member never is + pool := newALBWith(t, "ordered-alb", "rr", passive, + up(origin(t, "ordered-a", "10.0.0.1:9000"), 1), up(origin(t, "ordered-b", "10.0.0.2:9000"), 1)) + u := FromBackend(pool) + flow := clientFlow(l4.ProtocolTCP, "198.51.100.1:5000", "") + st := statsOf(pool)["10.0.0.1:9000"] + toA := func() l4.Route { + for range 4 { + r, _ := u.Pick(flow) + if r.Addr() == "10.0.0.1:9000" { + return r + } + r.Dialed(time.Millisecond, nil) + r.Closed(nil) + } + t.Fatal("the member was never picked") + return nil + } + toA().Dialed(time.Millisecond, errRefused) + long := toA() + long.Dialed(time.Millisecond, nil) + if st.ConnectFailures() != 0 { + t.Fatal("a successful dial did not end the run of failed dials") + } + toA().Dialed(time.Millisecond, errRefused) + if st.Ejected(time.Now()) { + t.Fatal("ejected for two failed dials that a successful one came between") + } + long.Closed(nil) + if st.ConnectFailures() != 1 { + t.Errorf("connect failures after the long connection closed = %d", st.ConnectFailures()) + } + toA().Dialed(time.Millisecond, errRefused) + if !st.Ejected(time.Now()) { + t.Error("two consecutive failed dials did not eject the member") + } + + // a udp socket always opens, so only a reply, or a session that ends clean, is a success + udp := newALBWith(t, "ordered-udp", "rr", passive, + up(origin(t, "ordered-ua", "10.0.0.1:9000"), 1), up(origin(t, "ordered-ub", "10.0.0.2:9000"), 1)) + uu := FromBackend(udp) + ust := statsOf(udp)["10.0.0.1:9000"] + session := func() l4.Route { + for range 4 { + r, _ := uu.Pick(clientFlow(l4.ProtocolUDP, "198.51.100.1:5000", "")) + r.Dialed(0, nil) + if r.Addr() == "10.0.0.1:9000" { + return r + } + r.Closed(nil) + } + t.Fatal("the member was never picked") + return nil + } + session().Closed(errRefused) + if ust.ConnectFailures() != 1 { + t.Fatalf("opening a udp socket reset the count: %d", ust.ConnectFailures()) + } + answered := session() + answered.FirstByte() + if ust.ConnectFailures() != 0 { + t.Error("a reply did not end the run of unreachable sessions") + } + answered.Closed(nil) + session().Closed(errRefused) + session().Closed(nil) + if ust.ConnectFailures() != 0 || ust.Ejected(time.Now()) { + t.Error("a session that ended clean did not end the run") + } +} + +// stubborn is a picker that cannot avoid members, and always offers the same one +type stubborn struct{ b *lb.Balancer } + +func (s stubborn) Needs() lb.Needs { return s.b.Needs() } + +func (s stubborn) Pick(f lb.Flow) (lb.Pick, bool) { return s.b.Pick(f) } + +// a picker that offers a retry the member the flow already failed on is refused, not dialed again +func TestRetryRefusesAMemberAlreadyTried(t *testing.T) { + pool := newALB(t, "stubborn-alb", "rr", up(origin(t, "stubborn-only", "10.0.0.1:9000"), 1)) + u := &upstream{picker: stubborn{pool.(*alb.Client).Picker().(*lb.Balancer)}, retries: 3} + flow := clientFlow(l4.ProtocolTCP, "198.51.100.1:5000", "") + first, ok := u.Pick(flow) + if !ok { + t.Fatal("refused") + } + first.Dialed(time.Millisecond, errRefused) + if again, ok := u.Retry(flow, first); ok { + t.Errorf("retried onto %s, which had just failed", again.Addr()) + } + for addr, st := range statsOf(pool) { + if st.Inflight() != 0 { + t.Errorf("%s left %d in flight", addr, st.Inflight()) + } + } +} + +// at the most retries a flow may have, under an affinity strategy over pools of one member +// each, every retry starts at the same pool and must pass every pool already spent: the last +// attempt passes ten of them, and still reaches a member that has not been tried +func TestConnectRetriesReachPastManySpentPools(t *testing.T) { + retries := func(o *ao.Options) { o.Stream = &ao.StreamOptions{ConnectRetries: ao.MaxConnectRetries} } + var children []spec + for i := range ao.MaxConnectRetries + 2 { + name := strconv.Itoa(i) + child := newALB(t, "single-"+name, "rr", up(origin(t, "single-leaf-"+name, "10.0.1."+name+":9000"), 1)) + children = append(children, up(child, 1)) + } + u := FromBackend(newALBWith(t, "affinity-outer", "hrw", retries, children...)) + for i := range 10 { + flow := clientFlow(l4.ProtocolTCP, "198.51.100."+strconv.Itoa(i)+":5000", "") + tried := exhaust(t, u, flow) + if len(tried) != ao.MaxConnectRetries+1 || !distinct(tried) { + t.Fatalf("client %d was offered %d members: %v", i, len(tried), tried) + } + } +} diff --git a/pkg/backends/alb/stream/stream.go b/pkg/backends/alb/stream/stream.go new file mode 100644 index 000000000..e54c272ae --- /dev/null +++ b/pkg/backends/alb/stream/stream.go @@ -0,0 +1,313 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package stream is the adapter between load-balanced pools and the tcp, tls and udp relay: +// it presents a backend to the relay as an upstream, committing each connection or session to +// the pool member that the backend's selection strategy picks. +package stream + +import ( + "errors" + "slices" + "sync" + "sync/atomic" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" + + "github.com/prometheus/client_golang/prometheus" +) + +// Results a pool member's connections and sessions are counted under +const ( + ResultProxied = "proxied" + ResultDialFailed = "dial_failed" + ResultUnreachable = "unreachable" +) + +// FromBackend returns an upstream over a backend. A load balancer that selects one member per +// flow commits each flow to a healthy member, following a member that is itself such a load +// balancer; any other backend is dialed at its origin host. It returns nil for a backend with +// neither. A member that cannot be dialed refuses its share rather than passing it to a sibling. +func FromBackend(b backends.Backend) l4.Upstream { + if b == nil { + return nil + } + cfg := b.Configuration() + if sp, ok := b.(spreader); ok && sp.Spread() != 0 { + return newSpread(sp, cfg) + } + if pp, ok := b.(lb.PickerProvider); ok { + if p := pp.Picker(); p != nil { + u := &upstream{picker: p} + if cfg != nil && cfg.ALBOptions != nil { + u.options = cfg.ALBOptions + if s := cfg.ALBOptions.Stream; s != nil { + u.retries = s.ConnectRetries + } + } + return u + } + } + if cfg == nil || cfg.Host == "" { + return nil + } + return l4.Static(cfg.Host) +} + +type upstream struct { + picker lb.Picker + // the options of the load balancer the listener is bound to; its members that are load + // balancers themselves are keyed and timed by their own + options *ao.Options + retries int + series seriesCache +} + +// seriesCache holds the series of each member a listener has dialed, resolved once per member +// rather than looked up by label on every connection +type seriesCache struct { + series sync.Map + seriesOf atomic.Int64 +} + +// memberSeries are one member's series for one listener +type memberSeries struct { + active prometheus.Gauge + connect prometheus.Observer + proxied, dialFailed, unreachable prometheus.Counter +} + +// maxCachedSeries bounds the cache: members come and go under discovery, and a member that +// has left is never looked up again, so the cache is simply dropped when it has grown large +const maxCachedSeries = 1024 + +// seriesFor returns the member's series. The cache is keyed by the member itself, not its +// name: a member that leaves has its series deleted, and one that returns under the same +// name is a new member whose series must be resolved afresh. +func (u *seriesCache) seriesFor(m *lb.Member, f l4.Flow, name string) *memberSeries { + if v, ok := u.series.Load(m); ok { + return v.(*memberSeries) + } + if u.seriesOf.Add(1) > maxCachedSeries { + u.series.Clear() + u.seriesOf.Store(1) + } + count := func(result string) prometheus.Counter { + return metrics.ProxyStreamMemberConnections.WithLabelValues(f.Listener, f.Protocol, name, result) + } + s := &memberSeries{ + active: metrics.ProxyStreamMemberActiveConnections.WithLabelValues(f.Listener, f.Protocol, name), + connect: metrics.ProxyStreamMemberConnectDuration.WithLabelValues(f.Listener, f.Protocol, name), + proxied: count(ResultProxied), + dialFailed: count(ResultDialFailed), + unreachable: count(ResultUnreachable), + } + v, _ := u.series.LoadOrStore(m, s) + return v.(*memberSeries) +} + +func (u *upstream) Pick(f l4.Flow) (l4.Route, bool) { + return u.pick(f, nil) +} + +// retryState is what a route that is a retry carries; a flow's first route has none +type retryState struct { + attempt int + // every member the flow failed to reach before this route + tried []*lb.Member +} + +// Retry offers another member when a connection could not reach the one it was given, as +// far as stream.connect_retries allows, and never one the flow has already been offered +func (u *upstream) Retry(f l4.Flow, failed l4.Route) (l4.Route, bool) { + prev, ok := failed.(*route) + if !ok { + return nil, false + } + next := &retryState{attempt: 1} + if prev.retry != nil { + next.attempt = prev.retry.attempt + 1 + next.tried = slices.Clip(prev.retry.tried) + } + if next.attempt > u.retries { + return nil, false + } + next.tried = append(next.tried, prev.pick.Member()) + return u.pick(f, next) +} + +func (u *upstream) pick(f l4.Flow, retry *retryState) (l4.Route, bool) { + r := &route{retry: retry, udp: f.Protocol == l4.ProtocolUDP} + var tried []*lb.Member + if retry != nil { + tried = retry.tried + } + flowOf := func(depth int, p lb.Picker, via *lb.Member) lb.Flow { + o := u.options + if via != nil { + o = optionsOf(via) + } + // a retry may try several members at one depth; the last one asked is the one kept + r.onConnect[depth] = timesConnect(o, f.Protocol) + if !p.Needs().Has(lb.NeedKey) { + return lb.Flow{} + } + return key(o, f) + } + var pk lb.LeafPick + var ok bool + if len(tried) == 0 { + pk, ok = lb.PickLeafFunc(u.picker, flowOf) + } else { + pk, ok = lb.RepickLeafFunc(u.picker, flowOf, tried...) + } + if !ok { + return nil, false + } + if retry != nil && slices.Contains(tried, pk.Member()) { + // a picker that cannot avoid members offered one the flow has already failed to reach + pk.Done(lb.OutcomeCanceled) + return nil, false + } + t, ok := pk.Member().Value.(*pool.Target) + if !ok || !t.Dialable() { + // the member holds its share and refuses it; that is no fault of the member's + pk.Done(lb.OutcomeCanceled) + return nil, false + } + r.pick, r.addr = pk, t.Addr() + r.series = u.series.seriesFor(pk.Member(), f, t.Name()) + return r, true +} + +// optionsOf returns the options of the load balancer that a pool member is, or nil +func optionsOf(m *lb.Member) *ao.Options { + if t, ok := m.Value.(*pool.Target); ok && t.Backend() != nil { + if cfg := t.Backend().Configuration(); cfg != nil { + return cfg.ALBOptions + } + } + return nil +} + +// key is the flow's affinity key as one load balancer is configured to read it: the server +// name a tls client offered, a PROXY protocol TLV, or else the client's address, never its port +func key(o *ao.Options, f l4.Flow) lb.Flow { + prefix := ao.DefaultIPv6Prefix + if o != nil { + switch o.HRW.KeySource.Kind { + case ao.KeySNI: + if f.ServerName == "" { + return lb.Flow{} + } + return lb.Flow{Key: lb.HashFold(f.ServerName), HasKey: true} + case ao.KeyProxyTLV: + if f.Proxy == nil { + return lb.Flow{} + } + v, ok := f.Proxy.ProxyTLV(o.HRW.KeySource.TLV) + if !ok || len(v) == 0 { + return lb.Flow{} + } + return lb.Flow{Key: lb.HashBytes(v), HasKey: true} + } + if o.HRW.IPv6Prefix > 0 { + prefix = o.HRW.IPv6Prefix + } + } + if !f.Client.IsValid() { + return lb.Flow{} + } + return lb.Flow{Key: lb.HashAddr(f.Client.Addr(), prefix), HasKey: true} +} + +// timesConnect reports whether one load balancer samples latency at the connect, which is a +// tcp or tls listener's default, rather than at the member's first byte or datagram +func timesConnect(o *ao.Options, protocol string) bool { + if o == nil { + return protocol != l4.ProtocolUDP + } + signal, err := o.LTSignalFor(protocol) + return err == nil && signal == ao.LTSignalConnect +} + +// route is one flow's commitment to a member, allocated once per connection or session +type route struct { + pick lb.LeafPick + addr string + series *memberSeries + // for each load balancer passed through, whether it times the connect or the first byte + onConnect [lb.MaxPickDepth]bool + udp bool + // nil unless the route is a retry + retry *retryState +} + +func (r *route) Addr() string { return r.addr } + +// Final is false: a pool may have another member to try +func (r *route) Final() bool { return false } + +func (r *route) Dialed(d time.Duration, err error) { + switch { + case err == nil: + if !r.udp { + // a udp socket opens whether or not anything listens; a reply is what reaches + r.pick.Reached() + } + for i := range r.pick.Depth() { + if r.onConnect[i] { + r.pick.Level(i).Established(d) + } + } + r.series.connect.Observe(d.Seconds()) + r.series.active.Inc() + case errors.Is(err, l4.ErrAbandoned): + r.pick.Done(lb.OutcomeCanceled) + default: + r.pick.Done(lb.OutcomeConnectFailed) + r.series.dialFailed.Inc() + } +} + +func (r *route) FirstByte() { + r.pick.Reached() + for i := range r.pick.Depth() { + if !r.onConnect[i] { + r.pick.Level(i).FirstByte() + } + } +} + +func (r *route) Closed(err error) { + r.series.active.Dec() + if err != nil { + r.pick.Done(lb.OutcomeConnectFailed) + r.series.unreachable.Inc() + return + } + if r.udp { + // a session that ended without a port-unreachable is all a one-way member ever shows + r.pick.Reached() + } + r.pick.Done(lb.OutcomeOK) + r.series.proxied.Inc() +} diff --git a/pkg/backends/alb/stream/stream_test.go b/pkg/backends/alb/stream/stream_test.go new file mode 100644 index 000000000..ac0a28e65 --- /dev/null +++ b/pkg/backends/alb/stream/stream_test.go @@ -0,0 +1,181 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package stream + +import ( + "errors" + "maps" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/alb" + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" +) + +func TestFromBackend(t *testing.T) { + if FromBackend(nil) != nil { + t.Error("nil backend yields an upstream") + } + if FromBackend(origin(t, "hostless", "")) != nil { + t.Error("a backend without an origin host yields an upstream") + } + r, ok := FromBackend(origin(t, "o", "10.0.0.1:9000")).Pick(l4.Flow{}) + if !ok || r.Addr() != "10.0.0.1:9000" || !r.Final() { + t.Errorf("origin backend = %v, %v", r, ok) + } + if _, ok := FromBackend(origin(t, "gone", "unresolved.kgw.invalid:1")).Pick(l4.Flow{}); ok { + t.Error("a backend under .invalid dials") + } + // a load balancer that does not commit a flow to one member has nothing to relay to + if FromBackend(newALB(t, "fanout", "fr", up(origin(t, "a", "10.0.0.1:1"), 1))) != nil { + t.Error("a fanout load balancer yields an upstream") + } + pooled := FromBackend(newALB(t, "alb", "rr", up(origin(t, "a", "10.0.0.1:1"), 1))) + r, ok = pooled.Pick(l4.Flow{}) + if !ok || r.Addr() != "10.0.0.1:1" || r.Final() { + t.Errorf("pooled route = %v, %v; a pool's route is not final", r, ok) + } +} + +func TestEveryWindowIsExact(t *testing.T) { + a := origin(t, "a", "10.0.0.1:1") + b := origin(t, "b", "10.0.0.2:1") + c := origin(t, "c", "10.0.0.3:1") + gone := origin(t, "gone", "unresolved.kgw.invalid:1") + hostless := origin(t, "hostless", "") + failing := origin(t, "failing", "10.0.0.9:1") + + t.Run("uniform", func(t *testing.T) { + u := FromBackend(newALB(t, "alb", "rr", up(a, 1), up(b, 1), up(c, 1))) + assertEveryWindowExact(t, pickN(u, 20), map[string]int{"10.0.0.1:1": 1, "10.0.0.2:1": 1, "10.0.0.3:1": 1}) + }) + t.Run("weighted", func(t *testing.T) { + u := FromBackend(newALB(t, "alb", "rr", up(a, 1), up(b, 3), up(c, 2))) + assertEveryWindowExact(t, pickN(u, 40), map[string]int{"10.0.0.1:1": 1, "10.0.0.2:1": 3, "10.0.0.3:1": 2}) + }) + t.Run("an unhealthy member holds no share", func(t *testing.T) { + u := FromBackend(newALB(t, "alb", "rr", up(a, 2), down(failing, 5), up(b, 1))) + assertEveryWindowExact(t, pickN(u, 20), map[string]int{"10.0.0.1:1": 2, "10.0.0.2:1": 1}) + }) + t.Run("a refusing member consumes and refuses exactly its weight", func(t *testing.T) { + u := FromBackend(newALB(t, "alb", "rr", up(a, 1), up(gone, 3))) + assertEveryWindowExact(t, pickN(u, 24), map[string]int{"10.0.0.1:1": 1, refusedKey: 3}) + }) + t.Run("a member with nothing to dial refuses its own share", func(t *testing.T) { + u := FromBackend(newALB(t, "alb", "rr", up(a, 1), up(hostless, 1))) + assertEveryWindowExact(t, pickN(u, 12), map[string]int{"10.0.0.1:1": 1, refusedKey: 1}) + }) +} + +func TestRefusesWithoutALiveMember(t *testing.T) { + a := origin(t, "a", "10.0.0.1:1") + b := origin(t, "b", "10.0.0.2:1") + for name, u := range map[string]l4.Upstream{ + "all unhealthy": FromBackend(newALB(t, "down", "rr", down(a, 1), down(b, 2))), + "empty": FromBackend(newALB(t, "empty", "rr")), + } { + for range 6 { + if r, ok := u.Pick(l4.Flow{}); ok { + t.Fatalf("%s: dialed %s", name, r.Addr()) + } + } + } + // a member that recovers is dialed by the very next flow + member := down(a, 1) + u := FromBackend(newALB(t, "recovering", "rr", member)) + if _, ok := u.Pick(l4.Flow{}); ok { + t.Fatal("dialed a failing member") + } + member.status.Set(healthcheck.StatusPassing) + if r, ok := u.Pick(l4.Flow{}); !ok || r.Addr() != "10.0.0.1:1" { + t.Error("a recovered member was not dialed") + } +} + +func TestNestedApportionment(t *testing.T) { + // outer 3:1 over two inner pools, the first of which is itself weighted 2:1 + newOuter := func(t *testing.T, inner func(backends.Backend, int) spec) l4.Upstream { + inner1 := newALB(t, "inner1", "rr", inner(origin(t, "a", "10.1.0.1:1"), 2), inner(origin(t, "b", "10.1.0.2:1"), 1)) + inner2 := newALB(t, "inner2", "rr", up(origin(t, "c", "10.2.0.1:1"), 1)) + return FromBackend(newALB(t, "outer", "rr", up(inner1, 3), up(inner2, 1))) + } + t.Run("each level apportions by its own weights", func(t *testing.T) { + // 12 selections hold 9 for inner1, a whole number of its rotations of 3 + got := tally(pickN(newOuter(t, up), 36)) + want := map[string]int{"10.1.0.1:1": 18, "10.1.0.2:1": 9, "10.2.0.1:1": 9} + if !maps.Equal(got, want) { + t.Errorf("nested apportionment = %v, want %v", got, want) + } + }) + t.Run("an inner pool with no live member refuses its share", func(t *testing.T) { + assertEveryWindowExact(t, pickN(newOuter(t, down), 16), map[string]int{refusedKey: 3, "10.2.0.1:1": 1}) + }) + t.Run("a pool nested beyond the depth bound is not followed", func(t *testing.T) { + l2 := newALB(t, "l2", "rr", up(origin(t, "z", "10.9.0.1:1"), 1)) + l1 := newALB(t, "l1", "rr", up(l2, 1)) + if r, ok := FromBackend(newALB(t, "l0", "rr", up(l1, 1))).Pick(l4.Flow{}); ok { + t.Errorf("a pool three deep was followed: %s", r.Addr()) + } + }) +} + +// what the relay reports of a route reaches the member's stats, at every level it passed +func TestRouteFeedbackReachesTheMember(t *testing.T) { + inner := newALB(t, "inner", "lc", up(origin(t, "a", "10.0.0.1:1"), 1)) + outer := newALB(t, "outer", "lc", up(inner, 1)) + u := FromBackend(outer) + leaf := inner.(*alb.Client).Pool().ConfiguredTargets()[0].Member().Stats() + top := outer.(*alb.Client).Pool().ConfiguredTargets()[0].Member().Stats() + inflight := func() [2]int64 { return [2]int64{top.Inflight(), leaf.Inflight()} } + + r, ok := u.Pick(l4.Flow{}) + if !ok || inflight() != [2]int64{1, 1} { + t.Fatalf("after a pick: %v, in flight %v", ok, inflight()) + } + r.Dialed(3*time.Millisecond, nil) + r.FirstByte() + if inflight() != [2]int64{1, 1} { + t.Errorf("an open connection is no longer in flight: %v", inflight()) + } + r.Closed(nil) + if inflight() != [2]int64{0, 0} || leaf.Failures() != 0 { + t.Errorf("after close: in flight %v, %d failures", inflight(), leaf.Failures()) + } + + // a failed dial ends the route and counts against the member, not against the pool above it + r, _ = u.Pick(l4.Flow{}) + r.Dialed(time.Millisecond, errors.New("connection refused")) + if inflight() != [2]int64{0, 0} || leaf.Failures() != 1 || top.Failures() != 0 { + t.Errorf("after a failed dial: in flight %v, failures %d and %d", inflight(), top.Failures(), leaf.Failures()) + } + // a route the relay gave up before dialing says nothing about the member + r, _ = u.Pick(l4.Flow{}) + r.Dialed(0, l4.ErrAbandoned) + if inflight() != [2]int64{0, 0} || leaf.Failures() != 1 { + t.Errorf("after an abandoned route: in flight %v, %d failures", inflight(), leaf.Failures()) + } + // nor does a member that must refuse its share + gone := newALB(t, "refusing", "lc", up(origin(t, "gone", "unresolved.kgw.invalid:1"), 1)) + if _, ok := FromBackend(gone).Pick(l4.Flow{}); ok { + t.Fatal("a refusing member was dialed") + } + st := gone.(*alb.Client).Pool().ConfiguredTargets()[0].Member().Stats() + if st.Inflight() != 0 || st.Failures() != 0 { + t.Errorf("a refused share left %d in flight and %d failures", st.Inflight(), st.Failures()) + } +} diff --git a/pkg/backends/alb/weights_characterization_test.go b/pkg/backends/alb/weights_characterization_test.go new file mode 100644 index 000000000..29686d14c --- /dev/null +++ b/pkg/backends/alb/weights_characterization_test.go @@ -0,0 +1,125 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package alb + +import ( + "net/http" + "net/http/httptest" + "runtime" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/backends" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + + "github.com/stretchr/testify/require" + "go.yaml.in/yaml/v3" +) + +// a weight written in a static pool reaches the pool's targets unchanged, in pool order, and +// the round robin mechanism apportions requests by it +func TestStaticYAMLWeightsReachThePool(t *testing.T) { + a := &ao.Options{} + require.NoError(t, yaml.Unmarshal([]byte(` +mechanism: rr +pool: + - a + - name: b + weight: 3 + - name: c + weight: 0 + - name: d + weight: 2 +`), a)) + o := bo.New() + o.Provider = providers.ALB + o.ALBOptions = a + cl, err := NewClient("weighted", o, nil, nil, nil, nil) + require.NoError(t, err) + c := cl.(*Client) + t.Cleanup(c.StopPool) + + clients := backends.Backends{"weighted": cl} + hits := make(map[string]*countingHandler) + for _, name := range a.Pool.Names() { + hits[name] = &countingHandler{} + mo := bo.New() + mo.Name = name + b, err := backends.New(name, mo, nil, hits[name], nil) + require.NoError(t, err) + clients[name] = b + } + require.NoError(t, c.ValidateAndStartPool(clients, nil)) + + want := map[string]int{"a": 1, "b": 3, "c": 1, "d": 2} + targets := c.Pool().ConfiguredTargets() + require.Len(t, targets, len(want)) + var total int + for i, tgt := range targets { + require.Equal(t, a.Pool[i].Name, tgt.Name(), "pool order is preserved") + require.Equal(t, want[tgt.Name()], tgt.Weight(), tgt.Name()) + total += tgt.Weight() + } + + const cycles = 4 + h := c.Handlers()[providers.ALB] + w := httptest.NewRecorder() + r := httptest.NewRequest(http.MethodGet, "http://example.com/", nil) + for range cycles * total { + h.ServeHTTP(w, r) + } + for name, weight := range want { + require.Equal(t, cycles*weight, hits[name].hits, name) + } +} + +// maxGoroutinesPerPool is what one started ALB may hold while idle: a pool runs no workers +const maxGoroutinesPerPool = 0 + +func TestIdleGoroutinesPerALB(t *testing.T) { + const albs = 100 + before := runtime.NumGoroutine() + clients := make([]*Client, albs) + for i := range clients { + clients[i] = newStaticPoolALB(t, benchTargets(8, 1)) + } + var got int + deadline := time.Now().Add(2 * time.Second) + for { + got = runtime.NumGoroutine() - before + if got <= albs*maxGoroutinesPerPool || time.Now().After(deadline) { + break + } + time.Sleep(5 * time.Millisecond) + } + t.Logf("%d idle ALBs hold %d goroutines (%.2f per ALB)", albs, got, float64(got)/albs) + if got > albs*maxGoroutinesPerPool { + t.Errorf("%d idle ALBs hold %d goroutines, want at most %d", albs, got, albs*maxGoroutinesPerPool) + } + for _, c := range clients { + c.StopPool() + } + deadline = time.Now().Add(2 * time.Second) + for runtime.NumGoroutine() > before && time.Now().Before(deadline) { + time.Sleep(5 * time.Millisecond) + } + if left := runtime.NumGoroutine() - before; left > 0 { + t.Errorf("%d goroutines outlived their pools", left) + } +} diff --git a/pkg/backends/backends.go b/pkg/backends/backends.go index 11acd5abe..f8f6f983a 100644 --- a/pkg/backends/backends.go +++ b/pkg/backends/backends.go @@ -25,18 +25,44 @@ import ( ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" bo "github.com/trickstercache/trickster/v2/pkg/backends/options" "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/observability/logging" + "github.com/trickstercache/trickster/v2/pkg/observability/logging/logger" ) type protocolHealthProber interface { HealthCheckProbe() healthcheck.Probe } +// healthCheckRefuser is implemented by a backend that cannot actively probe some origins, and +// says why; such a backend is left unprobed rather than probed in a way that can only fail +type healthCheckRefuser interface { + HealthCheckUnsupported() string +} + // healthCheckFinalizer returns the effective options a probe registers with, // possibly a clone that leaves the input's declarative state untouched. type healthCheckFinalizer interface { FinalizeHealthCheckOptions(*ho.Options) *ho.Options } +// healthStatusOwner is a virtual backend that may keep a health status of its own, as an ALB +// whose availability follows its pool does; nil when it keeps none. +type healthStatusOwner interface { + HealthStatus() *healthcheck.Status +} + +// externalRegistrar is a health checker that reports a status it does not drive +type externalRegistrar interface { + RegisterExternal(name, description string, s *healthcheck.Status) +} + +func ownStatus(c Backend) *healthcheck.Status { + if o, ok := c.(healthStatusOwner); ok { + return o.HealthStatus() + } + return nil +} + // Backends represents a map of Backends keyed by Name type Backends map[string]Backend @@ -45,47 +71,30 @@ type Backends map[string]Backend // sets the initial status of the provided targets (e.g., after a config reload) func (b Backends) StartHealthChecks(knownStatuses healthcheck.StatusLookup) (healthcheck.HealthChecker, error) { hc := healthcheck.New() - registrar, ok := hc.(healthcheck.Registrar) - if !ok { - hc.Shutdown() - return nil, errors.New("health checker does not support protocol probe registration") - } for k, c := range b { bo := c.Configuration() if k == "frontend" { continue } if !HasOrigin(bo.Provider) { - // Backends with no upstream to probe get a synthetic passing - // status so they surface in the health page and in outer ALB - // pool reporting. + // No upstream to probe: a backend that keeps its own status reports it, so the + // health page agrees with routing; any other gets a synthetic passing status. + if st := ownStatus(c); st != nil { + if er, ok := hc.(externalRegistrar); ok { + er.RegisterExternal(k, bo.Provider, st) + continue + } + } hc.RegisterVirtual(k, bo.Provider) continue } - hco := bo.HealthCheck - if hco == nil { - continue - } - bo.HealthCheck = c.DefaultHealthCheckConfig() - if bo.HealthCheck == nil { - bo.HealthCheck = hco - } else { - bo.HealthCheck.Overlay(hco) - } - probeOpts := bo.HealthCheck - if f, ok := c.(healthCheckFinalizer); ok && probeOpts != nil { - probeOpts = f.FinalizeHealthCheckOptions(probeOpts) - } - var st *healthcheck.Status - var err error - if prober, ok := c.(protocolHealthProber); ok { - st, err = registrar.RegisterProbe(k, bo.Provider, probeOpts, prober.HealthCheckProbe()) - } else { - st, err = hc.Register(k, bo.Provider, probeOpts, c.HealthCheckHTTPClient()) - } + st, err := RegisterHealthCheck(hc, k, bo.Provider, c) if err != nil { return nil, err } + if st == nil { + continue + } if oldSt, ok := knownStatuses[k]; ok { if v := oldSt.Get(); v != healthcheck.StatusInitializing { st.Set(v) @@ -96,6 +105,53 @@ func (b Backends) StartHealthChecks(knownStatuses healthcheck.StatusLookup) (hea return hc, nil } +// ErrNoProbeRegistrar is returned when a backend probes by protocol and the health checker +// cannot register such a probe. +var ErrNoProbeRegistrar = errors.New("health checker does not support protocol probe registration") + +// RegisterHealthCheck registers the active health check of a configured or discovered backend +// and returns its status. The status is nil, with no error, for a backend that configures no +// health check or whose origin cannot be probed; the latter is logged when a check was asked for. +func RegisterHealthCheck(hc healthcheck.HealthChecker, name, description string, c Backend, +) (*healthcheck.Status, error) { + bo := c.Configuration() + hco := bo.HealthCheck + if hco == nil { + return nil, nil + } + bo.HealthCheck = c.DefaultHealthCheckConfig() + if bo.HealthCheck == nil { + bo.HealthCheck = hco + } else { + bo.HealthCheck.Overlay(hco) + } + probeOpts := bo.HealthCheck + if f, ok := c.(healthCheckFinalizer); ok && probeOpts != nil { + probeOpts = f.FinalizeHealthCheckOptions(probeOpts) + } + if u, ok := c.(healthCheckRefuser); ok { + if why := u.HealthCheckUnsupported(); why != "" { + if probeOpts != nil && probeOpts.Interval > 0 { + logger.Warn("backend health check is not run", logging.Pairs{"backendName": name, "detail": why}) + } + return nil, nil + } + } + var probe healthcheck.Probe + if prober, ok := c.(protocolHealthProber); ok { + // a backend may probe by protocol for some origins and by request for the rest + probe = prober.HealthCheckProbe() + } + if probe == nil { + return hc.Register(name, description, probeOpts, c.HealthCheckHTTPClient()) + } + registrar, ok := hc.(healthcheck.Registrar) + if !ok { + return nil, ErrNoProbeRegistrar + } + return registrar.RegisterProbe(name, description, probeOpts, probe) +} + // Get returns the named origin func (b Backends) Get(backendName string) Backend { if c, ok := b[backendName]; ok { diff --git a/pkg/backends/backends_test.go b/pkg/backends/backends_test.go index 067c201cb..f77ee2c80 100644 --- a/pkg/backends/backends_test.go +++ b/pkg/backends/backends_test.go @@ -18,8 +18,10 @@ package backends import ( "context" + "errors" "net/http/httptest" "testing" + "time" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" @@ -178,3 +180,123 @@ func TestHasOrigin(t *testing.T) { t.Error("expected static not to be virtual") } } + +// choosyBackend probes by protocol only when it has a probe to offer, and may refuse to be +// probed at all +type choosyBackend struct { + testBackend + probe healthcheck.Probe + refusal string +} + +func (tb *choosyBackend) HealthCheckProbe() healthcheck.Probe { return tb.probe } + +func (tb *choosyBackend) HealthCheckUnsupported() string { return tb.refusal } + +func TestStartHealthChecksByWhatABackendOffers(t *testing.T) { + newBackend := func(name string) Backend { + o := bo.New() + o.HealthCheck = ho.New() + o.HealthCheck.Interval = 0 + c, err := New(name, o, nil, lm.NewRouter(), nil) + if err != nil { + t.Fatal(err) + } + return c + } + var probed int + b := Backends{ + // a backend with a protocol probe is probed with it + "protocol": &choosyBackend{Backend: newBackend("protocol"), + probe: func(context.Context) error { probed++; return nil }}, + // one with none to offer for its origin falls back to the request probe + "request": &choosyBackend{Backend: newBackend("request")}, + // one that cannot be probed is left out rather than probed in a way that must fail + "refuses": &choosyBackend{Backend: newBackend("refuses"), refusal: "no probe for this origin"}, + } + b["refuses"].Configuration().HealthCheck.Interval = 1 + hc, err := b.StartHealthChecks(nil) + if err != nil { + t.Fatal(err) + } + defer hc.Shutdown() + statuses := hc.Statuses() + if statuses["protocol"] == nil || statuses["request"] == nil { + t.Fatalf("registered = %v", statuses) + } + if statuses["refuses"] != nil { + t.Error("a backend that cannot be probed was registered for a probe") + } + w := httptest.NewRecorder() + statuses["protocol"].Prober()(w) + if probed != 1 { + t.Errorf("the protocol probe ran %d times", probed) + } +} + +type statusOwner struct { + Backend + status *healthcheck.Status +} + +func (o *statusOwner) HealthStatus() *healthcheck.Status { return o.status } + +// requestOnlyChecker is a health checker that cannot register a protocol probe +type requestOnlyChecker struct{ healthcheck.HealthChecker } + +func TestVirtualBackendsReportTheirOwnStatus(t *testing.T) { + virtual := func(name string) Backend { + o := bo.New() + o.Provider = providers.ALB + c, err := New(name, o, nil, lm.NewRouter(), nil) + if err != nil { + t.Fatal(err) + } + return c + } + own := healthcheck.NewStatus("follows", providers.ALB, "", healthcheck.StatusFailing, time.Time{}, nil) + hc, err := Backends{ + "follows": &statusOwner{Backend: virtual("follows"), status: own}, + "keeps": &statusOwner{Backend: virtual("keeps")}, + "synthetic": virtual("synthetic"), + }.StartHealthChecks(nil) + if err != nil { + t.Fatal(err) + } + defer hc.Shutdown() + statuses := hc.Statuses() + if statuses["follows"] != own { + t.Error("a virtual backend's own status is not the one reported") + } + for _, name := range []string{"keeps", "synthetic"} { + if st := statuses[name]; st == nil || st.Get() != healthcheck.StatusPassing { + t.Errorf("%s: status = %v", name, st) + } + } +} + +func TestRegisterHealthCheckNeedsAProbeRegistrar(t *testing.T) { + o := bo.New() + o.HealthCheck = ho.New() + c, err := New("protocol", o, nil, lm.NewRouter(), nil) + if err != nil { + t.Fatal(err) + } + hc := healthcheck.New() + defer hc.Shutdown() + probed := &choosyBackend{Backend: c, probe: func(context.Context) error { return nil }} + if _, err := RegisterHealthCheck(requestOnlyChecker{hc}, "protocol", "test", probed); !errors.Is(err, ErrNoProbeRegistrar) { + t.Errorf("error = %v", err) + } + if _, err := (Backends{"protocol": probed}).StartHealthChecks(nil); err != nil { + t.Errorf("a full health checker refused a protocol probe: %v", err) + } + bare, err := New("bare", bo.New(), nil, lm.NewRouter(), nil) + if err != nil { + t.Fatal(err) + } + bare.Configuration().HealthCheck = nil + if st, err := RegisterHealthCheck(hc, "bare", "test", bare); st != nil || err != nil { + t.Errorf("a backend with no health check: %v, %v", st, err) + } +} diff --git a/pkg/backends/clickhouse/clickhouse_test.go b/pkg/backends/clickhouse/clickhouse_test.go index 24ca3c025..2a91b18eb 100644 --- a/pkg/backends/clickhouse/clickhouse_test.go +++ b/pkg/backends/clickhouse/clickhouse_test.go @@ -228,6 +228,9 @@ func TestNativeListenerAdapterValidation(t *testing.T) { if err := a.ValidateUserRouter(nil, "", nil); err == nil { t.Fatal("accepted native user routing") } + if err := a.ValidateBalancer(nil, "", nil); err == nil { + t.Fatal("accepted native session balancing") + } if resolver := a.RouteResolver(native.BuildRequest{}); resolver != nil { t.Fatal("unexpected native route resolver") } diff --git a/pkg/backends/clickhouse/native_listener.go b/pkg/backends/clickhouse/native_listener.go index 2bfe8c3ee..47f94548b 100644 --- a/pkg/backends/clickhouse/native_listener.go +++ b/pkg/backends/clickhouse/native_listener.go @@ -66,6 +66,10 @@ func (nativeListenerAdapter) ValidateUserRouter(*config.Config, string, *bo.Opti return errors.New("ClickHouse native user routing is not supported") } +func (nativeListenerAdapter) ValidateBalancer(*config.Config, string, *bo.Options) error { + return errors.New("ClickHouse native session balancing is not supported") +} + func (nativeListenerAdapter) RouteResolver(native.BuildRequest) backends.RouteResolver { return nil } func nativeBackend(c *config.Config, name string) (string, *bo.Options, error) { diff --git a/pkg/backends/healthcheck/external_test.go b/pkg/backends/healthcheck/external_test.go index d4a4e4163..545456a5f 100644 --- a/pkg/backends/healthcheck/external_test.go +++ b/pkg/backends/healthcheck/external_test.go @@ -17,8 +17,13 @@ package healthcheck import ( + "context" + "sync/atomic" "testing" "time" + + ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" ) func TestRegisterExternal(t *testing.T) { @@ -63,3 +68,49 @@ func TestUnregisterVirtual(t *testing.T) { t.Fatal("expected virtual status to be removed") } } + +// a status that takes over a name retires the active probe registered under it: the probe loop +// has stopped by the time the registration returns, and never runs again +func TestStatusRegistrationRetiresAnActiveProbe(t *testing.T) { + for _, how := range []string{"external", "virtual"} { + hc := New().(*healthChecker) + var probes atomic.Int64 + o := ho.New() + o.Interval = timeconv.Duration(5 * time.Millisecond) + o.Timeout = timeconv.Duration(time.Second) + if _, err := hc.RegisterProbe("member", "test", o, func(context.Context) error { + probes.Add(1) + return nil + }); err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(3 * time.Second) + for probes.Load() == 0 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if probes.Load() == 0 { + t.Fatal("the probe never ran") + } + st := NewStatus("member", "test", "", StatusPassing, time.Time{}, nil) + if how == "external" { + hc.RegisterExternal("member", "test", st) + } else { + st = hc.RegisterVirtual("member", "test") + } + after := probes.Load() + hc.mtx.RLock() + _, stillActive := hc.targets["member"] + hc.mtx.RUnlock() + if stillActive { + t.Errorf("%s: the retired probe is still registered", how) + } + time.Sleep(50 * time.Millisecond) + if got := probes.Load(); got != after { + t.Errorf("%s: the retired probe ran %d more times", how, got-after) + } + if hc.Statuses()["member"] != st { + t.Errorf("%s: the status that took the name over is not the one reported", how) + } + hc.Shutdown() + } +} diff --git a/pkg/backends/healthcheck/healthcheck.go b/pkg/backends/healthcheck/healthcheck.go index c8d96d082..7568cdc45 100644 --- a/pkg/backends/healthcheck/healthcheck.go +++ b/pkg/backends/healthcheck/healthcheck.go @@ -158,10 +158,7 @@ func (hc *healthChecker) RegisterProbe(name, description string, o *ho.Options, func (hc *healthChecker) registerTarget(t *target) *Status { hc.mtx.Lock() - if t2, ok := hc.targets[t.name]; ok && t2 != nil { - // synchronous stop so the old probe loop exits before the new one starts - t2.Stop() - } + hc.retireTargetLocked(t.name) hc.targets[t.name] = t hc.statuses[t.name] = t.status hc.mtx.Unlock() @@ -172,9 +169,22 @@ func (hc *healthChecker) registerTarget(t *target) *Status { return t.status } +// retireTargetLocked stops and forgets the active probe registered under name, if there is one, +// so whatever takes the name over never runs beside it. The stop is synchronous: the old probe +// loop has exited by the time it returns. The caller holds the lock. +func (hc *healthChecker) retireTargetLocked(name string) { + if t, ok := hc.targets[name]; ok { + if t != nil { + t.Stop() + } + delete(hc.targets, name) + } +} + func (hc *healthChecker) RegisterVirtual(name, description string) *Status { s := NewStatus(name, description, "", StatusPassing, time.Time{}, nil) hc.mtx.Lock() + hc.retireTargetLocked(name) hc.statuses[name] = s hc.mtx.Unlock() hc.notifyRegistrations() @@ -184,12 +194,14 @@ func (hc *healthChecker) RegisterVirtual(name, description string) *Status { // RegisterExternal records a caller-managed Status (e.g., one driven by a // discovery provider's readiness reporting) so it surfaces in the health // page and status lookups. The caller owns status transitions; no probe is -// started. Remove it with Unregister. +// started, and an active probe already registered under the name is stopped. +// Remove it with Unregister. func (hc *healthChecker) RegisterExternal(name, description string, s *Status) { if name == "" || s == nil { return } hc.mtx.Lock() + hc.retireTargetLocked(name) hc.statuses[name] = s hc.mtx.Unlock() hc.notifyRegistrations() diff --git a/pkg/backends/healthcheck/status.go b/pkg/backends/healthcheck/status.go index 2787bd995..32acb0749 100644 --- a/pkg/backends/healthcheck/status.go +++ b/pkg/backends/healthcheck/status.go @@ -26,6 +26,7 @@ import ( "sync/atomic" "time" + "github.com/trickstercache/trickster/v2/pkg/lb" "github.com/trickstercache/trickster/v2/pkg/observability/logging" "github.com/trickstercache/trickster/v2/pkg/observability/logging/logger" "github.com/trickstercache/trickster/v2/pkg/observability/metrics" @@ -48,8 +49,10 @@ type Status struct { detail string failingSince time.Time subscribers []chan bool - mtx sync.Mutex - prober func(http.ResponseWriter) + // copy-on-write: Set reads it under mtx and calls it outside + onChange []*changeSubscription + mtx sync.Mutex + prober func(http.ResponseWriter) } func NewStatus( @@ -95,17 +98,78 @@ func (s *Status) Headers() http.Header { return h } -// Set updates the status +// Set updates the status. Change callbacks run first, synchronously and outside the Status +// lock, so a pool has republished by the time channel subscribers are notified. func (s *Status) Set(i int32) { - s.status.Store(i) + prev := s.status.Swap(i) s.mtx.Lock() subs := slices.Clone(s.subscribers) + callbacks := s.onChange s.mtx.Unlock() + if prev != i { + for _, c := range callbacks { + s.notifyChange(c, prev, i) + } + } for _, ch := range subs { s.notifySubscriber(ch) } } +type changeSubscription struct { + status *Status + fn func(prev, next int32) + dead atomic.Bool +} + +// Unsubscribe removes the callback without waiting on one that is already running, so it is +// safe to call from inside the callback. A callback captured by a concurrent Set is skipped. +func (c *changeSubscription) Unsubscribe() { + if c.dead.Swap(true) { + return + } + s := c.status + s.mtx.Lock() + defer s.mtx.Unlock() + kept := make([]*changeSubscription, 0, len(s.onChange)) + for _, o := range s.onChange { + if o != c { + kept = append(kept, o) + } + } + s.onChange = kept +} + +// OnChange registers fn to be called with the previous and new status after each change. It +// is called synchronously by Set, outside the Status lock, and recovered on its own. +func (s *Status) OnChange(fn func(prev, next int32)) lb.Subscription { + c := &changeSubscription{status: s, fn: fn} + s.mtx.Lock() + defer s.mtx.Unlock() + next := make([]*changeSubscription, len(s.onChange)+1) + copy(next, s.onChange) + next[len(s.onChange)] = c + s.onChange = next + return c +} + +// notifyChange isolates one callback by recover, so a panic in it reaches neither its +// siblings nor the probe loop calling Set +func (s *Status) notifyChange(c *changeSubscription, prev, next int32) { + if c.dead.Load() { + return + } + safego.Run(func(r any, _ []byte) { + logger.Error("healthcheck status change callback panic", logging.Pairs{ + "target": s.name, + "panic": fmt.Sprintf("%v", r), + }) + metrics.HealthcheckStatusNotifyPanicRecovered.WithLabelValues(s.name).Inc() + }, func() { + c.fn(prev, next) + }) +} + // notifySubscriber sends a non-blocking notification to ch. Each send is // isolated by recover so a closed-channel panic on one subscriber does not // stop notifying the rest or propagate up to the probe loop caller. diff --git a/pkg/backends/healthcheck/status_onchange_test.go b/pkg/backends/healthcheck/status_onchange_test.go new file mode 100644 index 000000000..826e1b499 --- /dev/null +++ b/pkg/backends/healthcheck/status_onchange_test.go @@ -0,0 +1,145 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package healthcheck + +import ( + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +var ( + _ lb.Health = (*Status)(nil) + _ lb.Notifier = (*Status)(nil) +) + +type transition struct{ prev, next int32 } + +func TestStatusOnChangeReportsTransitions(t *testing.T) { + s := &Status{} + var got []transition + sub := s.OnChange(func(prev, next int32) { got = append(got, transition{prev, next}) }) + + s.Set(StatusPassing) + s.Set(StatusPassing) // not a change + s.Set(StatusFailing) + want := []transition{{StatusUnchecked, StatusPassing}, {StatusPassing, StatusFailing}} + if len(got) != len(want) || got[0] != want[0] || got[1] != want[1] { + t.Fatalf("transitions = %v, want %v", got, want) + } + + sub.Unsubscribe() + sub.Unsubscribe() + s.Set(StatusPassing) + if len(got) != len(want) { + t.Errorf("an unsubscribed callback ran: %v", got) + } + if len(s.onChange) != 0 { + t.Errorf("%d subscriptions left behind", len(s.onChange)) + } +} + +// the callback runs outside the Status lock: it can read the Status, register, and +// unsubscribe itself without deadlocking +func TestStatusOnChangeCallbackMayReenter(t *testing.T) { + s := &Status{} + var sub lb.Subscription + var calls, late int + sub = s.OnChange(func(_, _ int32) { + calls++ + _ = s.Detail() + s.SetDetail("seen") + s.OnChange(func(_, _ int32) { late++ }) + sub.Unsubscribe() + }) + done := make(chan struct{}) + go func() { + defer close(done) + s.Set(StatusPassing) + s.Set(StatusFailing) + }() + select { + case <-done: + case <-time.After(3 * time.Second): + t.Fatal("Set deadlocked on a re-entrant callback") + } + if calls != 1 || late != 1 { + t.Errorf("calls = %d, registered-during-callback calls = %d", calls, late) + } +} + +// a panicking callback reaches neither its siblings, the channel subscribers, nor the caller +func TestStatusOnChangePanicIsIsolated(t *testing.T) { + s := NewStatus("panicky", "", "", StatusUnchecked, time.Time{}, nil) + var before, after int + s.OnChange(func(_, _ int32) { before++ }) + s.OnChange(func(_, _ int32) { panic("callback blew up") }) + s.OnChange(func(_, _ int32) { after++ }) + ch := make(chan bool, 1) + s.RegisterSubscriber(ch) + + s.Set(StatusPassing) + s.Set(StatusFailing) + if before != 2 || after != 2 { + t.Errorf("sibling callbacks ran %d and %d times, want 2 and 2", before, after) + } + select { + case <-ch: + default: + t.Error("the channel subscriber was not notified") + } + if s.Get() != StatusFailing { + t.Error("the status was not stored") + } +} + +// Unsubscribe racing Set neither deadlocks nor lets the callback run once it has returned +// and an in-flight Set has drained +func TestStatusOnChangeUnsubscribeRacesSet(t *testing.T) { + for range 200 { + s := &Status{} + var stopped atomic.Bool + var lateCalls atomic.Int32 + entered := make(chan struct{}, 1) + sub := s.OnChange(func(_, _ int32) { + select { + case entered <- struct{}{}: + default: + } + if stopped.Load() { + lateCalls.Add(1) + } + }) + var wg sync.WaitGroup + wg.Go(func() { + for i := range 100 { + s.Set(int32(i%2) - 1) + } + }) + <-entered + sub.Unsubscribe() + wg.Wait() + stopped.Store(true) + s.Set(StatusPassing) + s.Set(StatusFailing) + if lateCalls.Load() != 0 { + t.Fatal("a callback ran after Unsubscribe and the racing Set had both returned") + } + } +} diff --git a/pkg/backends/influxdb/native_listener.go b/pkg/backends/influxdb/native_listener.go index 1449d592a..5f4864a7a 100644 --- a/pkg/backends/influxdb/native_listener.go +++ b/pkg/backends/influxdb/native_listener.go @@ -88,6 +88,10 @@ func (nativeListenerAdapter) ValidateUserRouter(*config.Config, string, *bo.Opti return errors.New("InfluxDB Flight SQL user routing is not supported") } +func (nativeListenerAdapter) ValidateBalancer(*config.Config, string, *bo.Options) error { + return errors.New("InfluxDB Flight SQL session balancing is not supported") +} + func (nativeListenerAdapter) RouteResolver(native.BuildRequest) backends.RouteResolver { return nil } // flightUpstreamAddress resolves the upstream Flight SQL host:port for a diff --git a/pkg/backends/influxdb/native_listener_test.go b/pkg/backends/influxdb/native_listener_test.go index 661d246d0..bb2a81f07 100644 --- a/pkg/backends/influxdb/native_listener_test.go +++ b/pkg/backends/influxdb/native_listener_test.go @@ -80,6 +80,9 @@ func TestFlightNativeListenerAdapterContract(t *testing.T) { if err := adapter.ValidateUserRouter(nil, "flight", nil); err == nil { t.Fatal("ValidateUserRouter() succeeded; Flight SQL has no user routing") } + if err := adapter.ValidateBalancer(nil, "flight", nil); err == nil { + t.Fatal("ValidateBalancer() succeeded; Flight SQL has no session balancing") + } if resolver := adapter.RouteResolver(native.BuildRequest{}); resolver != nil { t.Fatal("RouteResolver() returned a resolver") } diff --git a/pkg/backends/mysql/native_balancer_test.go b/pkg/backends/mysql/native_balancer_test.go new file mode 100644 index 000000000..b971bc28f --- /dev/null +++ b/pkg/backends/mysql/native_balancer_test.go @@ -0,0 +1,245 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package mysql + +import ( + "context" + "net" + "net/netip" + "slices" + "strings" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/backends" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/config" + + vtmysql "vitess.io/vitess/go/mysql" +) + +// balancedTestConfig is the routed test config with its user router replaced by a strategy +// that balances the listener's sessions over both targets +func balancedTestConfig(mechanism string) *config.Config { + c := routedRestartTestConfig() + lb := c.Backends["mysql-users"] + lb.ALBOptions = ao.New() + lb.ALBOptions.MechanismName = mechanism + lb.ALBOptions.Pool = ao.Members("mysql-b", "mysql-a", "mysql-b") + return c +} + +func TestNativeBalancerIsRecognized(t *testing.T) { + c := balancedTestConfig("rr") + lb := c.Backends["mysql-users"] + if !isNativeBalancer(c, lb) || !isNativeRouter(c, lb) || isNativeUserRouter(lb) { + t.Fatal("a round robin pool of mysql backends is not a session balancer") + } + if got := routeTargetNames(lb); !slices.Equal(got, []string{"mysql-a", "mysql-b"}) { + t.Fatalf("route targets = %v", got) + } + a := nativeListenerAdapter{} + if err := a.ValidateBalancer(c, "mysql-users", lb); err != nil { + t.Fatal(err) + } + for name, breakIt := range map[string]func(*config.Config){ + "nil config": nil, + "fanout mechanism": func(c *config.Config) { c.Backends["mysql-users"].ALBOptions.MechanismName = "fr" }, + "latency mechanism": func(c *config.Config) { c.Backends["mysql-users"].ALBOptions.MechanismName = "lt" }, + "router mechanism": func(c *config.Config) { c.Backends["mysql-users"].ALBOptions.MechanismName = "ur" }, + "empty pool": func(c *config.Config) { c.Backends["mysql-users"].ALBOptions.Pool = nil }, + "missing member": func(c *config.Config) { delete(c.Backends, "mysql-b") }, + "foreign member": func(c *config.Config) { c.Backends["mysql-b"].Provider = providers.ReverseProxyShort }, + "not a load balancer": func(c *config.Config) { c.Backends["mysql-users"].Provider = providers.MySQL }, + } { + broken := balancedTestConfig("rr") + if breakIt == nil { + if isNativeBalancer(nil, broken.Backends["mysql-users"]) { + t.Errorf("%s: recognized", name) + } + continue + } + breakIt(broken) + if isNativeBalancer(broken, broken.Backends["mysql-users"]) { + t.Errorf("%s: recognized as a session balancer", name) + } + if err := a.ValidateBalancer(broken, "mysql-users", broken.Backends["mysql-users"]); err == nil { + t.Errorf("%s: validated", name) + } + } + // the listener authenticates its own clients, so the load balancer must name them + anonymous := balancedTestConfig("rr") + anonymous.Backends["mysql-users"].AuthOptions = nil + err := a.ValidateBalancer(anonymous, "mysql-users", anonymous.Backends["mysql-users"]) + if err == nil || !strings.Contains(err.Error(), "authenticator_name") { + t.Fatalf("a load balancer with no listener-facing users: %v", err) + } +} + +func TestNativeBalancerBuildsARoutedServer(t *testing.T) { + a := nativeListenerAdapter{} + c := balancedTestConfig("lc") + protocolConfig, routed, err := nativeProtocolConfig(c, "mysql-users") + if err != nil || !routed || protocolConfig.DownstreamUsers["alice"] == "" { + t.Fatalf("nativeProtocolConfig = %+v, %v, %v", protocolConfig, routed, err) + } + first, err := a.Describe(c, "mysql-users") + if err != nil { + t.Fatal(err) + } + // the strategy and its weights are swapped on reload; the set of targets is not + reweighted := balancedTestConfig("p2c") + reweighted.Backends["mysql-users"].ALBOptions.Pool[1].Weight = 5 + if again, _ := a.Describe(reweighted, "mysql-users"); again.RestartKey != first.RestartKey { + t.Error("a change of strategy or weight restarts the listener") + } + smaller := balancedTestConfig("lc") + smaller.Backends["mysql-users"].ALBOptions.Pool = ao.Members("mysql-a") + if again, _ := a.Describe(smaller, "mysql-users"); again.RestartKey == first.RestartKey { + t.Error("a change of targets does not restart the listener") + } + + request := nativeTestBuildRequest(c, "mysql-users", nil) + request.BackendClients = backends.Backends{ + "mysql-users": &nativeRuntimeBackend{resolver: staticRouteResolver{}}, + "mysql-a": &nativeRuntimeBackend{protocolConfig: ProtocolConfig{BackendName: "mysql-a"}}, + "mysql-b": &nativeRuntimeBackend{protocolConfig: ProtocolConfig{BackendName: "mysql-b"}}, + } + resolver, targets := nativeRouteRuntime(request) + if resolver == nil || len(targets) != 2 { + t.Fatalf("nativeRouteRuntime = %v, %v", resolver, targets) + } + server, err := a.Build(request) + if err != nil { + t.Fatal(err) + } + if err := server.Shutdown(context.Background()); err != nil { + t.Fatal(err) + } +} + +type countedResolver struct { + target backends.Backend + unavailable bool + input backends.RouteInput + released int +} + +type failingStatus struct{} + +func (failingStatus) Get() int32 { return -1 } + +func (r *countedResolver) ResolveRoute(in backends.RouteInput) (backends.RouteDecision, bool) { + r.input = in + d := backends.RouteDecision{ + Target: backends.RouteTarget{Backend: r.target}, Release: func() { r.released++ }, + } + if r.unavailable { + d.Target.Status = failingStatus{} + } + return d, true +} + +func TestRoutedSessionsAreReleasedOnce(t *testing.T) { + salt := []byte("12345678901234567890") + response := vtmysql.ScrambleMysqlNativePassword(salt, []byte("password")) + users := map[string]string{"client": "password"} + peer := &net.TCPAddr{IP: net.ParseIP("::ffff:198.51.100.7"), Port: 40000} + + t.Run("closed after activation", func(t *testing.T) { + routed, _ := newRoutedHandlerForTest(t) + r := &countedResolver{target: newRouteBackend(t, "mysql-a")} + c := &vtmysql.Conn{ConnectionID: 21} + routed.setControl(c.ConnectionID, newTestControl(t)) + routed.NewConnection(c) + if _, err := newCredentialAuth(users, "lb", r).UserEntryWithHash(c, salt, "client", response, peer); err != nil { + t.Fatal(err) + } + if r.input.Client != netip.MustParseAddr("198.51.100.7") || r.input.Username != "client" { + t.Errorf("route input = %+v", r.input) + } + if _, err := routed.activate(c); err != nil { + t.Fatal(err) + } + if r.released != 0 { + t.Fatal("an open session was released") + } + routed.ConnectionClosed(c) + if r.released != 1 { + t.Fatalf("released %d times", r.released) + } + }) + t.Run("closed before its first command", func(t *testing.T) { + routed, _ := newRoutedHandlerForTest(t) + r := &countedResolver{target: newRouteBackend(t, "mysql-a")} + c := &vtmysql.Conn{ConnectionID: 22} + routed.NewConnection(c) + if _, err := newCredentialAuth(users, "lb", r).UserEntryWithHash(c, salt, "client", response, nil); err != nil { + t.Fatal(err) + } + routed.ConnectionClosed(c) + if r.released != 1 { + t.Fatalf("released %d times", r.released) + } + }) + t.Run("target unknown to the listener", func(t *testing.T) { + routed, _ := newRoutedHandlerForTest(t) + r := &countedResolver{target: newRouteBackend(t, "mysql-missing")} + c := &vtmysql.Conn{ConnectionID: 23} + routed.NewConnection(c) + if _, err := newCredentialAuth(users, "lb", r).UserEntryWithHash(c, salt, "client", response, nil); err != nil { + t.Fatal(err) + } + if _, err := routed.activate(c); err == nil { + t.Fatal("activated an unknown target") + } + routed.ConnectionClosed(c) + if r.released != 1 { + t.Fatalf("released %d times", r.released) + } + }) + t.Run("target unavailable", func(t *testing.T) { + r := &countedResolver{target: newRouteBackend(t, "mysql-a"), unavailable: true} + if _, err := newCredentialAuth(users, "lb", r).UserEntryWithHash(&vtmysql.Conn{}, salt, "client", response, nil); err == nil { + t.Fatal("authenticated onto an unavailable target") + } + if r.released != 1 { + t.Fatalf("released %d times", r.released) + } + }) +} + +type namedAddr string + +func (namedAddr) Network() string { return "test" } + +func (a namedAddr) String() string { return string(a) } + +func TestClientAddr(t *testing.T) { + for want, remote := range map[string]net.Addr{ + "198.51.100.7": &net.TCPAddr{IP: net.ParseIP("198.51.100.7"), Port: 1}, + "2001:db8::1": namedAddr("[2001:db8::1]:3306"), + "203.0.113.4": namedAddr("[::ffff:203.0.113.4]:3306"), + } { + if got := clientAddr(remote); got != netip.MustParseAddr(want) { + t.Errorf("clientAddr(%v) = %v, want %s", remote, got, want) + } + } + if clientAddr(nil).IsValid() || clientAddr(namedAddr("pipe")).IsValid() { + t.Error("an address that is not an IP's produced one") + } +} diff --git a/pkg/backends/mysql/native_listener.go b/pkg/backends/mysql/native_listener.go index 4bc846846..c672b95ec 100644 --- a/pkg/backends/mysql/native_listener.go +++ b/pkg/backends/mysql/native_listener.go @@ -24,6 +24,8 @@ import ( "strings" "github.com/trickstercache/trickster/v2/pkg/backends" + albregistry "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/registry" + albtypes "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" uropt "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/ur/options" "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" mo "github.com/trickstercache/trickster/v2/pkg/backends/mysql/options" @@ -124,6 +126,17 @@ func (nativeListenerAdapter) ValidateUserRouter(c *config.Config, name string, b return nil } +func (nativeListenerAdapter) ValidateBalancer(c *config.Config, name string, backend *bo.Options) error { + if !isNativeBalancer(c, backend) { + return fmt.Errorf("mysql load balancer %q requires a mechanism that balances sessions "+ + "over a pool of direct mysql backends", name) + } + if _, err := DownstreamCredentialsFromOptions(backend); err != nil { + return fmt.Errorf("mysql load balancer %q: %w", name, err) + } + return nil +} + func (nativeListenerAdapter) Describe(c *config.Config, listenerName string) (native.Descriptor, error) { protocolConfig, _, err := nativeProtocolConfig(c, listenerName) if err != nil { @@ -179,7 +192,7 @@ func nativeProtocolConfig(c *config.Config, listenerName string) (*ProtocolConfi if !o.UsesListener(listenerName) { continue } - if isNativeUserRouter(o) { + if isNativeRouter(c, o) { users, err := DownstreamCredentialsFromOptions(o) if err != nil { return nil, false, err @@ -204,12 +217,33 @@ func nativeProtocolConfig(c *config.Config, listenerName string) (*ProtocolConfi return nil, false, nil } +// isNativeRouter reports whether o is an ALB that commits each MySQL session to a backend: +// by the user it authenticated as, or by a strategy that balances sessions over a pool +func isNativeRouter(c *config.Config, o *bo.Options) bool { + return isNativeUserRouter(o) || isNativeBalancer(c, o) +} + func isNativeUserRouter(o *bo.Options) bool { return o != nil && o.Provider == providers.ALB && o.ALBOptions != nil && o.ALBOptions.MechanismName == names.MechanismUR && o.ALBOptions.UserRouter != nil && o.ALBOptions.UserRouter.TargetProvider == providers.MySQL } +func isNativeBalancer(c *config.Config, o *bo.Options) bool { + if c == nil || o == nil || o.Provider != providers.ALB || o.ALBOptions == nil || + o.ALBOptions.UserRouter != nil || len(o.ALBOptions.Pool) == 0 || + o.ALBOptions.MechanismName == names.MechanismUR || + !albregistry.Supports(o.ALBOptions.MechanismName, albtypes.PlaneNative) { + return false + } + for _, m := range o.ALBOptions.Pool { + if member := c.Backends[m.Name]; member == nil || member.Provider != providers.MySQL { + return false + } + } + return true +} + type routeResolverProvider interface { RouteResolver() backends.RouteResolver } @@ -221,7 +255,7 @@ type nativeRouteProvider interface { func nativeRouteRuntime(request native.BuildRequest) (backends.RouteResolver, map[string]ProtocolConfig) { routerName := backendForListener(request.Config, request.ListenerName) routerOptions := request.Config.Backends[routerName] - if !isNativeUserRouter(routerOptions) { + if !isNativeRouter(request.Config, routerOptions) { return nil, nil } client := request.BackendClients.Get(routerName) @@ -264,9 +298,15 @@ func backendForListener(c *config.Config, listenerName string) string { } func routeTargetNames(o *bo.Options) []string { - if o == nil || o.ALBOptions == nil || o.ALBOptions.UserRouter == nil { + if o == nil || o.ALBOptions == nil { return nil } + if o.ALBOptions.UserRouter == nil { + // a load balancer's sessions may be committed to any member of its pool + result := o.ALBOptions.Pool.Names() + slices.Sort(result) + return slices.Compact(result) + } seen := make(map[string]struct{}) if name := o.ALBOptions.UserRouter.DefaultBackend; name != "" { seen[name] = struct{}{} diff --git a/pkg/backends/mysql/protocol.go b/pkg/backends/mysql/protocol.go index 8de46f296..46453c009 100644 --- a/pkg/backends/mysql/protocol.go +++ b/pkg/backends/mysql/protocol.go @@ -24,6 +24,7 @@ import ( "fmt" "maps" "net" + "net/netip" "net/url" "os" "slices" @@ -546,7 +547,7 @@ func (a *credentialAuth) HandleUser(user string) bool { } func (a *credentialAuth) UserEntryWithHash(c *vtmysql.Conn, salt []byte, user string, - authResponse []byte, _ net.Addr, + authResponse []byte, remote net.Addr, ) (vtmysql.Getter, error) { password, ok := a.users[user] expected := vtmysql.ScrambleMysqlNativePassword(salt, password) @@ -557,9 +558,10 @@ func (a *credentialAuth) UserEntryWithHash(c *vtmysql.Conn, salt []byte, user st } if a.resolver != nil { decision, resolved := a.resolver.ResolveRoute(backends.RouteInput{ - RouterName: a.backend, Username: user, Authenticated: true, + RouterName: a.backend, Username: user, Authenticated: true, Client: clientAddr(remote), }) if !resolved || !decision.Target.Available() { + releaseRoute(decision) outcome := decision.Outcome if outcome == "" { outcome = backends.RouteOutcomeNoRoute @@ -616,6 +618,29 @@ type upstreamSession struct { // target for the same connection. type routedConnection struct { target *protocolHandler + // release returns the session to the resolver that counted it; nil when it counts none + release func() +} + +func releaseRoute(decision backends.RouteDecision) { + if decision.Release != nil { + decision.Release() + } +} + +// clientAddr is the address a session arrived from, or the zero Addr when it is not an IP's +func clientAddr(remote net.Addr) netip.Addr { + switch a := remote.(type) { + case *net.TCPAddr: + return a.AddrPort().Addr().Unmap() + case nil: + return netip.Addr{} + default: + if ap, err := netip.ParseAddrPort(a.String()); err == nil { + return ap.Addr().Unmap() + } + } + return netip.Addr{} } // routedProtocolHandler adapts Vitess's protocol-specific callbacks to a @@ -675,7 +700,8 @@ func (h *routedProtocolHandler) activate(c *vtmysql.Conn) (*protocolHandler, err return routed.target, nil } var target *protocolHandler - if decision, ok := c.ClientData.(backends.RouteDecision); ok && decision.Target.Backend != nil { + decision, _ := c.ClientData.(backends.RouteDecision) + if decision.Target.Backend != nil { target = h.targets[decision.Target.Backend.Name()] } control := h.takeControl(c.ConnectionID) @@ -683,13 +709,14 @@ func (h *routedProtocolHandler) activate(c *vtmysql.Conn) (*protocolHandler, err // Activation failure is terminal for the connection. Recording it // releases the pending control and blocks a second target selection. c.ClientData = &routedConnection{} + releaseRoute(decision) c.MarkForClose() return nil, errNoRoute() } if control != nil { target.setControl(c.ConnectionID, control) } - c.ClientData = &routedConnection{target: target} + c.ClientData = &routedConnection{target: target, release: decision.Release} target.NewConnection(c) return target, nil } @@ -729,6 +756,15 @@ func (h *routedProtocolHandler) ConnectionClosed(c *vtmysql.Conn) { if target, err := h.target(c); err == nil { target.ConnectionClosed(c) } + switch routed := c.ClientData.(type) { + case *routedConnection: + if routed.release != nil { + routed.release() + } + case backends.RouteDecision: + // the session ended between its authentication and its first command + releaseRoute(routed) + } h.mtx.Lock() delete(h.controls, c.ConnectionID) h.mtx.Unlock() diff --git a/pkg/backends/options/options.go b/pkg/backends/options/options.go index 46fe7ca94..0206a86a0 100644 --- a/pkg/backends/options/options.go +++ b/pkg/backends/options/options.go @@ -744,6 +744,9 @@ func (l Lookup) ValidateConfigMappings(c co.Lookup, ncl negative.Lookups, if err := o.ALBOptions.ValidatePool(o.Name, l.Keys()); err != nil { return err } + if _, err := o.ALBOptions.Validate(); err != nil { + return fmt.Errorf("invalid alb options for backend %q: %w", o.Name, err) + } for _, m := range o.ALBOptions.Pool { if t, ok := l[m.Name]; ok && t != nil && t.IsTemplate { return NewErrTemplatePoolMember(m.Name, o.Name) @@ -937,7 +940,7 @@ func (o *Options) Initialize(name string) error { } if o.Provider == providers.ALB { if o.ALBOptions != nil { - if err := o.ALBOptions.Initialize(""); err != nil { + if err := o.ALBOptions.Initialize(o.Name); err != nil { return err } } diff --git a/pkg/backends/reverseproxy/health.go b/pkg/backends/reverseproxy/health.go index 1a0c3199e..27abce04b 100644 --- a/pkg/backends/reverseproxy/health.go +++ b/pkg/backends/reverseproxy/health.go @@ -17,15 +17,57 @@ package reverseproxy import ( + "context" + "net" + + "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" ) +// Origin schemes of a member that a stream listener relays to rather than an HTTP origin +const ( + schemeTCP = "tcp" + schemeUDP = "udp" +) + // DefaultHealthCheckConfig returns the default HealthCheck Config for this backend provider func (c *Client) DefaultHealthCheckConfig() *ho.Options { o := ho.New() u := c.BaseUpstreamURL() + if u.Scheme == schemeTCP || u.Scheme == schemeUDP { + // a stream member is not probed with a request, so it has no request to describe + return o + } o.Scheme = u.Scheme o.Host = u.Host o.Path = u.Path return o } + +// HealthCheckProbe returns the probe of a tcp origin: the member is healthy when a connection +// to it can be opened. It is nil for every other origin, which is probed with an HTTP request. +func (c *Client) HealthCheckProbe() healthcheck.Probe { + u := c.BaseUpstreamURL() + if u == nil || u.Scheme != schemeTCP || u.Host == "" { + return nil + } + addr := u.Host + return func(ctx context.Context) error { + var d net.Dialer + conn, err := d.DialContext(ctx, schemeTCP, addr) + if err != nil { + return err + } + return conn.Close() + } +} + +// HealthCheckUnsupported names why an origin cannot be actively probed, or is empty. A udp +// origin has no handshake to test and no request to send: whether it is up shows only in +// discovery readiness and in the datagrams it refuses. +func (c *Client) HealthCheckUnsupported() string { + if u := c.BaseUpstreamURL(); u != nil && u.Scheme == schemeUDP { + return "a udp origin has no generic health probe" + } + return "" +} diff --git a/pkg/backends/reverseproxy/health_test.go b/pkg/backends/reverseproxy/health_test.go index c7b077d9c..e2935d9ef 100644 --- a/pkg/backends/reverseproxy/health_test.go +++ b/pkg/backends/reverseproxy/health_test.go @@ -17,7 +17,10 @@ package reverseproxy import ( + "context" + "net" "testing" + "time" bo "github.com/trickstercache/trickster/v2/pkg/backends/options" @@ -34,3 +37,65 @@ func TestDefaultHealthCheckConfig(t *testing.T) { t.Error("expected / for path", dho.Path) } } + +func streamClient(t *testing.T, originURL string) *Client { + t.Helper() + o := bo.New() + o.OriginURL = originURL + require.NoError(t, o.Initialize("stream-member")) + c, err := NewClient("stream-member", o, nil, nil, nil, nil) + require.NoError(t, err) + return c.(*Client) +} + +// a tcp member is healthy when a connection to it opens; nothing is sent +func TestTCPConnectProbe(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + accepted := make(chan struct{}, 4) + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + accepted <- struct{}{} + _ = conn.Close() + } + }() + c := streamClient(t, "tcp://"+ln.Addr().String()) + probe := c.HealthCheckProbe() + require.NotNil(t, probe) + require.Empty(t, c.HealthCheckUnsupported()) + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + require.NoError(t, probe(ctx)) + <-accepted + + // the probe takes no HTTP request options, or the registrar would refuse it + dho := c.DefaultHealthCheckConfig() + require.Empty(t, dho.Scheme) + require.Empty(t, dho.Host) + require.Empty(t, dho.Path) + + _ = ln.Close() + require.Error(t, probe(ctx), "nothing is listening any more") + expired, stop := context.WithCancel(context.Background()) + stop() + require.Error(t, probe(expired)) +} + +// every other origin keeps its request probe; a udp origin has none to run at all +func TestProbeByOrigin(t *testing.T) { + http := streamClient(t, "http://127.0.0.1:9090/base") + require.Nil(t, http.HealthCheckProbe()) + require.Empty(t, http.HealthCheckUnsupported()) + dho := http.DefaultHealthCheckConfig() + require.Equal(t, "http", dho.Scheme) + require.Equal(t, "127.0.0.1:9090", dho.Host) + + udp := streamClient(t, "udp://127.0.0.1:5353") + require.Nil(t, udp.HealthCheckProbe()) + require.NotEmpty(t, udp.HealthCheckUnsupported()) + require.Empty(t, udp.DefaultHealthCheckConfig().Scheme) +} diff --git a/pkg/backends/routes.go b/pkg/backends/routes.go index fe74c9ced..b03052b61 100644 --- a/pkg/backends/routes.go +++ b/pkg/backends/routes.go @@ -16,6 +16,8 @@ package backends +import "net/netip" + // RouteHealthStatus exposes the protocol-neutral health state used to admit a // route target. It is intentionally smaller than healthcheck.Status so route // selection does not depend on a particular health-check transport. @@ -42,6 +44,8 @@ type RouteInput struct { Username string Credential string Authenticated bool + // Client is the address the session arrived from, when the protocol adapter knows it + Client netip.Addr // FallbackOnMappedUnavailable preserves HTTP User Router availability // semantics. Session protocols leave it false to prevent cross-target failover. FallbackOnMappedUnavailable bool @@ -66,6 +70,9 @@ type RouteDecision struct { OutboundUsername string OutboundCredential string ReplaceCredentials bool + // Release, when set, must be called once when the routed session ends, however it ends. + // A resolver that balances sessions across targets counts the ones in progress by it. + Release func() } // RouteResolver selects a runtime backend target for an authenticated identity. diff --git a/pkg/config/loader.go b/pkg/config/loader.go index 3f3881ade..0799507fe 100644 --- a/pkg/config/loader.go +++ b/pkg/config/loader.go @@ -96,6 +96,11 @@ func LoadWithOverlay(args []string, overlay *Overlay) (*Config, error) { if err != nil { return nil, err } + if o.ALBOptions != nil { + if w := o.ALBOptions.PoolRepeatWarning(k); w != "" { + c.addLoaderWarning(w) + } + } } if len(c.Discovery) > 0 { diff --git a/pkg/config/loader_pool_repeats_test.go b/pkg/config/loader_pool_repeats_test.go new file mode 100644 index 000000000..8286bccc9 --- /dev/null +++ b/pkg/config/loader_pool_repeats_test.go @@ -0,0 +1,75 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package config + +import ( + "os" + "path/filepath" + "strings" + "testing" + + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + + "github.com/stretchr/testify/require" +) + +const repeatedPoolConfig = ` +backends: + a: + provider: reverseproxycache + origin_url: http://127.0.0.1:9001 + b: + provider: reverseproxycache + origin_url: http://127.0.0.1:9002 + lb: + provider: alb + alb: + mechanism: rr + pool: POOL +` + +func loadPoolConfig(t *testing.T, pool string) (*Config, error) { + t.Helper() + path := filepath.Join(t.TempDir(), "trickster.yaml") + yml := strings.Replace(repeatedPoolConfig, "POOL", pool, 1) + require.NoError(t, os.WriteFile(path, []byte(yml), 0o600)) + return Load([]string{"-config", path}) +} + +func TestLoadDedupesRepeatedPoolMembers(t *testing.T) { + c, err := loadPoolConfig(t, "[a, a, b]") + require.NoError(t, err) + require.Equal(t, ao.Members("a", "b"), c.Backends["lb"].ALBOptions.Pool) + var warnings []string + for _, w := range c.LoaderWarnings { + if strings.Contains(w, "repeating a pool member") { + warnings = append(warnings, w) + } + } + require.Len(t, warnings, 1, "one warning per alb") + require.Contains(t, warnings[0], `alb "lb"`) + require.Contains(t, warnings[0], "{name: a, weight: 2}") + + c, err = loadPoolConfig(t, "[{name: a, weight: 2}, b]") + require.NoError(t, err) + for _, w := range c.LoaderWarnings { + require.NotContains(t, w, "repeating a pool member") + } + + _, err = loadPoolConfig(t, "[{name: a, weight: 2}, b, {name: a, weight: 3}]") + require.ErrorIs(t, err, ao.ErrConflictingPoolWeights) +} diff --git a/pkg/config/types/statusranges.go b/pkg/config/types/statusranges.go new file mode 100644 index 000000000..828cfd2ea --- /dev/null +++ b/pkg/config/types/statusranges.go @@ -0,0 +1,116 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package types + +import ( + "errors" + "fmt" + + "go.yaml.in/yaml/v3" +) + +const ( + minStatusCode = 100 + maxStatusCode = 599 +) + +// ErrInvalidStatusRange is returned for a status code range outside 100-599 or with its +// start above its end. +var ErrInvalidStatusRange = errors.New("invalid status code range") + +// StatusRange is an inclusive range of HTTP status codes. In YAML it is either a bare code: +// +// status_codes: [200, 204] +// +// or a mapping, and a list may mix the two: +// +// status_codes: [{start: 200, end: 299}, 304] +type StatusRange struct { + Start int `yaml:"start"` + End int `yaml:"end"` +} + +// StatusRanges is a list of status code ranges. Ranges may overlap or touch; the list means +// their union. +type StatusRanges []StatusRange + +// StatusCodes returns a StatusRanges of single codes. +func StatusCodes(codes ...int) StatusRanges { + out := make(StatusRanges, len(codes)) + for i, c := range codes { + out[i] = StatusRange{Start: c, End: c} + } + return out +} + +// UnmarshalYAML accepts either a bare status code or a {start, end} mapping. +func (r *StatusRange) UnmarshalYAML(value *yaml.Node) error { + if value.Kind == yaml.ScalarNode { + var code int + if err := value.Decode(&code); err != nil { + return err + } + r.Start, r.End = code, code + return nil + } + type loadStatusRange StatusRange + var lr loadStatusRange + if err := value.Decode(&lr); err != nil { + return err + } + *r = StatusRange(lr) + return nil +} + +// MarshalYAML renders a single-code range as a bare code, so a list of codes reads back out +// the way it was written. +func (r StatusRange) MarshalYAML() (any, error) { + if r.Start == r.End { + return r.Start, nil + } + type dumpStatusRange StatusRange + return dumpStatusRange(r), nil +} + +// Validate checks that every range lies within 100-599 and starts no higher than it ends. +func (l StatusRanges) Validate() error { + for _, r := range l { + if r.Start < minStatusCode || r.End > maxStatusCode || r.Start > r.End { + return fmt.Errorf("%w: %d-%d (codes are %d-%d, start first)", + ErrInvalidStatusRange, r.Start, r.End, minStatusCode, maxStatusCode) + } + } + return nil +} + +// StatusTable answers whether a status code is in a set with one array index. +type StatusTable [maxStatusCode + 1]bool + +// Compile returns the lookup table of the ranges' union. Anything outside 100-599 is ignored. +func (l StatusRanges) Compile() *StatusTable { + t := &StatusTable{} + for _, r := range l { + for code := max(r.Start, minStatusCode); code <= min(r.End, maxStatusCode); code++ { + t[code] = true + } + } + return t +} + +// Contains reports whether code is in the set; a code outside 0-599 never is. +func (t *StatusTable) Contains(code int) bool { + return t != nil && code >= 0 && code <= maxStatusCode && t[code] +} diff --git a/pkg/config/types/statusranges_test.go b/pkg/config/types/statusranges_test.go new file mode 100644 index 000000000..6ef552a48 --- /dev/null +++ b/pkg/config/types/statusranges_test.go @@ -0,0 +1,106 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package types + +import ( + "testing" + + "github.com/stretchr/testify/require" + "go.yaml.in/yaml/v3" +) + +func TestStatusRangesYAML(t *testing.T) { + var l StatusRanges + require.NoError(t, yaml.Unmarshal([]byte("[{start: 200, end: 299}, 304, {start: 405, end: 405}]"), &l)) + require.Equal(t, StatusRanges{{200, 299}, {304, 304}, {405, 405}}, l) + + // a list of bare codes, the only form that existed before ranges, round-trips unchanged + var codes StatusRanges + require.NoError(t, yaml.Unmarshal([]byte("[200, 204]"), &codes)) + require.Equal(t, StatusCodes(200, 204), codes) + out, err := yaml.Marshal(codes) + require.NoError(t, err) + require.Equal(t, "- 200\n- 204\n", string(out)) + + out, err = yaml.Marshal(l) + require.NoError(t, err) + var again StatusRanges + require.NoError(t, yaml.Unmarshal(out, &again)) + require.Equal(t, l, again) + require.Contains(t, string(out), "- 304\n") + require.Contains(t, string(out), "- 405\n", "a single-code range marshals as a bare code") + require.Contains(t, string(out), "start: 200") + + require.Error(t, yaml.Unmarshal([]byte("[ok]"), &again)) + require.Error(t, yaml.Unmarshal([]byte("[{start: low}]"), &again)) +} + +func TestStatusRangesValidate(t *testing.T) { + require.NoError(t, StatusRanges(nil).Validate()) + require.NoError(t, StatusRanges{{100, 599}, {200, 200}}.Validate()) + for name, l := range map[string]StatusRanges{ + "below 100": {{99, 200}}, + "above 599": {{200, 600}}, + "backwards": {{300, 200}}, + "zero start": {{0, 200}}, + "second item": {{200, 299}, {700, 700}}, + } { + require.ErrorIs(t, l.Validate(), ErrInvalidStatusRange, name) + } +} + +func TestStatusTable(t *testing.T) { + // overlapping and adjacent ranges are one set + table := StatusRanges{{200, 250}, {240, 299}, {300, 304}, {429, 429}}.Compile() + for _, code := range []int{200, 245, 299, 300, 304, 429} { + require.True(t, table.Contains(code), code) + } + for _, code := range []int{-1, 0, 99, 199, 305, 428, 430, 599, 600, 100000} { + require.False(t, table.Contains(code), code) + } + // out-of-range bounds are clipped rather than indexed + wide := StatusRanges{{-5, 120}, {590, 9000}}.Compile() + require.True(t, wide.Contains(100)) + require.True(t, wide.Contains(599)) + require.False(t, wide.Contains(99)) + require.False(t, StatusRanges(nil).Compile().Contains(200)) + var none *StatusTable + require.False(t, none.Contains(200)) + require.Zero(t, testing.AllocsPerRun(100, func() { _ = table.Contains(204) })) +} + +func FuzzStatusRangesUnmarshal(f *testing.F) { + f.Add("[200, 204]") + f.Add("[{start: 200, end: 299}, 304]") + f.Add("[{start: 9, end: -1}]") + f.Add("{}") + f.Fuzz(func(t *testing.T, doc string) { + var l StatusRanges + if err := yaml.Unmarshal([]byte(doc), &l); err != nil { + return + } + // whatever parsed must compile without indexing out of range, valid or not + table := l.Compile() + if l.Validate() != nil { + return + } + for _, r := range l { + if !table.Contains(r.Start) || !table.Contains(r.End) { + t.Errorf("valid range %d-%d is missing from its table", r.Start, r.End) + } + } + }) +} diff --git a/pkg/config/validate/alb_strategies_test.go b/pkg/config/validate/alb_strategies_test.go new file mode 100644 index 000000000..8ddf5704a --- /dev/null +++ b/pkg/config/validate/alb_strategies_test.go @@ -0,0 +1,89 @@ +/* + * Copyright 2018 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package validate + +import ( + "os" + "path/filepath" + "strings" + "testing" + + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + "github.com/trickstercache/trickster/v2/pkg/config" + "github.com/trickstercache/trickster/v2/pkg/config/types" + + "github.com/stretchr/testify/require" +) + +const albStrategyConfig = ` +backends: + a: + provider: reverseproxycache + origin_url: http://127.0.0.1:9001 + b: + provider: reverseproxycache + origin_url: http://127.0.0.1:9002 + lb: + provider: alb + alb: + pool: [a, {name: b, weight: 3}] +ALB +` + +func validateALB(t *testing.T, alb string) error { + t.Helper() + indented := " " + strings.ReplaceAll(strings.TrimSpace(alb), "\n", "\n ") + path := filepath.Join(t.TempDir(), "trickster.yaml") + yml := strings.Replace(albStrategyConfig, "ALB", indented, 1) + require.NoError(t, os.WriteFile(path, []byte(yml), 0o600)) + c, err := config.Load([]string{"-config", path}) + if err != nil { + return err + } + return Validate(c) +} + +func TestALBStrategyConfigs(t *testing.T) { + for name, alb := range map[string]string{ + "p2c": "mechanism: p2c", + "lc": "mechanism: least_connections", + "hrw": "mechanism: hrw", + "hrw keyed": "mechanism: hrw\nhrw:\n key: header:X-Tenant\n ipv6_prefix: 56", + "lt": "mechanism: lt", + "lt tuned": "mechanism: lt\nlt:\n status_codes: [{start: 200, end: 499}]\n decay: 30s\n signal: first_write", + "fgr codes": "mechanism: fgr\nfgr:\n status_codes: [200, 204]", + "fgr ranges": "mechanism: fgr\nfgr:\n status_codes: [{start: 200, end: 299}, 304]", + "rr as today": "mechanism: rr", + } { + require.NoError(t, validateALB(t, alb), name) + } + for name, test := range map[string]struct { + alb string + want error + }{ + "hrw block on rr": {"mechanism: rr\nhrw:\n key: host", ao.ErrHRWOnlyForHRW}, + "lt block on p2c": {"mechanism: p2c\nlt:\n decay: 5s", ao.ErrLTOnlyForLT}, + "bad key source": {"mechanism: hrw\nhrw:\n key: port", ao.ErrInvalidKeySource}, + "bad ipv6 prefix": {"mechanism: hrw\nhrw:\n ipv6_prefix: 200", ao.ErrInvalidIPv6Prefix}, + "stream lt signal": {"mechanism: lt\nlt:\n signal: connect", ao.ErrInvalidLTSignal}, + "bad lt range": {"mechanism: lt\nlt:\n status_codes: [{start: 100, end: 900}]", types.ErrInvalidStatusRange}, + "bad fgr range": {"mechanism: fgr\nfgr:\n status_codes: [42]", types.ErrInvalidStatusRange}, + "output_format": {"mechanism: rr\noutput_format: prometheus", ao.ErrOutputFormatOnlyForTSM}, + } { + require.ErrorIs(t, validateALB(t, test.alb), test.want, name) + } +} diff --git a/pkg/config/validate/listeners_test.go b/pkg/config/validate/listeners_test.go index 6993bbe9d..7872fbd41 100644 --- a/pkg/config/validate/listeners_test.go +++ b/pkg/config/validate/listeners_test.go @@ -32,6 +32,7 @@ import ( autho "github.com/trickstercache/trickster/v2/pkg/proxy/authenticator/options" l4o "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" tlsopts "github.com/trickstercache/trickster/v2/pkg/proxy/tls/options" + "github.com/trickstercache/trickster/v2/pkg/util/sets" ) func mysqlBackend(listenerName string) *bo.Options { @@ -571,6 +572,10 @@ func TestListenersStreamProtocols(t *testing.T) { c.Backends = backends return c } + behindProxy := func(c *config.Config) *config.Config { + c.Listeners["relay"].ProxyProtocol = true + return c + } albBackend := func(mechanism string) *bo.Options { b := bo.New() b.Provider = providers.ALB @@ -578,6 +583,17 @@ func TestListenersStreamProtocols(t *testing.T) { b.ALBOptions = &ao.Options{MechanismName: mechanism, Pool: ao.PoolMemberList{{Name: "m1"}}} return b } + albWith := func(mechanism string, set func(*ao.Options)) *bo.Options { + b := albBackend(mechanism) + set(b.ALBOptions) + return b + } + keyed := func(kind ao.KeyKind, spelled string) func(*ao.Options) { + return func(o *ao.Options) { + o.HRW = ao.HRWOptions{Key: spelled, KeySource: ao.KeySource{Kind: kind, Name: "X"}} + } + } + signal := func(s string) func(*ao.Options) { return func(o *ao.Options) { o.LT.Signal = s } } member := bo.New() member.OriginURL = "tcp://member.example.com:9000" cases := []struct { @@ -604,7 +620,52 @@ func TestListenersStreamProtocols(t *testing.T) { "b": streamBackend("relay", providers.ReverseProxyShort, "X.example.com."), }), "already routed"}, {"wrong_provider", newConfig(listener.ProtocolTCP, bo.Lookup{"p": streamBackend("relay", providers.Prometheus)}), "cannot map to backend"}, - {"alb_not_rr", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("fr"), "m1": member}), "requires alb backend"}, + {"tcp_alb_round_robin", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("round_robin"), "m1": member}), ""}, + {"udp_alb_rr", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albBackend("rr"), "m1": member}), ""}, + // the mechanisms a stream listener may use come from the registry, and the error names them + {"alb_fanout", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("fr"), "m1": member}), + "a mechanism that serves a tcp listener: hrw, lc, lt, p2c, race, rr"}, + {"udp_alb_fanout", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albBackend("fr"), "m1": member}), + "a mechanism that serves a udp listener: hrw, lc, lt, mirror, p2c, rr"}, + // a mechanism that commits a flow to several members serves only the protocols it can + {"tcp_alb_race", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("race"), "m1": member}), ""}, + {"tls_alb_connect_race", newConfig(listener.ProtocolTLS, bo.Lookup{"pool": albBackend("connect_race"), "m1": member}), ""}, + {"udp_alb_mirror", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albBackend("mirror"), "m1": member}), ""}, + {"udp_alb_race", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albBackend("race"), "m1": member}), + "mechanism \"race\" requires a tcp or tls listener"}, + {"tcp_alb_mirror", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("udp_mirror"), "m1": member}), + "mechanism \"udp_mirror\" requires a udp listener"}, + {"alb_router", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("ur"), "m1": member}), "requires alb backend"}, + {"tcp_alb_p2c", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("p2c"), "m1": member}), ""}, + {"udp_alb_lc", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albBackend("lc"), "m1": member}), ""}, + {"tls_alb_lt", newConfig(listener.ProtocolTLS, bo.Lookup{"pool": albBackend("lt"), "m1": member}), ""}, + {"tcp_alb_hrw", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("hrw"), "m1": member}), ""}, + // a key must be something the listener can read: the client address on any of them, + // the server name on tls alone, and nothing of a request + {"tls_hrw_sni", newConfig(listener.ProtocolTLS, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeySNI, "sni")), "m1": member}), ""}, + {"tcp_hrw_sni", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeySNI, "sni")), "m1": member}), + "cannot read alb backend \"pool\"'s hrw.key \"sni\""}, + {"udp_hrw_sni", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeySNI, "sni")), "m1": member}), "cannot read"}, + {"tcp_hrw_header", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeyHeader, "header:X")), "m1": member}), + "use client_ip, sni on a tls listener, or proxy_tlv:"}, + // a PROXY protocol TLV is there to read only where the header is accepted, which udp never does + {"tcp_hrw_tlv", behindProxy(newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeyProxyTLV, "proxy_tlv:0xEA")), "m1": member})), ""}, + {"tls_hrw_tlv", behindProxy(newConfig(listener.ProtocolTLS, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeyProxyTLV, "proxy_tlv:0xEA")), "m1": member})), ""}, + {"tcp_hrw_tlv_no_proxy_protocol", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeyProxyTLV, "proxy_tlv:0xEA")), "m1": member}), + "cannot read alb backend \"pool\"'s hrw.key \"proxy_tlv:0xEA\""}, + {"udp_hrw_tlv", behindProxy(newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albWith("hrw", keyed(ao.KeyProxyTLV, "proxy_tlv:0xEA")), "m1": member})), + "cannot read"}, + // and a latency signal must be one the protocol has + {"tcp_lt_first_byte", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albWith("lt", signal("first_byte")), "m1": member}), ""}, + {"udp_lt_first_reply", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albWith("lt", signal("first_reply")), "m1": member}), ""}, + {"udp_lt_connect", newConfig(listener.ProtocolUDP, bo.Lookup{"pool": albWith("lt", signal("connect")), "m1": member}), + "\"connect\" on a udp listener (use first_reply)"}, + {"tcp_lt_first_write", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albWith("lt", signal("first_write")), "m1": member}), + "use connect or first_byte"}, + {"tcp_alb_stream_block", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albWith("p2c", func(o *ao.Options) { + o.Stream = &ao.StreamOptions{ConnectRetries: 2, PassiveHealth: &ao.PassiveHealthOptions{Failures: 3}} + }), "m1": member}), ""}, + {"alb_unknown_mechanism", newConfig(listener.ProtocolTCP, bo.Lookup{"pool": albBackend("nope"), "m1": member}), "requires alb backend"}, {"unsupported_still_refused", newConfig("sctp", bo.Lookup{"p": streamBackend("relay", providers.ReverseProxyShort)}), "unsupported protocol"}, } for _, tc := range cases { @@ -684,6 +745,95 @@ func TestListenersReservePortsByTransport(t *testing.T) { } } +// a load balancer that serves no stream listener is held to what a request can offer +func TestRequestALBsRefuseStreamOnlySettings(t *testing.T) { + alb := func(set func(*ao.Options), listeners ...string) *config.Config { + c := config.NewConfig() + c.Listeners["relay"] = listener.New("relay") + c.Listeners["relay"].Protocol = listener.ProtocolTLS + c.Listeners["web"] = listener.New("web") + b := bo.New() + b.Provider = providers.ALB + b.ListenerNames = listeners + b.ALBOptions = &ao.Options{MechanismName: "hrw"} + set(b.ALBOptions) + notALB := bo.New() + c.Backends = bo.Lookup{"lb": b, "origin": notALB, "unset": nil} + return c + } + for name, test := range map[string]struct { + set func(*ao.Options) + want string + }{ + "plain": {func(*ao.Options) {}, ""}, + "header key": {func(o *ao.Options) { o.HRW.KeySource = ao.KeySource{Kind: ao.KeyHeader, Name: "X"} }, ""}, + "first_write": {func(o *ao.Options) { o.LT.Signal = ao.LTSignalFirstWrite }, ""}, + "stream block": {func(o *ao.Options) { o.Stream = &ao.StreamOptions{} }, "'stream' options apply only"}, + "sni key": {func(o *ao.Options) { o.HRW = ao.HRWOptions{Key: "sni", KeySource: ao.KeySource{Kind: ao.KeySNI}} }, + "cannot be read from a request"}, + "tlv key": {func(o *ao.Options) { + o.HRW = ao.HRWOptions{Key: "proxy_tlv:5", KeySource: ao.KeySource{Kind: ao.KeyProxyTLV, TLV: 5}} + }, + "cannot be read from a request"}, + "connect signal": {func(o *ao.Options) { o.LT.Signal = ao.LTSignalConnect }, "on a http listener (use first_write)"}, + } { + err := requestALBs(alb(test.set), sets.NewStringSet()) + switch { + case test.want == "" && err != nil: + t.Errorf("%s: unexpected error: %v", name, err) + case test.want != "" && (err == nil || !strings.Contains(err.Error(), test.want)): + t.Errorf("%s: error = %v, want %q", name, err, test.want) + } + // the same load balancer on a stream listener alone is somebody else's to judge + if err := requestALBs(alb(test.set, "relay"), sets.New([]string{"lb"})); err != nil { + t.Errorf("%s: a stream load balancer was held to request rules: %v", name, err) + } + // one that serves requests as well is held to both, but for a stream block, which the + // stream listener it also serves has a use for + err = requestALBs(alb(test.set, "relay", "web"), sets.New([]string{"lb"})) + switch { + case (test.want == "" || name == "stream block") && err != nil: + t.Errorf("%s: on both planes: unexpected error: %v", name, err) + case test.want != "" && name != "stream block" && (err == nil || !strings.Contains(err.Error(), test.want)): + t.Errorf("%s: on both planes: error = %v, want %q", name, err, test.want) + } + } +} + +// a mechanism that commits a flow to several members needs a stream listener of its own +func TestSpreadMechanismsNeedAStreamListener(t *testing.T) { + alb := func(mechanism string, pool ...string) *bo.Options { + b := bo.New() + b.Provider = providers.ALB + b.ALBOptions = &ao.Options{MechanismName: mechanism, Pool: ao.Members(pool...)} + return b + } + c := config.NewConfig() + c.Listeners["relay"] = listener.New("relay") + c.Listeners["relay"].Protocol = listener.ProtocolTCP + c.Backends = bo.Lookup{"racer": alb("race", "origin"), "origin": bo.New()} + if err := requestALBs(c, sets.NewStringSet()); err == nil || + !strings.Contains(err.Error(), "mechanism \"race\" requires a tcp or tls listener") { + t.Errorf("a race on a request listener: %v", err) + } + c.Backends["racer"].ListenerNames = []string{"relay"} + if err := requestALBs(c, sets.New([]string{"racer"})); err != nil { + t.Errorf("a race on a stream listener: %v", err) + } + c.Backends["racer"].ListenerNames = []string{"relay", "default"} + if err := requestALBs(c, sets.New([]string{"racer"})); err == nil || + !strings.Contains(err.Error(), "cannot serve http listener \"default\"") { + t.Errorf("a race on a stream and a request listener: %v", err) + } + c.Backends["racer"].ListenerNames = []string{"relay"} + // and cannot be reached through another load balancer, which has one member to hand it + c.Backends["outer"] = alb("rr", "racer") + if err := requestALBs(c, sets.New([]string{"racer", "outer"})); err == nil || + !strings.Contains(err.Error(), "cannot be a member of another alb's pool") { + t.Errorf("a race as a pool member: %v", err) + } +} + const pgTestListener = "pg1" func postgresListenerConfig(backend *bo.Options) *config.Config { diff --git a/pkg/config/validate/session_balancer_test.go b/pkg/config/validate/session_balancer_test.go new file mode 100644 index 000000000..56aa31f26 --- /dev/null +++ b/pkg/config/validate/session_balancer_test.go @@ -0,0 +1,186 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package validate + +import ( + "strings" + "testing" + + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" + "github.com/trickstercache/trickster/v2/pkg/backends/providers" + "github.com/trickstercache/trickster/v2/pkg/config" + "github.com/trickstercache/trickster/v2/pkg/config/listener" + configtypes "github.com/trickstercache/trickster/v2/pkg/config/types" + autho "github.com/trickstercache/trickster/v2/pkg/proxy/authenticator/options" +) + +// replicaConfig maps a mysql listener to a load balancer over two mysql backends that have +// no listener of their own +func replicaConfig(mechanism string) *config.Config { + c := config.NewConfig() + c.Listeners["mysql1"] = listener.New("mysql1") + c.Listeners["mysql1"].Protocol = listener.ProtocolMySQL + c.Listeners["mysql1"].ListenPort = 8486 + lb := bo.New() + lb.Provider = providers.ALB + lb.ListenerName = "mysql1" + lb.ALBOptions = &ao.Options{MechanismName: mechanism, Pool: ao.Members("replica-a", "replica-b")} + lb.AuthenticatorName = "mysql-listener-clients" + lb.AuthOptions = &autho.Options{Users: configtypes.EnvStringMap{"client": "password"}} + c.Backends = bo.Lookup{"replicas": lb, "replica-a": mysqlBackend(""), "replica-b": mysqlBackend("")} + return c +} + +func TestNativeListenerBalancesSessions(t *testing.T) { + for _, mechanism := range []string{"rr", "round_robin", "p2c", "lc", "hrw"} { + c := replicaConfig(mechanism) + if err := Listeners(c); err != nil { + t.Fatalf("%s: %v", mechanism, err) + } + if !c.Listeners["mysql1"].Active { + t.Errorf("%s: the listener is not active", mechanism) + } + if names := c.Backends["replica-a"].ListenerNames; len(names) != 0 { + t.Errorf("%s: a replica reached through the pool was given listeners %v", mechanism, names) + } + } + keyed := func(key string, kind ao.KeyKind) func(*config.Config) { + return func(c *config.Config) { + c.Backends["replicas"].ALBOptions.HRW = ao.HRWOptions{Key: key, KeySource: ao.KeySource{Kind: kind}} + } + } + for name, test := range map[string]struct { + mechanism string + adjust func(*config.Config) + want string + }{ + "user key": {"hrw", keyed("user", ao.KeyUser), ""}, + "host key": {"hrw", keyed("host", ao.KeyHost), "use client_ip or user"}, + "fanout": {"fr", nil, "routes or balances sessions: hrw, lc, p2c, rr, ur"}, + "no timing": {"lt", nil, "routes or balances sessions"}, + "foreign member": {"rr", func(c *config.Config) { + c.Backends["replica-b"].Provider = providers.ReverseProxyShort + c.Backends["replica-b"].OriginURL = "http://example.com" + }, "to be a mysql backend: \"replica-b\" is not"}, + "missing member": {"rr", func(c *config.Config) { delete(c.Backends, "replica-b") }, "\"replica-b\" is not"}, + "discovery": {"rr", func(c *config.Config) { + c.Backends["replicas"].ALBOptions.Discovery = &ao.DiscoveryOptions{} + }, "'discovery' is not supported on a mysql listener"}, + "stream block": {"rr", func(c *config.Config) { + c.Backends["replicas"].ALBOptions.Stream = &ao.StreamOptions{} + }, "'stream' options apply only"}, + "no listener users": {"rr", func(c *config.Config) { + c.Backends["replicas"].AuthenticatorName, c.Backends["replicas"].AuthOptions = "", nil + }, "requires an authenticator_name"}, + } { + c := replicaConfig(test.mechanism) + if test.adjust != nil { + test.adjust(c) + } + err := Listeners(c) + switch { + case test.want == "" && err != nil: + t.Errorf("%s: %v", name, err) + case test.want != "" && (err == nil || !strings.Contains(err.Error(), test.want)): + t.Errorf("%s: error = %v, want %q", name, err, test.want) + } + } +} + +// the same pool on an http listener is an ordinary load balancer, held to a request's rules +func TestSessionKeysAreForNativeListeners(t *testing.T) { + c := config.NewConfig() + lb := bo.New() + lb.Provider = providers.ALB + lb.ALBOptions = &ao.Options{ + MechanismName: "hrw", Pool: ao.Members("origin"), + HRW: ao.HRWOptions{Key: "user", KeySource: ao.KeySource{Kind: ao.KeyUser}}, + } + origin := bo.New() + origin.Provider = providers.ReverseProxyShort + origin.OriginURL = "http://example.com" + c.Backends = bo.Lookup{"lb": lb, "origin": origin} + if err := Listeners(c); err == nil || !strings.Contains(err.Error(), "hrw.key \"user\"") { + t.Fatalf("error = %v", err) + } +} + +// a load balancer mapped to listeners of two kinds must suit both +func TestALBsOnTwoPlanesAreHeldToBoth(t *testing.T) { + web := func(c *config.Config) { + c.Listeners["web"] = listener.New("web") + c.Listeners["web"].ListenPort = 18480 + } + t.Run("http and mysql", func(t *testing.T) { + for key, want := range map[string]string{ + "client_ip": "", + "user": "hrw.key \"user\" cannot be read from a request, which http listener \"web\" serves", + } { + c := replicaConfig("hrw") + web(c) + lb := c.Backends["replicas"] + lb.ListenerName, lb.ListenerNames = "", []string{"mysql1", "web"} + ks, err := ao.ParseKeySource(key) + if err != nil { + t.Fatal(err) + } + lb.ALBOptions.HRW = ao.HRWOptions{Key: key, KeySource: ks} + err = Listeners(c) + if (want == "") != (err == nil) || (err != nil && !strings.Contains(err.Error(), want)) { + t.Errorf("%s: error = %v, want %q", key, err, want) + } + } + }) + t.Run("http and tls", func(t *testing.T) { + for name, test := range map[string]struct { + mechanism string + set func(*ao.Options) + want string + }{ + "client_ip": {"hrw", func(*ao.Options) {}, ""}, + "sni": {"hrw", func(o *ao.Options) { o.HRW = ao.HRWOptions{Key: "sni", KeySource: ao.KeySource{Kind: ao.KeySNI}} }, + "hrw.key \"sni\" cannot be read from a request"}, + "host": {"hrw", func(o *ao.Options) { o.HRW = ao.HRWOptions{Key: "host", KeySource: ao.KeySource{Kind: ao.KeyHost}} }, + "cannot read alb backend \"lb\"'s hrw.key \"host\""}, + "default signal": {"lt", func(*ao.Options) {}, ""}, + "connect signal": {"lt", func(o *ao.Options) { o.LT.Signal = ao.LTSignalConnect }, "on a http listener"}, + "write signal": {"lt", func(o *ao.Options) { o.LT.Signal = ao.LTSignalFirstWrite }, "on a tls listener"}, + "stream block": {"rr", func(o *ao.Options) { o.Stream = &ao.StreamOptions{ConnectRetries: 1} }, ""}, + "race": {"race", func(*ao.Options) {}, "cannot serve http listener \"web\""}, + } { + c := config.NewConfig() + web(c) + c.Listeners["relay"] = listener.New("relay") + c.Listeners["relay"].Protocol = listener.ProtocolTLS + c.Listeners["relay"].ListenPort = 9443 + lb := bo.New() + lb.Provider = providers.ALB + lb.ListenerNames = []string{"relay", "web"} + lb.ALBOptions = &ao.Options{MechanismName: test.mechanism, Pool: ao.Members("m1")} + test.set(lb.ALBOptions) + member := bo.New() + member.Provider = providers.ReverseProxyShort + member.OriginURL = "tcp://member.example.com:9000" + c.Backends = bo.Lookup{"lb": lb, "m1": member} + err := Listeners(c) + if (test.want == "") != (err == nil) || (err != nil && !strings.Contains(err.Error(), test.want)) { + t.Errorf("%s: error = %v, want %q", name, err, test.want) + } + } + }) +} diff --git a/pkg/config/validate/validate.go b/pkg/config/validate/validate.go index c6b0fe77c..38f9e981e 100644 --- a/pkg/config/validate/validate.go +++ b/pkg/config/validate/validate.go @@ -25,8 +25,10 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends" "github.com/trickstercache/trickster/v2/pkg/backends/alb" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/rr" - albnames "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" + albregistry "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/registry" + albtypes "github.com/trickstercache/trickster/v2/pkg/backends/alb/mech/types" + ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" + bo "github.com/trickstercache/trickster/v2/pkg/backends/options" "github.com/trickstercache/trickster/v2/pkg/backends/providers" providerregistry "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry" "github.com/trickstercache/trickster/v2/pkg/backends/rule" @@ -267,6 +269,8 @@ func Listeners(c *config.Config) error { mappedProviders := make(map[string]map[string]string, len(c.Listeners)) nativeListeners := providerregistry.NativeListeners() nativeTargets := nativeUserRouterTargets(c, nativeListeners) + // the load balancers that serve a stream or native listener; every other one serves requests + streamALBs := sets.NewStringSet() for backendName, backend := range c.Backends { if backend == nil || backend.IsTemplate { // templates are never routed, so they map to no listener @@ -337,7 +341,7 @@ func Listeners(c *config.Config) error { return fmt.Errorf("listener %q configures stream options for protocol %q", name, options.Protocol) } if options.IsStream() { - if err := streamListener(c, name, options, mappedProviders[name]); err != nil { + if err := streamListener(c, name, options, mappedProviders[name], streamALBs); err != nil { return err } } @@ -363,6 +367,14 @@ func Listeners(c *config.Config) error { if provider == providers.ALB && backend != nil && backend.ALBOptions != nil && backend.ALBOptions.UserRouter != nil { targetProvider = strings.ToLower(backend.ALBOptions.UserRouter.TargetProvider) + } else if provider == providers.ALB && options.Protocol != listener.ProtocolHTTP && nativeAdapter != nil { + // a selection strategy balances the listener's sessions over its pool + if err := sessionBalancer(c, name, options.Protocol, backendName, backend, nativeAdapter); err != nil { + return err + } + streamALBs.Set(backendName) + // its pool members were each checked against the adapter's providers + continue } if options.Protocol != listener.ProtocolHTTP && nativeAdapter != nil && !nativeAdapter.ServesProvider(targetProvider) { @@ -430,7 +442,41 @@ func Listeners(c *config.Config) error { } } } - return nil + return requestALBs(c, streamALBs) +} + +// sessionBalancer validates an ALB that a native protocol listener maps to without a user +// router: its strategy must serve sessions, over a pool of the listener's own kind of backend +func sessionBalancer(c *config.Config, listenerName, protocol, backendName string, backend *bo.Options, + adapter native.Adapter, +) error { + if backend == nil || backend.ALBOptions == nil || + !albregistry.Supports(backend.ALBOptions.MechanismName, albtypes.PlaneNative) { + return fmt.Errorf("listener %q with protocol %q requires alb backend %q to use a mechanism that "+ + "routes or balances sessions: %s", listenerName, protocol, backendName, + strings.Join(albregistry.Supporting(albtypes.PlaneNative), ", ")) + } + o := backend.ALBOptions + for _, m := range o.Pool { + if member := c.Backends[m.Name]; member == nil || + !adapter.ServesProvider(strings.ToLower(member.Provider)) { + return fmt.Errorf("listener %q with protocol %q requires every pool member of alb backend %q "+ + "to be a %s backend: %q is not", listenerName, protocol, backendName, + strings.Join(adapter.Providers(), " or "), m.Name) + } + } + if o.Discovery != nil { + return fmt.Errorf("alb backend %q: 'discovery' is not supported on a %s listener", backendName, protocol) + } + if o.Stream != nil { + return fmt.Errorf("alb backend %q: 'stream' options apply only to an alb that serves a "+ + "tcp, tls or udp listener", backendName) + } + if !o.HRW.KeySource.OnNative() { + return fmt.Errorf("listener %q with protocol %q cannot read alb backend %q's hrw.key %q: "+ + "use client_ip or user", listenerName, protocol, backendName, o.HRW.Key) + } + return adapter.ValidateBalancer(c, backendName, backend) } // streamProviders are the providers a stream listener may relay to: one with an origin to dial, @@ -439,8 +485,66 @@ var streamProviders = sets.New([]string{ providers.ReverseProxyShort, providers.ReverseProxy, providers.Proxy, providers.ALB, }) +// requestALBs holds every load balancer that serves an http listener to what a request can +// offer: no server name or session to key on and no connect to time. One that also serves a +// stream or native listener was held to that listener's rules as well, so its settings must +// suit every listener it serves. streamALBs names those that serve a stream or native listener. +func requestALBs(c *config.Config, streamALBs sets.Set[string]) error { + members := c.Backends.PoolMembers() + for _, backendName := range slices.Sorted(maps.Keys(c.Backends)) { + backend := c.Backends[backendName] + if backend == nil || backend.Provider != providers.ALB || backend.ALBOptions == nil { + continue + } + o := backend.ALBOptions + // a mechanism that commits a flow to several members at once is the relay's to carry out + streamOnly := albregistry.Supports(o.MechanismName, albtypes.PlaneStream) && + !albregistry.Supports(o.MechanismName, albtypes.PlaneHTTP) + if streamOnly && members.Contains(backendName) { + return fmt.Errorf("alb backend %q: mechanism %q cannot be a member of another alb's pool", + backendName, o.MechanismName) + } + if o.Stream != nil && !streamALBs.Contains(backendName) { + return fmt.Errorf("alb backend %q: 'stream' options apply only to an alb that serves a "+ + "tcp, tls or udp listener", backendName) + } + httpListener := servesHTTPListener(c, backend) + if httpListener == "" { + continue + } + if streamOnly { + return fmt.Errorf("alb backend %q: mechanism %q requires a %s listener, and cannot serve "+ + "http listener %q", backendName, o.MechanismName, + strings.Join(albregistry.StreamProtocols(o.MechanismName), " or "), httpListener) + } + if !o.HRW.KeySource.OnHTTP() { + return fmt.Errorf("alb backend %q: hrw.key %q cannot be read from a request, which http "+ + "listener %q serves", backendName, o.HRW.Key, httpListener) + } + if _, err := o.LTSignalFor(listener.ProtocolHTTP); err != nil { + return fmt.Errorf("alb backend %q: %w", backendName, err) + } + } + return nil +} + +// servesHTTPListener returns the name of an http listener the backend is mapped to, or "" when +// it has none. A backend that names no listener serves the default one, which is http. +func servesHTTPListener(c *config.Config, backend *bo.Options) string { + if len(backend.ListenerNames) == 0 { + return listener.DefaultFrontendName + } + for _, name := range backend.ListenerNames { + lo := c.Listeners[name] + if lo == nil || lo.Protocol == "" || lo.Protocol == listener.ProtocolHTTP { + return name + } + } + return "" +} + func streamListener(c *config.Config, name string, options *listener.Options, - mapped map[string]string, + mapped map[string]string, streamALBs sets.Set[string], ) error { if err := options.Stream.Validate(); err != nil { return fmt.Errorf("listener %q: %w", name, err) @@ -462,10 +566,30 @@ func streamListener(c *config.Config, name string, options *listener.Options, name, options.Protocol, backendName, provider) } if provider == providers.ALB && (backend.ALBOptions == nil || - (backend.ALBOptions.MechanismName != albnames.MechanismRR && - backend.ALBOptions.MechanismName != rr.Name)) { - return fmt.Errorf("listener %q with protocol %q requires alb backend %q to use the %s mechanism", - name, options.Protocol, backendName, albnames.MechanismRR) + !albregistry.Supports(backend.ALBOptions.MechanismName, albtypes.PlaneStream)) { + return fmt.Errorf("listener %q with protocol %q requires alb backend %q to use a mechanism that "+ + "serves a %s listener: %s", name, options.Protocol, backendName, options.Protocol, + strings.Join(albregistry.ServingProtocol(options.Protocol), ", ")) + } + if provider == providers.ALB { + streamALBs.Set(backendName) + o := backend.ALBOptions + if !albregistry.ServesProtocol(o.MechanismName, options.Protocol) { + return fmt.Errorf("listener %q with protocol %q cannot serve alb backend %q: mechanism %q requires a %s listener", + name, options.Protocol, backendName, o.MechanismName, + strings.Join(albregistry.StreamProtocols(o.MechanismName), " or ")) + } + if !o.HRW.KeySource.OnStream(ao.StreamListener{ + TLS: options.Protocol == listener.ProtocolTLS, + ProxyProtocol: options.ProxyProtocol && options.Protocol != listener.ProtocolUDP, + }) { + return fmt.Errorf("listener %q with protocol %q cannot read alb backend %q's hrw.key %q: use "+ + "client_ip, sni on a tls listener, or proxy_tlv: on a tcp or tls listener with proxy_protocol", + name, options.Protocol, backendName, o.HRW.Key) + } + if _, err := o.LTSignalFor(options.Protocol); err != nil { + return fmt.Errorf("listener %q: alb backend %q: %w", name, backendName, err) + } } if members.Contains(backendName) { continue @@ -491,15 +615,26 @@ func streamListener(c *config.Config, name string, options *listener.Options, return nil } +// nativeUserRouterTargets returns the backends that a native protocol listener reaches through +// an ALB, which therefore need no listener of their own func nativeUserRouterTargets(c *config.Config, nativeListeners native.Registry) map[string]bool { targets := make(map[string]bool) if c == nil { return targets } for _, backend := range c.Backends { - if backend == nil || backend.Provider != providers.ALB || backend.ALBOptions == nil || - backend.ALBOptions.UserRouter == nil || - nativeListeners.GetByProvider(strings.ToLower(backend.ALBOptions.UserRouter.TargetProvider)) == nil { + if backend == nil || backend.Provider != providers.ALB || backend.ALBOptions == nil { + continue + } + if backend.ALBOptions.UserRouter == nil { + if servesNativeListener(c, backend, nativeListeners) { + for _, m := range backend.ALBOptions.Pool { + targets[m.Name] = true + } + } + continue + } + if nativeListeners.GetByProvider(strings.ToLower(backend.ALBOptions.UserRouter.TargetProvider)) == nil { continue } if name := backend.ALBOptions.UserRouter.DefaultBackend; name != "" { @@ -514,6 +649,18 @@ func nativeUserRouterTargets(c *config.Config, nativeListeners native.Registry) return targets } +// servesNativeListener reports whether any listener the backend names speaks a native protocol +func servesNativeListener(c *config.Config, backend *bo.Options, nativeListeners native.Registry) bool { + backend.NormalizeListenerNames() + for _, name := range backend.ListenerNames { + if lo := c.Listeners[name]; lo != nil && !strings.EqualFold(lo.Protocol, listener.ProtocolHTTP) && + nativeListeners.Get(strings.ToLower(lo.Protocol)) != nil { + return true + } + } + return false +} + func addWarning(c *config.Config, warning string) { if slices.Contains(c.LoaderWarnings, warning) { return diff --git a/pkg/daemon/setup/listeners.go b/pkg/daemon/setup/listeners.go index 92b617277..552cac32e 100644 --- a/pkg/daemon/setup/listeners.go +++ b/pkg/daemon/setup/listeners.go @@ -24,6 +24,7 @@ import ( "time" "github.com/trickstercache/trickster/v2/pkg/backends" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/stream" providerregistry "github.com/trickstercache/trickster/v2/pkg/backends/providers/registry" "github.com/trickstercache/trickster/v2/pkg/config" listenerconfig "github.com/trickstercache/trickster/v2/pkg/config/listener" @@ -40,6 +41,7 @@ import ( ch "github.com/trickstercache/trickster/v2/pkg/proxy/handlers/trickster/config" ph "github.com/trickstercache/trickster/v2/pkg/proxy/handlers/trickster/purge" "github.com/trickstercache/trickster/v2/pkg/proxy/l4" + l4observe "github.com/trickstercache/trickster/v2/pkg/proxy/l4/observe" "github.com/trickstercache/trickster/v2/pkg/proxy/listener" listenerhttp3 "github.com/trickstercache/trickster/v2/pkg/proxy/listener/http3" "github.com/trickstercache/trickster/v2/pkg/proxy/listener/native" @@ -354,7 +356,7 @@ func streamConfig(conf *config.Config, desired desiredListener, clients backends !o.UsesListener(desired.listenerName) { continue } - up := l4.FromBackend(clients.Get(backendName)) + up := stream.FromBackend(clients.Get(backendName)) if up == nil { logger.Error("stream listener backend has no dialable origin", logging.Pairs{ keys.ListenerName: desired.listenerName, keys.BackendName: backendName, @@ -380,6 +382,7 @@ func streamConfig(conf *config.Config, desired desiredListener, clients backends return &l4.Config{ Table: table, Options: desired.options.Stream, MaxConnections: desired.options.ConnectionsLimit, + Observer: l4observe.Listener(desired.listenerName, desired.options.Protocol), } } diff --git a/pkg/kube/gateway/annotations/annotations.go b/pkg/kube/gateway/annotations/annotations.go index fa50af5fb..6fb0c37fa 100644 --- a/pkg/kube/gateway/annotations/annotations.go +++ b/pkg/kube/gateway/annotations/annotations.go @@ -73,6 +73,11 @@ const ( // HealthMode selects how the discovered members of a generated ALB are // judged healthy in the endpoint routing mode: probe or provider HealthMode = Prefix + "health-mode" + // LoadBalancing selects the mechanism that spreads traffic across a Service's endpoints + // in the endpoint routing mode: rr, p2c, lc, lt or hrw + LoadBalancing = Prefix + "load-balancing" + // LoadBalancingKey is what the hrw mechanism keeps together, such as client_ip + LoadBalancingKey = Prefix + "load-balancing-key" ) // Problem is one rejected annotation, for logging and for the Events and @@ -111,7 +116,7 @@ func (s *Set) ConfiguresPolicy() bool { return p.Handler != "" || p.CacheName != "" || p.NegativeCacheName != "" || p.TimeoutMS > 0 || p.MaxTTLMS > 0 || p.CORSMode != "" || p.CollapsedForwarding != "" || p.RewriteTarget != "" || - p.HealthMode != "" || + p.HealthMode != "" || p.LoadBalancing != "" || p.LoadBalancingKey != "" || len(p.RequestHeaders) > 0 || len(p.ResponseHeaders) > 0 || len(p.CORSHeaders) > 0 } @@ -204,6 +209,10 @@ func (s *Set) apply(key, value string) (err error) { return nil case HealthMode: s.Policy.HealthMode, err = translate.HealthMode(value) + case LoadBalancing: + s.Policy.LoadBalancing, err = translate.LoadBalancing(value) + case LoadBalancingKey: + s.Policy.LoadBalancingKey, err = translate.LoadBalancingKey(value) default: return errors.New(reasonUnknown) } diff --git a/pkg/kube/gateway/annotations/annotations_test.go b/pkg/kube/gateway/annotations/annotations_test.go index f1c2a6d12..e0ff08cc5 100644 --- a/pkg/kube/gateway/annotations/annotations_test.go +++ b/pkg/kube/gateway/annotations/annotations_test.go @@ -40,6 +40,8 @@ func TestParseFullSet(t *testing.T) { UseRegex: "true", RewriteTarget: "/v2/${1}", HealthMode: "probe", + LoadBalancing: "hrw", + LoadBalancingKey: "header:X-Tenant", }) require.Empty(t, problems) require.True(t, set.UseRegex) @@ -62,6 +64,8 @@ func TestParseFullSet(t *testing.T) { require.Equal(t, map[string]string{"+Vary": "Accept-Encoding"}, p.ResponseHeaders) require.Equal(t, "/v2/${1}", p.RewriteTarget) require.Equal(t, "probe", p.HealthMode) + require.Equal(t, "hrw", p.LoadBalancing) + require.Equal(t, "header:X-Tenant", p.LoadBalancingKey) } // An annotation outside this controller's namespace is another @@ -99,6 +103,8 @@ func TestParseRejections(t *testing.T) { {"use regex", UseRegex, "yes please", "must be a boolean"}, {"rewrite whitespace", RewriteTarget, "/a b", "must not contain whitespace"}, {"health mode", HealthMode, "guess", "must be"}, + {"load balancing", LoadBalancing, "fr", "must be one of rr, p2c, lc, lt, hrw"}, + {"load balancing key", LoadBalancingKey, "port", "invalid key source"}, {"header shape", RequestHeaders, "X-A 1", "must be 'Name: value'"}, {"header name", RequestHeaders, "X A: 1", "not a valid header name"}, {"header operator only", ResponseHeaders, "-: 1", "not a valid header name"}, diff --git a/pkg/kube/gateway/cachepolicy/cachepolicy_test.go b/pkg/kube/gateway/cachepolicy/cachepolicy_test.go index c728bdc20..3b5444194 100644 --- a/pkg/kube/gateway/cachepolicy/cachepolicy_test.go +++ b/pkg/kube/gateway/cachepolicy/cachepolicy_test.go @@ -171,6 +171,8 @@ func TestIndexLowersEveryField(t *testing.T) { ResponseHeaders: map[string]string{"+Vary": "Accept-Encoding"}, CORS: &CORS{Mode: "merge", Headers: map[string]string{"Access-Control-Allow-Origin": "*"}}, HealthMode: "probe", + LoadBalancing: "p2c", + LoadBalancingKey: "client_ip", ResultHeader: "Hide", } x := New([]*CachePolicy{p}, Config{Known: known()}) @@ -196,6 +198,7 @@ func TestIndexLowersEveryField(t *testing.T) { ResponseHeaders: map[string]string{"+Vary": "Accept-Encoding"}, CORSMode: "merge", CORSHeaders: map[string]string{"Access-Control-Allow-Origin": "*"}, HealthMode: "probe", ResultHeader: ir.ResultHeaderHide, + LoadBalancing: "p2c", LoadBalancingKey: "client_ip", }, *got) require.Equal(t, ir.KindCachePolicy, got.Source.Kind) require.Equal(t, "uid-full", got.Source.UID) @@ -231,6 +234,8 @@ func TestIndexRefusesAnInvalidSpecWhole(t *testing.T) { "cors.mode": func(s *Spec) { s.CORS = &CORS{Mode: "sometimes"} }, "cors.headers": func(s *Spec) { s.CORS = &CORS{Headers: map[string]string{"bad name": "1"}} }, "healthMode": func(s *Spec) { s.HealthMode = "guess" }, + "loadBalancing": func(s *Spec) { s.LoadBalancing = "tsm" }, + "loadBalancingKey": func(s *Spec) { s.LoadBalancingKey = "header:" }, "resultHeader": func(s *Spec) { s.ResultHeader = "Maybe" }, } for field, mutate := range cases { diff --git a/pkg/kube/gateway/cachepolicy/index.go b/pkg/kube/gateway/cachepolicy/index.go index 0a682b322..c40bc0e32 100644 --- a/pkg/kube/gateway/cachepolicy/index.go +++ b/pkg/kube/gateway/cachepolicy/index.go @@ -203,6 +203,8 @@ func (x *Index) lower(p *CachePolicy) (ir.Policy, error) { {"requestHeaders", headerMap(&out.RequestHeaders, s.RequestHeaders)}, {"responseHeaders", headerMap(&out.ResponseHeaders, s.ResponseHeaders)}, {"healthMode", parse(&out.HealthMode, s.HealthMode, translate.HealthMode)}, + {"loadBalancing", parse(&out.LoadBalancing, s.LoadBalancing, translate.LoadBalancing)}, + {"loadBalancingKey", parse(&out.LoadBalancingKey, s.LoadBalancingKey, translate.LoadBalancingKey)}, {"resultHeader", parse(&out.ResultHeader, s.ResultHeader, translate.ResultHeader)}, } if s.CORS != nil { diff --git a/pkg/kube/gateway/cachepolicy/types.go b/pkg/kube/gateway/cachepolicy/types.go index 5ecf15683..ed18a16a9 100644 --- a/pkg/kube/gateway/cachepolicy/types.go +++ b/pkg/kube/gateway/cachepolicy/types.go @@ -114,6 +114,12 @@ type Spec struct { CORS *CORS `json:"cors,omitempty"` // HealthMode is probe or provider, for the endpoint routing mode HealthMode string `json:"healthMode,omitempty"` + // LoadBalancing is rr, p2c, lc, lt or hrw: how traffic is spread across a Service's + // endpoints in the endpoint routing mode + LoadBalancing string `json:"loadBalancing,omitempty"` + // LoadBalancingKey is what hrw keeps together: client_ip, host, sni, header:, + // cookie: or query: + LoadBalancingKey string `json:"loadBalancingKey,omitempty"` // ResultHeader is Expose or Hide: whether X-Trickster-Result reaches the client ResultHeader string `json:"resultHeader,omitempty"` } diff --git a/pkg/kube/gateway/compile/document.go b/pkg/kube/gateway/compile/document.go index 4a4b6e8f1..dff9dba4e 100644 --- a/pkg/kube/gateway/compile/document.go +++ b/pkg/kube/gateway/compile/document.go @@ -290,6 +290,12 @@ type albDoc struct { Mechanism string `yaml:"mechanism,omitempty"` Pool []*albPoolDoc `yaml:"pool,omitempty"` Discovery *albDiscoveryDoc `yaml:"discovery,omitempty"` + HRW *albHRWDoc `yaml:"hrw,omitempty"` +} + +// albHRWDoc is what a generated ALB keeps together when its mechanism is hrw +type albHRWDoc struct { + Key string `yaml:"key,omitempty"` } // albDiscoveryDoc binds a generated ALB's pool to the generated discoverer: diff --git a/pkg/kube/gateway/compile/policy.go b/pkg/kube/gateway/compile/policy.go index e50adb599..f517148f2 100644 --- a/pkg/kube/gateway/compile/policy.go +++ b/pkg/kube/gateway/compile/policy.go @@ -19,6 +19,7 @@ package compile import ( "time" + albnames "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" ho "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck/options" "github.com/trickstercache/trickster/v2/pkg/backends/providers" @@ -61,6 +62,10 @@ type effective struct { // judged healthy, and the probe it runs when that is by probing healthMode string healthCheck *ho.Options + // loadBalancing and loadBalancingKey choose how an endpoint mode ALB spreads traffic + // across a Service's endpoints; the pool across a rule's backendRefs is always round robin + loadBalancing string + loadBalancingKey string // tsProvider is the time series provider the generated backend is, or empty for a plain // reverse proxy; a provider accelerates its own API paths and always caches tsProvider string @@ -121,6 +126,23 @@ func resolveMember(opts *kubecfg.Options, rule, member *ir.Policy) effective { return resolve(opts, &merged) } +// endpointALB returns the ALB that balances one Service's endpoints, with the policy's +// mechanism. A key that the listener cannot read, such as a request header on a tcp route, +// is left at the default rather than compiled into a configuration that would not load. +func (e effective) endpointALB(discovery *albDiscoveryDoc, readable func(ao.KeySource) bool) *albDoc { + alb := &albDoc{Mechanism: albnames.MechanismRR, Discovery: discovery} + if e.loadBalancing != "" { + alb.Mechanism = e.loadBalancing + } + if alb.Mechanism != albnames.MechanismHRW || e.loadBalancingKey == "" { + return alb + } + if ks, err := ao.ParseKeySource(e.loadBalancingKey); err == nil && readable(ks) { + alb.HRW = &albHRWDoc{Key: e.loadBalancingKey} + } + return alb +} + func resolve(opts *kubecfg.Options, p *ir.Policy) effective { e := effective{routingMode: opts.RoutingMode(), healthMode: ao.HealthModeProvider} if d := opts.Defaults; d != nil { @@ -146,6 +168,12 @@ func resolve(opts *kubecfg.Options, p *ir.Policy) effective { if p.HealthMode != "" { e.healthMode = p.HealthMode } + if p.LoadBalancing != "" { + e.loadBalancing = p.LoadBalancing + } + if p.LoadBalancingKey != "" { + e.loadBalancingKey = p.LoadBalancingKey + } if p.CacheName != "" { e.cacheName = p.CacheName } diff --git a/pkg/kube/gateway/compile/stream.go b/pkg/kube/gateway/compile/stream.go index 06c3469f4..96faa8079 100644 --- a/pkg/kube/gateway/compile/stream.go +++ b/pkg/kube/gateway/compile/stream.go @@ -82,7 +82,7 @@ func compileStreamRule(doc *document, r ir.Route, group ir.BackendGroup, listene b.AnyHostRouting = len(r.Hostnames) == 0 } if len(group.Members) == 1 && !group.Members[0].Invalid { - front, err := streamMember(doc, group, group.Members[0], eff, opts, listeners) + front, err := streamMember(doc, r, group, group.Members[0], eff, opts, listeners) if err != nil { return err } @@ -98,7 +98,7 @@ func compileStreamRule(doc *document, r ir.Route, group ir.BackendGroup, listene front = unresolvedStreamBackend() } else { var err error - if front, err = streamMember(doc, group, m, eff, opts, listeners); err != nil { + if front, err = streamMember(doc, r, group, m, eff, opts, listeners); err != nil { return err } } @@ -116,7 +116,7 @@ func compileStreamRule(doc *document, r ir.Route, group ir.BackendGroup, listene return nil } -func streamMember(doc *document, g ir.BackendGroup, m ir.BackendMember, eff effective, +func streamMember(doc *document, r ir.Route, g ir.BackendGroup, m ir.BackendMember, eff effective, opts *kubecfg.Options, listeners []string, ) (*backendDoc, error) { // the member is a reverse proxy backend for its origin alone: the stream listener dials the @@ -147,18 +147,19 @@ func streamMember(doc *document, g ir.BackendGroup, m ir.BackendMember, eff effe return &backendDoc{ Provider: providers.ALB, ListenerNames: listeners, - ALB: &albDoc{ - Mechanism: albnames.MechanismRR, - Discovery: &albDiscoveryDoc{ - DiscovererName: doc.discoverer(opts), - TemplateBackend: tmplName, - HealthMode: ao.HealthModeProvider, - Query: &queryDoc{ - Kind: do.KindEndpointSlices, Namespace: m.Service.Namespace, - Service: m.Service.Name, Port: m.Service.PortName, Scheme: scheme, - }, + // the endpoints are balanced by the policy's mechanism; a key must be one the route's + // listener can read, which is the client address, or the server name on a tls route + ALB: eff.endpointALB(&albDiscoveryDoc{ + DiscovererName: doc.discoverer(opts), + TemplateBackend: tmplName, + HealthMode: ao.HealthModeProvider, + Query: &queryDoc{ + Kind: do.KindEndpointSlices, Namespace: m.Service.Namespace, + Service: m.Service.Name, Port: m.Service.PortName, Scheme: scheme, }, - }, + }, func(ks ao.KeySource) bool { + return ks.OnStream(ao.StreamListener{TLS: r.Protocol == ir.ProtocolTLS}) + }), }, nil } diff --git a/pkg/kube/gateway/compile/stream_test.go b/pkg/kube/gateway/compile/stream_test.go index 0e8cfe78a..90376ddbd 100644 --- a/pkg/kube/gateway/compile/stream_test.go +++ b/pkg/kube/gateway/compile/stream_test.go @@ -167,3 +167,92 @@ func TestCompileStreamSkipsWhatItCannotServe(t *testing.T) { require.NoError(t, err) require.Empty(t, doc.Backends) } + +// a policy's mechanism balances each Service's endpoints; the weights between a rule's +// backendRefs stay with round robin, since Gateway API makes them an exact apportionment +func TestCompileLoadBalancingPolicy(t *testing.T) { + withPolicy := func(m *ir.IR, p ir.Policy) *ir.IR { + p.Name = "lb" + m.Policies = []ir.Policy{p} + for i := range m.Routes { + for j := range m.Routes[i].Rules { + m.Routes[i].Rules[j].Policy = "lb" + } + } + return m + } + t.Run("tcp endpoints, weighted rule", func(t *testing.T) { + m := withPolicy(streamShape(ir.ProtocolTCP, tcpMember(0, "a-svc", 3), tcpMember(1, "b-svc", 1)), + ir.Policy{LoadBalancing: "p2c"}) + doc, err := buildDocument(m, endpointOpts(t), nil) + require.NoError(t, err) + outer := doc.Backends["kgw--tcproute.data.db_r0"] + require.Equal(t, "rr", outer.ALB.Mechanism, "backendRef weights are apportioned exactly") + require.Len(t, outer.ALB.Pool, 2) + for _, name := range []string{"kgw--tcproute.data.db_r0_b0", "kgw--tcproute.data.db_r0_b1"} { + inner := doc.Backends[name] + require.Equal(t, "p2c", inner.ALB.Mechanism, name) + require.NotNil(t, inner.ALB.Discovery, name) + require.Nil(t, inner.ALB.HRW, name) + } + }) + t.Run("stream keys are what the listener can read", func(t *testing.T) { + for _, test := range []struct { + protocol, key, want string + }{ + {ir.ProtocolTCP, "client_ip", "client_ip"}, + {ir.ProtocolTLS, "sni", "sni"}, + // there is no server name on a tcp or udp route, and no header on any of them + {ir.ProtocolTCP, "sni", ""}, + {ir.ProtocolUDP, "sni", ""}, + {ir.ProtocolTLS, "header:X-Tenant", ""}, + {ir.ProtocolTCP, "", ""}, + } { + m := withPolicy(streamShape(test.protocol, tcpMember(0, "a-svc", 1)), + ir.Policy{LoadBalancing: "hrw", LoadBalancingKey: test.key}) + doc, err := buildDocument(m, endpointOpts(t), nil) + require.NoError(t, err) + front := doc.Backends["kgw--tcproute.data.db_r0"] + require.Equal(t, "hrw", front.ALB.Mechanism) + if test.want == "" { + require.Nil(t, front.ALB.HRW, "%s key %q", test.protocol, test.key) + continue + } + require.Equal(t, &albHRWDoc{Key: test.want}, front.ALB.HRW, "%s key %q", test.protocol, test.key) + } + }) + t.Run("a key without hrw is not compiled", func(t *testing.T) { + m := withPolicy(streamShape(ir.ProtocolTCP, tcpMember(0, "a-svc", 1)), + ir.Policy{LoadBalancing: "lc", LoadBalancingKey: "client_ip"}) + doc, err := buildDocument(m, endpointOpts(t), nil) + require.NoError(t, err) + front := doc.Backends["kgw--tcproute.data.db_r0"] + require.Equal(t, "lc", front.ALB.Mechanism) + require.Nil(t, front.ALB.HRW) + }) + t.Run("http endpoints", func(t *testing.T) { + g := group("shop", "web", 0, svcMember(0, "shop", "a", 80, 3), svcMember(1, "shop", "b", 80, 1)) + m := withPolicy(&ir.IR{ + Listeners: []ir.Listener{httpListener()}, + Routes: []ir.Route{route("shop", "web", ir.Rule{BackendGroup: g.Name})}, + Backends: []ir.BackendGroup{g}, + }, ir.Policy{LoadBalancing: "hrw", LoadBalancingKey: "header:X-Tenant"}) + doc, err := buildDocument(m, endpointOpts(t), nil) + require.NoError(t, err) + require.Equal(t, "rr", doc.Backends["kgw--httproute.shop.web_r0"].ALB.Mechanism) + inner := doc.Backends["kgw--httproute.shop.web_r0_b0"] + require.Equal(t, "hrw", inner.ALB.Mechanism) + require.Equal(t, &albHRWDoc{Key: "header:X-Tenant"}, inner.ALB.HRW) + // a request has no server name to key on + m.Policies[0].LoadBalancingKey = "sni" + doc, err = buildDocument(m, endpointOpts(t), nil) + require.NoError(t, err) + require.Nil(t, doc.Backends["kgw--httproute.shop.web_r0_b0"].ALB.HRW) + }) + t.Run("service mode has no endpoints to balance", func(t *testing.T) { + m := withPolicy(streamShape(ir.ProtocolTCP, tcpMember(0, "a-svc", 1)), ir.Policy{LoadBalancing: "p2c"}) + doc, err := buildDocument(m, serviceOpts(t), nil) + require.NoError(t, err) + require.Nil(t, doc.Backends["kgw--tcproute.data.db_r0"].ALB) + }) +} diff --git a/pkg/kube/gateway/compile/targets.go b/pkg/kube/gateway/compile/targets.go index 2f15773a0..668e2c1c6 100644 --- a/pkg/kube/gateway/compile/targets.go +++ b/pkg/kube/gateway/compile/targets.go @@ -22,7 +22,6 @@ import ( "strings" "time" - albnames "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" "github.com/trickstercache/trickster/v2/pkg/backends/providers" kubecfg "github.com/trickstercache/trickster/v2/pkg/config/kubernetes" @@ -73,15 +72,12 @@ func newMemberTarget(doc *document, g ir.BackendGroup, m ir.BackendMember, } front := &backendDoc{ Provider: providers.ALB, - ALB: &albDoc{ - Mechanism: albnames.MechanismRR, - Discovery: &albDiscoveryDoc{ - DiscovererName: doc.discoverer(opts), - TemplateBackend: tmplName, - HealthMode: eff.healthMode, - Query: query, - }, - }, + ALB: eff.endpointALB(&albDiscoveryDoc{ + DiscovererName: doc.discoverer(opts), + TemplateBackend: tmplName, + HealthMode: eff.healthMode, + Query: query, + }, ao.KeySource.OnHTTP), } return &memberTarget{ front: front, frontHandler: providers.ALB, diff --git a/pkg/kube/gateway/gateway/classes.go b/pkg/kube/gateway/gateway/classes.go index 86efb48f5..9b1bcab7b 100644 --- a/pkg/kube/gateway/gateway/classes.go +++ b/pkg/kube/gateway/gateway/classes.go @@ -42,6 +42,8 @@ const ( ParamAuthenticatorName = "authenticator_name" ParamTimeout = "timeout" ParamHealthMode = "health_mode" + ParamLoadBalancing = "load_balancing" + ParamLoadBalancingKey = "load_balancing_key" ) var ( @@ -134,6 +136,14 @@ var classParams = map[string]paramSetter{ p.HealthMode, err = translate.HealthMode(v) return err }, + ParamLoadBalancing: func(_ *translator, p *ir.Policy, v string) (err error) { + p.LoadBalancing, err = translate.LoadBalancing(v) + return err + }, + ParamLoadBalancingKey: func(_ *translator, p *ir.Policy, v string) (err error) { + p.LoadBalancingKey, err = translate.LoadBalancingKey(v) + return err + }, ParamTimeout: func(_ *translator, p *ir.Policy, v string) error { d, err := timeconv.ParsePositiveDuration(v) if err != nil { diff --git a/pkg/kube/gateway/gateway/gateway_test.go b/pkg/kube/gateway/gateway/gateway_test.go index 28b29452a..46f4d154c 100644 --- a/pkg/kube/gateway/gateway/gateway_test.go +++ b/pkg/kube/gateway/gateway/gateway_test.go @@ -2037,6 +2037,8 @@ func TestApplyParameterRejections(t *testing.T) { "unknown key": {"colour", "blue"}, "bad routing mode": {ParamRoutingMode, "sideways"}, "bad health mode": {ParamHealthMode, "guess"}, + "bad load balancing": {ParamLoadBalancing, "fr"}, + "bad balancing key": {ParamLoadBalancingKey, "port"}, "bad timeout": {ParamTimeout, "45"}, "negative timeout": {ParamTimeout, "-1s"}, "unknown cache": {ParamCacheName, "nope"}, @@ -2055,6 +2057,10 @@ func TestApplyParameterRejections(t *testing.T) { var p ir.Policy require.NoError(t, tr.applyParameter(&p, ParamNegativeCacheName, "api-errors")) require.Equal(t, "api-errors", p.NegativeCacheName) + require.NoError(t, tr.applyParameter(&p, ParamLoadBalancing, "hrw")) + require.NoError(t, tr.applyParameter(&p, ParamLoadBalancingKey, "sni")) + require.Equal(t, "hrw", p.LoadBalancing) + require.Equal(t, "sni", p.LoadBalancingKey) // with nothing to check against, any name is accepted free := &translator{} require.NoError(t, free.applyParameter(&p, ParamReqRewriterName, "anything")) diff --git a/pkg/kube/gateway/gateway/stream_test.go b/pkg/kube/gateway/gateway/stream_test.go index 9188f61cc..dd48f256d 100644 --- a/pkg/kube/gateway/gateway/stream_test.go +++ b/pkg/kube/gateway/gateway/stream_test.go @@ -29,6 +29,7 @@ import ( "github.com/trickstercache/trickster/v2/pkg/backends" "github.com/trickstercache/trickster/v2/pkg/backends/alb" "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" + "github.com/trickstercache/trickster/v2/pkg/backends/alb/stream" "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" "github.com/trickstercache/trickster/v2/pkg/config" "github.com/trickstercache/trickster/v2/pkg/discovery" @@ -211,7 +212,7 @@ func streamTable(t *testing.T, conf *config.Config, clients backends.Backends, l if o.IsTemplate || members.Contains(name) || !o.UsesListener(listener) { continue } - up := l4.FromBackend(clients.Get(name)) + up := stream.FromBackend(clients.Get(name)) require.NotNil(t, up, name) hosts := o.Hosts if !sni || len(hosts) == 0 { diff --git a/pkg/kube/gateway/internal/translate/translate_test.go b/pkg/kube/gateway/internal/translate/translate_test.go index e60343d53..2be5669ad 100644 --- a/pkg/kube/gateway/internal/translate/translate_test.go +++ b/pkg/kube/gateway/internal/translate/translate_test.go @@ -271,3 +271,25 @@ func TestServicePortSelectsByTransport(t *testing.T) { _, ok := ServicePort(only, PortRef{Number: 53, Protocol: corev1.ProtocolUDP}) require.False(t, ok) } + +func TestLoadBalancing(t *testing.T) { + for _, v := range []string{"rr", "p2c", "lc", "lt", "hrw"} { + got, err := LoadBalancing(v) + require.NoError(t, err) + require.Equal(t, v, got) + } + // only a mechanism that commits to one endpoint can balance a Service's endpoints + for _, v := range []string{"", "fr", "tsm", "ur", "round_robin", "RR"} { + _, err := LoadBalancing(v) + require.ErrorContains(t, err, "must be one of", v) + } + for _, v := range []string{"client_ip", "sni", "host", "header:X-Tenant", "cookie:session", "query:tenant"} { + got, err := LoadBalancingKey(" " + v + " ") + require.NoError(t, err) + require.Equal(t, v, got) + } + for _, v := range []string{"port", "header:", "cookie:a b"} { + _, err := LoadBalancingKey(v) + require.Error(t, err, v) + } +} diff --git a/pkg/kube/gateway/internal/translate/values.go b/pkg/kube/gateway/internal/translate/values.go index eb8b0f75f..ab15e084f 100644 --- a/pkg/kube/gateway/internal/translate/values.go +++ b/pkg/kube/gateway/internal/translate/values.go @@ -23,6 +23,7 @@ import ( "slices" "strings" + albnames "github.com/trickstercache/trickster/v2/pkg/backends/alb/names" ao "github.com/trickstercache/trickster/v2/pkg/backends/alb/options" "github.com/trickstercache/trickster/v2/pkg/backends/providers" kubecfg "github.com/trickstercache/trickster/v2/pkg/config/kubernetes" @@ -51,6 +52,27 @@ func HealthMode(v string) (string, error) { return v, nil } +// LoadBalancing parses the mechanism that balances a Service's endpoints: one that commits +// each request or connection to a single member, by its short name +func LoadBalancing(v string) (string, error) { + if slices.Contains(loadBalancingMechanisms, v) { + return v, nil + } + return "", fmt.Errorf("must be one of %s", strings.Join(loadBalancingMechanisms, ", ")) +} + +var loadBalancingMechanisms = []string{ + albnames.MechanismRR, albnames.MechanismP2C, albnames.MechanismLC, albnames.MechanismLT, albnames.MechanismHRW, +} + +// LoadBalancingKey parses what the hrw mechanism keeps together +func LoadBalancingKey(v string) (string, error) { + if _, err := ao.ParseKeySource(v); err != nil { + return "", err + } + return strings.TrimSpace(v), nil +} + // Handler parses a path handler name func Handler(v string) (string, error) { switch v { diff --git a/pkg/kube/gateway/ir/ir.go b/pkg/kube/gateway/ir/ir.go index b73610fe2..1ad3abf23 100644 --- a/pkg/kube/gateway/ir/ir.go +++ b/pkg/kube/gateway/ir/ir.go @@ -556,6 +556,12 @@ type Policy struct { AuthenticatorName string `json:"authenticator_name,omitempty"` // HealthMode is the health mode of generated discovery-backed ALBs HealthMode string `json:"health_mode,omitempty"` + // LoadBalancing is the mechanism that spreads traffic across a Service's endpoints in the + // endpoint routing mode; empty is round robin. The weights between a rule's backendRefs + // are always apportioned by round robin, whatever this says. + LoadBalancing string `json:"load_balancing,omitempty"` + // LoadBalancingKey is what the hrw mechanism keeps together, such as client_ip + LoadBalancingKey string `json:"load_balancing_key,omitempty"` // Provider makes the generated backend a time series provider (prometheus, influxdb, ...) // whose own API paths it then accelerates; only a cache policy sets it Provider string `json:"provider,omitempty"` @@ -587,6 +593,8 @@ func (p Policy) Overlay(o *Policy) Policy { overlayString(&out.ReqRewriterName, o.ReqRewriterName) overlayString(&out.AuthenticatorName, o.AuthenticatorName) overlayString(&out.HealthMode, o.HealthMode) + overlayString(&out.LoadBalancing, o.LoadBalancing) + overlayString(&out.LoadBalancingKey, o.LoadBalancingKey) overlayString(&out.Provider, o.Provider) overlayString(&out.ResultHeader, o.ResultHeader) if o.TimeoutMS > 0 { diff --git a/pkg/kube/gateway/ir/policy_test.go b/pkg/kube/gateway/ir/policy_test.go index e77a2653a..0a531bd0f 100644 --- a/pkg/kube/gateway/ir/policy_test.go +++ b/pkg/kube/gateway/ir/policy_test.go @@ -124,6 +124,7 @@ func TestPolicyOverlayFillsEveryStringField(t *testing.T) { Handler: "h", CacheName: "c", RoutingMode: "r", NegativeCacheName: "n", CORSMode: "m", CollapsedForwarding: "cf", RewriteTarget: "/t", TracingName: "tr", ReqRewriterName: "rw", AuthenticatorName: "a", HealthMode: "probe", + LoadBalancing: "hrw", LoadBalancingKey: "client_ip", Provider: "graphite", ResultHeader: ResultHeaderHide, } got := Policy{}.Overlay(over) diff --git a/pkg/lb/balancer.go b/pkg/lb/balancer.go new file mode 100644 index 000000000..91d44fd20 --- /dev/null +++ b/pkg/lb/balancer.go @@ -0,0 +1,360 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import ( + "math" + "slices" + "sync/atomic" + "time" +) + +// Pick is one committed selection, returned by value. Its holder reports what became of the +// flow; each report costs nothing unless the strategy declared a need for it. +type Pick struct { + member *Member + balancer *Balancer + // monotonic nanoseconds since the process's epoch; set only for a strategy that needs latency + start int64 +} + +// epoch anchors the monotonic clock readings that picks are timed with +var epoch = time.Now() + +func monotonic() int64 { + return int64(time.Since(epoch)) +} + +// Member returns the selected member, whose Value is what the caller dispatches to. +func (p Pick) Member() *Member { + return p.member +} + +// Established reports that the member was reached and how long that took, as a latency +// sample. A caller reports whichever of Established and FirstByte is its latency signal. +func (p Pick) Established(d time.Duration) { + if p.tracks(NeedLatency) { + p.member.stats.observe(float64(d), time.Now().UnixNano(), p.balancer.decay) + } +} + +// FirstByte reports the first sign of a response from the member; the time since the pick is +// a latency sample. +func (p Pick) FirstByte() { + if p.tracks(NeedLatency) { + p.member.stats.observe(float64(monotonic()-p.start), time.Now().UnixNano(), p.balancer.decay) + } +} + +// Reached reports that the member was reached, which ends its run of failures to reach it. A +// caller that reports OutcomeConnectFailed reports this as soon as a connect succeeds, not when +// the flow ends: only then are the failures that count toward ejection consecutive connects. +func (p Pick) Reached() { + if p.member != nil && p.balancer != nil && p.balancer.ejects && p.member.stats.connectFails.Load() != 0 { + p.member.stats.connectFails.Store(0) + } +} + +// Done reports that the flow ended. It must be called exactly once per Pick, including when +// the work panics, or the member's in-flight count leaks. A failed outcome records a latency +// penalty, so a member that fails fast never looks fast. +func (p Pick) Done(o Outcome) { + if p.member == nil || p.balancer == nil || (p.balancer.needs == 0 && !p.balancer.ejects) { + return + } + st := p.member.stats + if p.balancer.needs.Has(NeedInflight) { + st.inflight.Add(-1) + } + switch o { + case OutcomeOK: + if st.fails.Load() != 0 { + st.fails.Store(0) + } + case OutcomeFailed, OutcomeConnectFailed: + st.fails.Add(1) + if p.balancer.needs.Has(NeedLatency) { + p.penalize(st) + } + // how a member answered never counts toward ejection, only whether it could be reached + if o == OutcomeConnectFailed && p.balancer.ejects && + int(st.connectFails.Add(1)) >= p.balancer.ejection.Failures { + p.balancer.eject(p.member) + } + case OutcomeCanceled: + } +} + +func (p Pick) tracks(n Needs) bool { + return p.member != nil && p.balancer != nil && p.balancer.needs.Has(n) +} + +// penalize records a sample no lower than the penalty, the time the failure took, and twice +// the current average, so repeated failures push a member further back, up to a ceiling +func (p Pick) penalize(st *Stats) { + b := p.balancer + sample := max(b.penalty, float64(monotonic()-p.start), 2*math.Float64frombits(st.latency.Load())) + st.latency.Store(math.Float64bits(min(sample, b.penalty*penaltyCeiling))) + st.stamp.Store(time.Now().UnixNano()) +} + +const ( + // DefaultLatencyDecay is the time constant of the latency average. + DefaultLatencyDecay = 10 * time.Second + // DefaultLatencyPenalty is the least latency a failed outcome is recorded as. + DefaultLatencyPenalty = 5 * time.Second + // penaltyCeiling caps a penalized average, as a multiple of the penalty + penaltyCeiling = 12 +) + +// LatencyOptions tune how latency samples are averaged. Zero values take the defaults. +type LatencyOptions struct { + // Decay is the time constant with which an average yields to lower samples. + Decay time.Duration + // Penalty is the least latency a failed outcome is recorded as. + Penalty time.Duration +} + +// LatencyTuner is optionally implemented by a Selector that needs latency, to tune how the +// balancer averages the samples it records for that strategy. +type LatencyTuner interface { + Latency() LatencyOptions +} + +// EjectionOptions configure passive ejection: taking a member out of selection when flows +// keep failing to reach it, without waiting for a health check to notice. Only +// OutcomeConnectFailed counts, never how a member answered. +type EjectionOptions struct { + // Failures is how many consecutive connect failures eject a member; zero disables ejection. + Failures int + // Duration is how long an ejected member stays out. When it ends the member is selected + // again only if its Health, if it has one, still meets the pool's floor. + Duration time.Duration + // MaxPercent is the most of a pool's members that may be out at once, 1-100; the last + // live member is never ejected whatever the percentage. Zero means 50. + MaxPercent int +} + +// BalancerOptions are the optional settings of a Balancer. +type BalancerOptions struct { + // Pool is the balancer's first pool; SetPool installs one later. + Pool *Pool + // Ejection configures passive ejection; the zero value leaves it off. + Ejection EjectionOptions + // Observer receives the balancer's events; nil discards them. + Observer Observer +} + +// Balancer binds a strategy to a swappable pool. It is the Picker behind every pick-one +// mechanism, on any plane, and is safe for concurrent use. +type Balancer struct { + selector Selector + needs Needs + // nanoseconds, as float64 so the sampling path converts nothing + decay float64 + penalty float64 + ejects bool + ejection EjectionOptions + observer Observer + pool atomic.Pointer[Pool] + prepared atomic.Pointer[preparedSnapshot] +} + +// preparedSnapshot pairs a Prepared with the snapshot it was built from. The pairing is by +// snapshot identity: generations restart with each pool, so they do not identify a snapshot. +type preparedSnapshot struct { + snap *Snapshot + prepared Prepared +} + +// NewBalancer returns a Balancer for selector, which it then owns. +func NewBalancer(selector Selector, opts ...BalancerOptions) *Balancer { + b := &Balancer{ + selector: selector, needs: selector.Needs(), + decay: float64(DefaultLatencyDecay), penalty: float64(DefaultLatencyPenalty), + } + if t, ok := selector.(LatencyTuner); ok { + lo := t.Latency() + if lo.Decay > 0 { + b.decay = float64(lo.Decay) + } + if lo.Penalty > 0 { + b.penalty = float64(lo.Penalty) + } + } + if len(opts) > 0 { + b.configure(opts[0]) + } + return b +} + +const ( + // DefaultEjectionDuration is how long an ejected member stays out when none is set. + DefaultEjectionDuration = 30 * time.Second + // DefaultEjectionMaxPercent is the most of a pool that may be ejected when none is set. + DefaultEjectionMaxPercent = 50 +) + +func (b *Balancer) configure(o BalancerOptions) { + if o.Pool != nil { + b.pool.Store(o.Pool) + } + b.observer = o.Observer + if o.Ejection.Failures <= 0 { + return + } + b.ejects, b.ejection = true, o.Ejection + if b.ejection.Duration <= 0 { + b.ejection.Duration = DefaultEjectionDuration + } + if b.ejection.MaxPercent <= 0 || b.ejection.MaxPercent > 100 { + b.ejection.MaxPercent = DefaultEjectionMaxPercent + } +} + +// eject takes a member out of selection for the ejection duration, if the pool can spare it. +// It runs on a failure path only. The member returns through a refresh of whichever pool is +// current when the time is up, since membership may have been swapped meanwhile. +func (b *Balancer) eject(m *Member) { + p := b.pool.Load() + if p == nil || !p.eject(m, time.Now().Add(b.ejection.Duration), b.ejection.MaxPercent) { + return + } + p.Refresh() + time.AfterFunc(b.ejection.Duration, func() { + if cur := b.pool.Load(); cur != nil { + cur.Refresh() + } + }) + if b.observer != nil { + b.observer.Observe(Event{Kind: EventEjected, Member: m.name}) + } +} + +// SetPool replaces the pool that picks are made from; nil leaves the balancer with none. The +// strategy's state carries over, so a rotation continues across a change of membership. +func (b *Balancer) SetPool(p *Pool) { + b.pool.Store(p) +} + +// Pool returns the balancer's current pool, or nil. +func (b *Balancer) Pool() *Pool { + return b.pool.Load() +} + +// Selector returns the balancer's strategy. +func (b *Balancer) Selector() Selector { + return b.selector +} + +// Needs returns the Needs of the balancer's strategy. +func (b *Balancer) Needs() Needs { + return b.needs +} + +// Pick commits one flow to an eligible member of the current pool. +func (b *Balancer) Pick(f Flow) (Pick, bool) { + p := b.pool.Load() + if p == nil { + return Pick{}, false + } + snap := p.Snapshot() + if len(snap.Members) == 0 { + return Pick{}, false + } + ps := b.prepared.Load() + if ps == nil || ps.snap != snap { + ps = b.prepare(snap) + } + return b.commit(ps.prepared.Select(f)) +} + +// prepare is the rare path taken on the first pick of each snapshot. Racing callers build +// equal values; one that stores a superseded value is corrected by the next pick. +func (b *Balancer) prepare(snap *Snapshot) *preparedSnapshot { + ps := &preparedSnapshot{snap: snap, prepared: b.selector.Prepare(snap)} + b.prepared.Store(ps) + return ps +} + +func (b *Balancer) commit(m *Member) (Pick, bool) { + if m == nil { + return Pick{}, false + } + if b.needs == 0 { + return Pick{member: m, balancer: b}, true + } + if b.needs.Has(NeedInflight) { + m.stats.inflight.Add(1) + } + pk := Pick{member: m, balancer: b} + if b.needs.Has(NeedLatency) { + pk.start = monotonic() + } + return pk, true +} + +// Alternatives returns the eligible members that skip does not report: the strategy's choice +// among them for the flow first, then the others in pool order from there. It is not a +// selection path: it prepares a one-off snapshot, once, however many members it returns. +func (b *Balancer) Alternatives(f Flow, skip func(*Member) bool) []*Member { + p := b.pool.Load() + if p == nil { + return nil + } + snap := p.Snapshot() + if len(snap.Members) == 0 { + return nil + } + others := make([]*Member, 0, len(snap.Members)) + for _, m := range snap.Members { + if skip == nil || !skip(m) { + others = append(others, m) + } + } + if len(others) == 0 { + return nil + } + first := b.selector.Prepare(&Snapshot{Members: others, Tier: snap.Tier, Gen: snap.Gen}).Select(f) + if first == nil { + return nil + } + i := slices.Index(others, first) + if i < 0 { + // the strategy's choice is honored as Pick honors it, even one it was never offered + return append([]*Member{first}, others...) + } + // rotate the choice to the front in place: three reversals, no second slice + slices.Reverse(others[:i]) + slices.Reverse(others[i:]) + slices.Reverse(others) + return others +} + +// Commit commits a flow to a member that Alternatives returned, as Pick would have. +func (b *Balancer) Commit(m *Member) (Pick, bool) { + return b.commit(m) +} + +// Repick commits the flow to an eligible member that is none of failed, for a caller retrying +// work those members could not take. It is not a selection path. +func (b *Balancer) Repick(f Flow, failed ...*Member) (Pick, bool) { + alts := b.Alternatives(f, func(m *Member) bool { return slices.Contains(failed, m) }) + if len(alts) == 0 { + return Pick{}, false + } + return b.commit(alts[0]) +} diff --git a/pkg/lb/balancer_test.go b/pkg/lb/balancer_test.go new file mode 100644 index 000000000..587f99eba --- /dev/null +++ b/pkg/lb/balancer_test.go @@ -0,0 +1,209 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb_test + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" +) + +// countingSelector always takes the first member, and counts how often it was prepared +type countingSelector struct { + needs lb.Needs + prepares int + refuse bool +} + +type firstMember struct { + members []*lb.Member + refuse bool +} + +func (s *countingSelector) Name() string { return "first" } + +func (s *countingSelector) Needs() lb.Needs { return s.needs } + +func (s *countingSelector) Prepare(snap *lb.Snapshot) lb.Prepared { + s.prepares++ + return &firstMember{members: snap.Members, refuse: s.refuse} +} + +func (p *firstMember) Select(lb.Flow) *lb.Member { + if p.refuse { + return nil + } + return p.members[0] +} + +func TestNeeds(t *testing.T) { + n := lb.NeedKey | lb.NeedLatency + if !n.Has(lb.NeedKey) || !n.Has(lb.NeedKey|lb.NeedLatency) || n.Has(lb.NeedInflight) || !n.Has(0) { + t.Errorf("needs %b misreports what it has", n) + } +} + +// a snapshot is prepared once, however many picks it serves, and again only when it changes +func TestBalancerPreparesOncePerSnapshot(t *testing.T) { + members, healths := lbtest.Members(1, 1, 1) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + defer p.Stop() + sel := &countingSelector{needs: lb.NeedInflight} + b := lb.NewBalancer(sel) + if b.Pool() != nil { + t.Error("a new balancer has a pool") + } + b.SetPool(p) + for range 100 { + pk, ok := b.Pick(lb.Flow{}) + if !ok || pk.Member() != members[0] { + t.Fatal("expected the first member") + } + } + if sel.prepares != 1 { + t.Errorf("prepared %d times for one snapshot", sel.prepares) + } + if got := members[0].Stats().Inflight(); got != 100 { + t.Errorf("in-flight = %d, want the 100 open picks", got) + } + healths[0].Set(-1) + pk, ok := b.Pick(lb.Flow{}) + if !ok || pk.Member() != members[1] || sel.prepares != 2 { + t.Errorf("after a transition: member %v, prepared %d times", pk.Member(), sel.prepares) + } + pk.Done(lb.OutcomeOK) + if got := members[1].Stats().Inflight(); got != 0 { + t.Errorf("in-flight after Done = %d", got) + } + + // a strategy that needs no in-flight count is not charged for one + free := lb.NewBalancer(&countingSelector{}, lb.BalancerOptions{Pool: p}) + if pk, ok := free.Pick(lb.Flow{}); !ok || pk.Member().Stats().Inflight() != 0 { + t.Error("a strategy without NeedInflight moved the in-flight count") + } +} + +// a strategy that selects nothing yields no pick rather than a pick of nothing +func TestBalancerSelectorMayDecline(t *testing.T) { + members, _ := lbtest.Members(1, 1) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + defer p.Stop() + b := lb.NewBalancer(&countingSelector{refuse: true, needs: lb.NeedInflight}, lb.BalancerOptions{Pool: p}) + if _, ok := b.Pick(lb.Flow{}); ok { + t.Error("picked although the strategy declined") + } + if _, ok := b.Repick(lb.Flow{}, members[0]); ok { + t.Error("repicked although the strategy declined") + } + if _, ok := lb.NewBalancer(&countingSelector{}).Repick(lb.Flow{}, nil); ok { + t.Error("repicked with no pool") + } +} + +// keyed spreads flows by key and declares every need, so the suite drives all of the balancer +type keyed struct{ members []*lb.Member } + +func (*keyed) Name() string { return "keyed" } + +func (*keyed) Needs() lb.Needs { return lb.NeedKey | lb.NeedInflight | lb.NeedLatency } + +func (*keyed) Prepare(s *lb.Snapshot) lb.Prepared { return &keyed{members: s.Members} } + +func (k *keyed) Select(f lb.Flow) *lb.Member { + return k.members[f.Key%uint64(len(k.members))] +} + +func TestBalancerConformance(t *testing.T) { + lbtest.Run(t, func() lb.Selector { return &keyed{} }, lbtest.Options{}) +} + +// tuned needs latency and asks for its own averaging +type tuned struct { + keyed + opts lb.LatencyOptions +} + +func (t *tuned) Latency() lb.LatencyOptions { return t.opts } + +func TestLatencyAccounting(t *testing.T) { + members, _ := lbtest.Members(1) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + defer p.Stop() + st := members[0].Stats() + b := lb.NewBalancer(&tuned{opts: lb.LatencyOptions{Decay: time.Minute, Penalty: 3 * time.Second}}, + lb.BalancerOptions{Pool: p}) + + pk, _ := b.Pick(lb.Flow{}) + pk.Established(40 * time.Millisecond) + if st.Latency() != 40*time.Millisecond || st.LastSample().IsZero() { + t.Fatalf("connect sample = %v", st.Latency()) + } + pk.Done(lb.OutcomeOK) + + pk, _ = b.Pick(lb.Flow{}) + time.Sleep(60 * time.Millisecond) + pk.FirstByte() + if got := st.Latency(); got < 60*time.Millisecond || got > 2*time.Second { + t.Errorf("first-byte sample = %v", got) + } + // a caller that gives up says nothing about the member + before := st.Latency() + pk.Done(lb.OutcomeCanceled) + if st.Latency() != before || st.Failures() != 0 { + t.Errorf("a canceled flow moved the stats: %v, %d failures", st.Latency(), st.Failures()) + } + + // failures record the penalty, then twice the average, up to the ceiling + for i, want := range []time.Duration{3 * time.Second, 6 * time.Second, 12 * time.Second, 24 * time.Second, 36 * time.Second, 36 * time.Second} { + pk, _ = b.Pick(lb.Flow{}) + if i%2 == 0 { + pk.Done(lb.OutcomeFailed) + } else { + pk.Done(lb.OutcomeConnectFailed) + } + if st.Latency() != want || st.Failures() != int32(i+1) { + t.Fatalf("failure %d: latency %v, want %v; %d failures", i+1, st.Latency(), want, st.Failures()) + } + } + pk, _ = b.Pick(lb.Flow{}) + pk.Done(lb.OutcomeOK) + if st.Failures() != 0 || st.Inflight() != 0 { + t.Errorf("after a success: %d failures, %d in flight", st.Failures(), st.Inflight()) + } + + // defaults apply when a strategy does not tune them + d := lb.NewBalancer(&tuned{}, lb.BalancerOptions{Pool: p}) + fresh, _ := lbtest.Members(1) + fp, _ := lb.NewPool(fresh, 1) + defer fp.Stop() + d.SetPool(fp) + pk, _ = d.Pick(lb.Flow{}) + pk.Done(lb.OutcomeFailed) + if got := fresh[0].Stats().Latency(); got != lb.DefaultLatencyPenalty { + t.Errorf("default penalty = %v", got) + } +} diff --git a/pkg/lb/boundary_test.go b/pkg/lb/boundary_test.go new file mode 100644 index 000000000..ac4ba13dc --- /dev/null +++ b/pkg/lb/boundary_test.go @@ -0,0 +1,64 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb_test + +import ( + "go/parser" + "go/token" + "io/fs" + "path/filepath" + "strconv" + "strings" + "testing" +) + +// the core must stay importable on its own: every file under this directory, tests included, +// may import the standard library and this package tree, and nothing else +func TestImportBoundary(t *testing.T) { + const self = "github.com/trickstercache/trickster/v2/pkg/lb" + fset := token.NewFileSet() + var files int + err := filepath.WalkDir(".", func(path string, d fs.DirEntry, err error) error { + if err != nil || d.IsDir() || !strings.HasSuffix(path, ".go") { + return err + } + f, err := parser.ParseFile(fset, path, nil, parser.ImportsOnly) + if err != nil { + return err + } + files++ + for _, imp := range f.Imports { + name, err := strconv.Unquote(imp.Path.Value) + if err != nil { + return err + } + if name == self || strings.HasPrefix(name, self+"/") { + continue + } + // a standard library import path has no dot in its first element + if first, _, _ := strings.Cut(name, "/"); strings.Contains(first, ".") { + t.Errorf("%s imports %s: only the standard library is allowed under pkg/lb", path, name) + } + } + return nil + }) + if err != nil { + t.Fatal(err) + } + if files == 0 { + t.Fatal("no source files were checked") + } +} diff --git a/pkg/lb/ejection_test.go b/pkg/lb/ejection_test.go new file mode 100644 index 000000000..6b09884fb --- /dev/null +++ b/pkg/lb/ejection_test.go @@ -0,0 +1,320 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb_test + +import ( + "slices" + "sync" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" +) + +type events struct { + mu sync.Mutex + got []lb.Event +} + +func (e *events) Observe(ev lb.Event) { + e.mu.Lock() + defer e.mu.Unlock() + e.got = append(e.got, ev) +} + +func (e *events) ejected() []string { + e.mu.Lock() + defer e.mu.Unlock() + var out []string + for _, ev := range e.got { + if ev.Kind == lb.EventEjected { + out = append(out, ev.Member) + } + } + return out +} + +func ejecting(t *testing.T, o lb.EjectionOptions, weights ...int) (*lb.Balancer, []*lb.Member, []*lbtest.Health, *events) { + t.Helper() + members, healths := lbtest.Members(weights...) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + ev := &events{} + return lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: p, Ejection: o, Observer: ev}), members, healths, ev +} + +// failTo drives picks until n of them have failed to connect to m +func failTo(t *testing.T, b *lb.Balancer, m *lb.Member, n int) { + t.Helper() + for range 1000 { + if n == 0 { + return + } + pk, ok := b.Pick(lb.Flow{}) + if !ok { + t.Fatal("no pick") + } + if pk.Member() == m { + pk.Done(lb.OutcomeConnectFailed) + n-- + continue + } + pk.Reached() + pk.Done(lb.OutcomeOK) + } + t.Fatalf("%s was not picked often enough to fail %d more times", m.Name(), n) +} + +func live(b *lb.Balancer) []string { + var out []string + for _, m := range b.Pool().Snapshot().Members { + out = append(out, m.Name()) + } + return out +} + +// consecutive connect failures eject a member, even under a strategy that otherwise tracks +// nothing; a success in between starts the count again, and other outcomes never count +func TestEjectionNeedsConsecutiveConnectFailures(t *testing.T) { + b, m, _, ev := ejecting(t, lb.EjectionOptions{Failures: 3, Duration: time.Hour}, 1, 1, 1) + failTo(t, b, m[1], 2) + for range 3 { + pk, _ := b.Pick(lb.Flow{}) + pk.Reached() + pk.Done(lb.OutcomeOK) + } + failTo(t, b, m[1], 2) + if len(live(b)) != 3 { + t.Fatalf("ejected after failures that were not consecutive: %v", live(b)) + } + for range 30 { + pk, _ := b.Pick(lb.Flow{}) + if pk.Member() == m[0] { + pk.Done(lb.OutcomeFailed) + } else { + pk.Done(lb.OutcomeCanceled) + } + } + if len(live(b)) != 3 { + t.Fatalf("an outcome other than a connect failure ejected a member: %v", live(b)) + } + failTo(t, b, m[1], 1) + if got := live(b); !slices.Equal(got, []string{m[0].Name(), m[2].Name()}) { + t.Fatalf("after three consecutive connect failures: %v", got) + } + if !m[1].Stats().Ejected(time.Now()) || m[0].Stats().Ejected(time.Now()) { + t.Error("the stats do not say who is ejected") + } + if got := ev.ejected(); !slices.Equal(got, []string{m[1].Name()}) { + t.Errorf("ejection events = %v", got) + } + // it takes no more flows + for range 20 { + pk, _ := b.Pick(lb.Flow{}) + if pk.Member() == m[1] { + t.Fatal("an ejected member was picked") + } + pk.Done(lb.OutcomeOK) + } +} + +// an ejection with no duration or share configured lasts 30s and spares half the pool +func TestEjectionDefaults(t *testing.T) { + b, m, _, _ := ejecting(t, lb.EjectionOptions{Failures: 1}, 1, 1, 1, 1) + for _, member := range m { + failTo(t, b, member, 1) + } + if got := len(live(b)); got != 2 { + t.Errorf("%d of 4 members live", got) + } + var out *lb.Member + for _, member := range m { + if member.Stats().Ejected(time.Now()) { + out = member + } + } + if out == nil || !out.Stats().Ejected(time.Now().Add(lb.DefaultEjectionDuration-time.Second)) || + out.Stats().Ejected(time.Now().Add(lb.DefaultEjectionDuration+time.Second)) { + t.Error("the default ejection does not last the default duration") + } +} + +func TestEjectionIsOffByDefault(t *testing.T) { + b, m, _, ev := ejecting(t, lb.EjectionOptions{}, 1, 1) + failTo(t, b, m[0], 50) + if len(live(b)) != 2 || len(ev.ejected()) != 0 { + t.Errorf("a balancer with no ejection configured ejected: %v", live(b)) + } + if m[0].Stats().Failures() != 0 { + t.Error("a strategy with no needs and no ejection paid for failure counting") + } +} + +// the pool is never emptied: not below one live member, and not past the configured share +func TestEjectionCap(t *testing.T) { + b, m, healths, _ := ejecting(t, lb.EjectionOptions{Failures: 1, Duration: time.Hour}, 1, 1, 1, 1) + for _, member := range m { + failTo(t, b, member, 1) + } + if got := len(live(b)); got != 2 { + t.Errorf("the default cap of half left %d of 4 members live", got) + } + + all, m2, _, _ := ejecting(t, lb.EjectionOptions{Failures: 1, Duration: time.Hour, MaxPercent: 100}, 1, 1, 1) + for range 3 { + for _, member := range m2 { + if slices.Contains(live(all), member.Name()) && len(live(all)) > 1 { + failTo(t, all, member, 1) + } + } + } + if got := len(live(all)); got != 1 { + t.Errorf("with no percentage cap, %d of 3 members stayed live; the last one always must", got) + } + pk, ok := all.Pick(lb.Flow{}) + if !ok { + t.Fatal("the pool was emptied") + } + pk.Done(lb.OutcomeConnectFailed) + if len(live(all)) != 1 { + t.Error("the last live member was ejected") + } + + // a member the health check has already taken out does not count as live + two, m3, h3, _ := ejecting(t, lb.EjectionOptions{Failures: 1, Duration: time.Hour, MaxPercent: 100}, 1, 1) + h3[0].Set(-1) + failTo(t, two, m3[1], 1) + if got := live(two); len(got) != 1 { + t.Errorf("ejected the only member its health check had left: %v", got) + } + _ = healths +} + +// an ejected member returns when its time is up, through whichever pool is current by then, +// and only if its health check still has it up +func TestEjectionEnds(t *testing.T) { + b, m, healths, _ := ejecting(t, lb.EjectionOptions{Failures: 1, Duration: 60 * time.Millisecond}, 1, 1, 1) + failTo(t, b, m[0], 1) + failTo(t, b, m[2], 0) + if len(live(b)) != 2 { + t.Fatalf("live = %v", live(b)) + } + // membership is swapped while the member is out: the new pool starts without it + swapped, err := lb.NewPool(m, 1) + if err != nil { + t.Fatal(err) + } + defer swapped.Stop() + b.SetPool(swapped) + if len(live(b)) != 2 { + t.Fatalf("a new pool forgot the ejection: %v", live(b)) + } + deadline := time.Now().Add(3 * time.Second) + for len(live(b)) != 3 { + if time.Now().After(deadline) { + t.Fatal("the ejected member never returned to the current pool") + } + time.Sleep(5 * time.Millisecond) + } + + failTo(t, b, m[1], 1) + healths[1].Set(-1) + time.Sleep(120 * time.Millisecond) + if slices.Contains(live(b), m[1].Name()) { + t.Error("a member returned from ejection although its health check has it down") + } + // with no pool left by the time an ejection ends, there is nothing to refresh + failTo(t, b, m[2], 1) + b.SetPool(nil) + time.Sleep(120 * time.Millisecond) +} + +// a flow that fails after its member has left the pool ejects nobody +func TestEjectionIgnoresADepartedMember(t *testing.T) { + b, m, _, ev := ejecting(t, lb.EjectionOptions{Failures: 1, Duration: time.Hour}, 1, 1, 1) + var held lb.Pick + for held.Member() != m[0] { + held.Done(lb.OutcomeOK) + held, _ = b.Pick(lb.Flow{}) + } + without, err := lb.NewPool(m[1:], 1) + if err != nil { + t.Fatal(err) + } + defer without.Stop() + b.SetPool(without) + held.Done(lb.OutcomeConnectFailed) + if len(ev.ejected()) != 0 || len(live(b)) != 2 { + t.Errorf("ejected %v from a pool it had left; live = %v", ev.ejected(), live(b)) + } + // nor does one that fails once the pool is stopped, or gone + held, _ = b.Pick(lb.Flow{}) + without.Stop() + held.Done(lb.OutcomeConnectFailed) + held, _ = b.Pick(lb.Flow{}) + b.SetPool(nil) + held.Done(lb.OutcomeConnectFailed) + if len(ev.ejected()) != 0 { + t.Errorf("ejected %v from a stopped or absent pool", ev.ejected()) + } +} + +// what counts is whether each connect succeeded, in order: a flow that was reached long ago +// and only now ends says nothing about the connects since, and how a member answered never counts +func TestEjectionCountsConnectsNotFlows(t *testing.T) { + b, m, _, _ := ejecting(t, lb.EjectionOptions{Failures: 2, Duration: time.Hour}, 1) + only := m[0] + fail := func() { + pk, _ := b.Pick(lb.Flow{}) + pk.Done(lb.OutcomeConnectFailed) + } + // fail, connect (and stay open), fail: the failures are not consecutive + fail() + long, _ := b.Pick(lb.Flow{}) + long.Reached() + if only.Stats().ConnectFailures() != 0 { + t.Fatal("a successful connect did not end the run of failures") + } + fail() + if only.Stats().ConnectFailures() != 1 { + t.Fatalf("connect failures = %d", only.Stats().ConnectFailures()) + } + // the long flow ending well does not excuse the failure since + long.Done(lb.OutcomeOK) + if only.Stats().ConnectFailures() != 1 { + t.Error("a flow's end reset the count of failed connects") + } + // a failed answer on a flow that was reached counts toward nothing + for range 5 { + pk, _ := b.Pick(lb.Flow{}) + pk.Done(lb.OutcomeFailed) + } + if only.Stats().ConnectFailures() != 1 || only.Stats().Failures() < 5 { + t.Errorf("connect failures = %d, failures = %d", only.Stats().ConnectFailures(), only.Stats().Failures()) + } + // reports on a pick that has no balancer, or one that does not eject, are harmless + lb.Pick{}.Reached() + plain := lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: b.Pool()}) + pk, _ := plain.Pick(lb.Flow{}) + pk.Reached() + pk.Done(lb.OutcomeOK) + lb.LeafPick{}.Reached() +} diff --git a/pkg/lb/example_test.go b/pkg/lb/example_test.go new file mode 100644 index 000000000..d64202a65 --- /dev/null +++ b/pkg/lb/example_test.go @@ -0,0 +1,125 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb_test + +import ( + "fmt" + "io" + "net" + "net/http" + "net/http/httptest" + "sync/atomic" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" +) + +// Balance requests across http.Handlers, giving one of them twice the share. +func Example_httpHandlers() { + named := func(name string) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { fmt.Fprint(w, name) }) + } + pool, err := lb.NewPool([]*lb.Member{ + lb.NewMember(lb.MemberOptions{Name: "a", Value: named("a")}), + lb.NewMember(lb.MemberOptions{Name: "b", Weight: 2, Value: named("b")}), + }, 0) + if err != nil { + panic(err) + } + defer pool.Stop() + // rr.New starts its rotation at a random turn; NewAt makes this example's order repeatable + balancer := lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: pool}) + + front := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + pick, ok := balancer.Pick(lb.Flow{}) + if !ok { + http.Error(w, "no backend available", http.StatusBadGateway) + return + } + defer pick.Done(lb.OutcomeOK) + pick.Member().Value.(http.Handler).ServeHTTP(w, r) + }) + + for range 6 { + w := httptest.NewRecorder() + front.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/", nil)) + fmt.Print(w.Body.String(), " ") + } + // Output: b b a b b a +} + +// upDown is the smallest possible health source: the pool follows it on Refresh +type upDown struct{ status atomic.Int32 } + +func (h *upDown) Get() int32 { return h.status.Load() } + +// Balance TCP dials across addresses, skipping a member whose health source reports it down. +func Example_tcpDials() { + serve := func(name string) string { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + panic(err) + } + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + _, _ = io.WriteString(conn, name) + _ = conn.Close() + } + }() + return ln.Addr().String() + } + health := map[string]*upDown{"one": {}, "two": {}} + pool, err := lb.NewPool([]*lb.Member{ + lb.NewMember(lb.MemberOptions{Name: "one", Health: health["one"], Value: serve("one")}), + lb.NewMember(lb.MemberOptions{Name: "two", Health: health["two"], Value: serve("two")}), + }, 0) + if err != nil { + panic(err) + } + defer pool.Stop() + balancer := lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: pool}) + + dial := func() string { + pick, ok := balancer.Pick(lb.Flow{}) + if !ok { + return "refused" + } + start := time.Now() + conn, err := net.Dial("tcp", pick.Member().Value.(string)) + if err != nil { + pick.Done(lb.OutcomeConnectFailed) + return "refused" + } + pick.Established(time.Since(start)) + defer pick.Done(lb.OutcomeOK) + defer conn.Close() + reply, _ := io.ReadAll(conn) + return string(reply) + } + + fmt.Println(dial(), dial(), dial()) + health["two"].status.Store(-1) + pool.Refresh() + fmt.Println(dial(), dial()) + // Output: + // two one two + // one one +} diff --git a/pkg/lb/hrw/hrw.go b/pkg/lb/hrw/hrw.go new file mode 100644 index 000000000..e8e034c88 --- /dev/null +++ b/pkg/lb/hrw/hrw.go @@ -0,0 +1,99 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package hrw is the highest random weight (rendezvous hashing) strategy: a flow's key and +// each member's name hash to a score, and the highest score wins. The same key reaches the +// same member from every process, and losing a member moves only that member's keys. It reads +// every member per pick. +package hrw + +import ( + "math" + "math/rand/v2" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Name is the strategy's name. +const Name = "highest_random_weight" + +// New returns a highest random weight selector. +func New() lb.Selector { + return selector{} +} + +type selector struct{} + +func (selector) Name() string { return Name } + +func (selector) Needs() lb.Needs { return lb.NeedKey } + +func (selector) Prepare(snap *lb.Snapshot) lb.Prepared { + members := snap.Members + p := &prepared{members: members, hashes: make([]uint64, len(members))} + uniform := true + for i, m := range members { + p.hashes[i] = m.Hash() + if m.Name() == "" { + // unnamed members share one name hash; their position tells them apart + p.hashes[i] = lb.Mix(uint64(i) + 1) // #nosec G115 -- a slice index + } + if m.Weight() != members[0].Weight() { + uniform = false + } + } + if !uniform { + p.weights = make([]float64, len(members)) + for i, m := range members { + p.weights[i] = float64(m.Weight()) + } + } + return p +} + +type prepared struct { + members []*lb.Member + hashes []uint64 + // nil when every member has the same weight, which skips the logarithm + weights []float64 +} + +func (p *prepared) Select(f lb.Flow) *lb.Member { + key := f.Key + if !f.HasKey { + // nothing identifies the flow, so it has no affinity to keep: spread it + key = rand.Uint64() // #nosec G404 -- load spreading, not a secret + } + best := 0 + if p.weights == nil { + var high uint64 + for i, h := range p.hashes { + if s := lb.Mix(key ^ h); s > high || i == 0 { + best, high = i, s + } + } + return p.members[best] + } + // weighted rendezvous: the score -w/ln(u), with u uniform in (0,1), gives each member a + // share of the key space proportional to its weight + high := math.Inf(-1) + for i, h := range p.hashes { + u := (float64(lb.Mix(key^h)>>11) + 0.5) / (1 << 53) + if s := -p.weights[i] / math.Log(u); s > high { + best, high = i, s + } + } + return p.members[best] +} diff --git a/pkg/lb/hrw/hrw_test.go b/pkg/lb/hrw/hrw_test.go new file mode 100644 index 000000000..9beda302e --- /dev/null +++ b/pkg/lb/hrw/hrw_test.go @@ -0,0 +1,173 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package hrw_test + +import ( + "strconv" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/hrw" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" +) + +func TestConformance(t *testing.T) { + lbtest.Run(t, hrw.New, lbtest.Options{}) +} + +func TestWeights(t *testing.T) { + lbtest.RunWeighted(t, hrw.New, lbtest.WeightOptions{Keyed: true}) +} + +func BenchmarkSelect(b *testing.B) { + lbtest.Bench(b, hrw.New) +} + +func named(names []string, weights []int) []*lb.Member { + members := make([]*lb.Member, len(names)) + for i, n := range names { + w := 1 + if weights != nil { + w = weights[i] + } + members[i] = lb.NewMember(lb.MemberOptions{Name: n, Weight: w}) + } + return members +} + +func balancerOver(t *testing.T, members []*lb.Member) *lb.Balancer { + t.Helper() + p, err := lb.NewPool(members, 0) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return lb.NewBalancer(hrw.New(), lb.BalancerOptions{Pool: p}) +} + +func owner(t *testing.T, b *lb.Balancer, key string) string { + t.Helper() + pk, ok := b.Pick(lb.Flow{Key: lb.HashString(key), HasKey: true}) + if !ok { + t.Fatal("no pick") + } + return pk.Member().Name() +} + +func TestIdentity(t *testing.T) { + s := hrw.New() + if s.Name() != hrw.Name || s.Needs() != lb.NeedKey { + t.Errorf("name %q needs %b", s.Name(), s.Needs()) + } +} + +// the mapping is a fixed function of key and member names: every replica, and every restart +// of one, must send a client to the same member. A change here reshuffles deployed affinity. +func TestGoldenMapping(t *testing.T) { + names := []string{"cache-a", "cache-b", "cache-c", "cache-d"} + uniform := balancerOver(t, named(names, nil)) + weighted := balancerOver(t, named(names, []int{1, 2, 3, 4})) + // the same members in another order, as a second replica might list them + shuffled := balancerOver(t, named([]string{"cache-c", "cache-a", "cache-d", "cache-b"}, nil)) + golden := []struct{ key, uniform, weighted string }{ + {"10.0.0.1", "cache-a", "cache-a"}, + {"10.0.0.2", "cache-d", "cache-d"}, + {"10.0.0.4", "cache-a", "cache-d"}, + {"192.0.2.10", "cache-a", "cache-a"}, + {"2001:db8::1", "cache-d", "cache-d"}, + {"tenant-42", "cache-d", "cache-d"}, + {"tenant-43", "cache-a", "cache-a"}, + {"tenant-44", "cache-b", "cache-d"}, + {"", "cache-d", "cache-d"}, + } + for _, g := range golden { + if got := owner(t, uniform, g.key); got != g.uniform { + t.Errorf("uniform owner of %q = %s, want %s", g.key, got, g.uniform) + } + if got := owner(t, weighted, g.key); got != g.weighted { + t.Errorf("weighted owner of %q = %s, want %s", g.key, got, g.weighted) + } + if got := owner(t, shuffled, g.key); got != g.uniform { + t.Errorf("owner of %q depends on member order: %s, want %s", g.key, got, g.uniform) + } + } +} + +// losing a member moves only the keys it owned; gaining one takes only about its share +func TestMinimalDisruption(t *testing.T) { + for name, weights := range map[string][]int{"uniform": nil, "weighted": {3, 1, 2, 1, 3}} { + t.Run(name, func(t *testing.T) { + names := []string{"m0", "m1", "m2", "m3", "m4"} + all := named(names, weights) + before := balancerOver(t, all) + after := balancerOver(t, all[:4]) + const keys = 20000 + var moved, owned int + for i := range keys { + key := "client-" + strconv.Itoa(i) + was, is := owner(t, before, key), owner(t, after, key) + if was == "m4" { + owned++ + continue + } + if was != is { + moved++ + } + } + if moved != 0 { + t.Errorf("%d keys that m4 never owned moved when it left", moved) + } + share := float64(all[4].Weight()) + var total float64 + for _, m := range all { + total += float64(m.Weight()) + } + if got, want := float64(owned)/keys, share/total; got < want-0.02 || got > want+0.02 { + t.Errorf("m4 owned %.3f of the keys, want about %.3f", got, want) + } + }) + } +} + +// a flow with nothing to key on has no affinity to keep, so it is spread, not piled on one member +func TestKeylessFlowsSpread(t *testing.T) { + b := balancerOver(t, named([]string{"a", "b", "c"}, nil)) + seen := make(map[string]int) + for range 3000 { + pk, _ := b.Pick(lb.Flow{}) + seen[pk.Member().Name()]++ + } + for _, n := range []string{"a", "b", "c"} { + if seen[n] < 800 { + t.Errorf("%s took %d of 3000 keyless flows", n, seen[n]) + } + } +} + +// unnamed members are told apart by position, so they do not all score alike +func TestUnnamedMembersShareKeys(t *testing.T) { + b := balancerOver(t, []*lb.Member{ + lb.NewMember(lb.MemberOptions{}), lb.NewMember(lb.MemberOptions{}), lb.NewMember(lb.MemberOptions{}), + }) + seen := make(map[*lb.Member]int) + for i := range 3000 { + pk, _ := b.Pick(lb.Flow{Key: lb.HashString(strconv.Itoa(i)), HasKey: true}) + seen[pk.Member()]++ + } + if len(seen) != 3 { + t.Errorf("keys reached %d of 3 unnamed members", len(seen)) + } +} diff --git a/pkg/lb/key.go b/pkg/lb/key.go new file mode 100644 index 000000000..bd874ba02 --- /dev/null +++ b/pkg/lb/key.go @@ -0,0 +1,80 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import "net/netip" + +// Flow keys are hashed with no per-process seed: replicas behind one front door must agree on +// where a key lands, and a restart must not move it. + +// Mix is a 64-bit finalizer: every input bit affects every output bit. Strategies use it to +// combine a flow key with a member hash. +func Mix(x uint64) uint64 { + x ^= x >> 30 + x *= 0xbf58476d1ce4e5b9 + x ^= x >> 27 + x *= 0x94d049bb133111eb + x ^= x >> 31 + return x +} + +// HashString returns the flow key of a string. +func HashString(s string) uint64 { + return Mix(hashString(s)) +} + +// HashBytes returns the flow key of a byte slice; equal to HashString of the same bytes. +func HashBytes(b []byte) uint64 { + h := fnvOffset64 + for _, c := range b { + h ^= uint64(c) + h *= fnvPrime64 + } + return Mix(h) +} + +// HashFold returns the flow key of a string with ASCII letters folded to lower case, for +// identifiers such as host names that compare without regard to case. +func HashFold(s string) uint64 { + h := fnvOffset64 + for i := range len(s) { + c := s[i] + if c >= 'A' && c <= 'Z' { + c += 'a' - 'A' + } + h ^= uint64(c) + h *= fnvPrime64 + } + return Mix(h) +} + +// HashAddr returns the flow key of an IP address, never of a port. An IPv6 address is masked +// to v6Prefix bits first, since privacy addressing rotates a client's low bits; a prefix +// outside 1-128 keeps the whole address. An IPv4-mapped IPv6 address keys as its IPv4 form. +func HashAddr(addr netip.Addr, v6Prefix int) uint64 { + addr = addr.Unmap() + if addr.Is6() && v6Prefix >= 1 && v6Prefix < 128 { + if p, err := addr.Prefix(v6Prefix); err == nil { + addr = p.Addr() + } + } + if addr.Is4() { + b := addr.As4() + return HashBytes(b[:]) + } + b := addr.As16() + return HashBytes(b[:]) +} diff --git a/pkg/lb/lb.go b/pkg/lb/lb.go new file mode 100644 index 000000000..62c453cd4 --- /dev/null +++ b/pkg/lb/lb.go @@ -0,0 +1,39 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package lb is a protocol-neutral load-balancing core: pool membership, health-filtered +// immutable snapshots and per-member runtime stats. It imports only the standard library. +package lb + +// Health reports a member's current health status. Higher is healthier; a pool compares the +// value against its floor. A member without a Health is always eligible. +type Health interface { + Get() int32 +} + +// Notifier is optionally implemented by a Health that can announce its transitions, which lets +// a pool rebuild its snapshot inside the transition instead of waiting for Refresh. +type Notifier interface { + // OnChange registers fn to be called, synchronously and outside any lock the Notifier + // holds, after each status change with the previous and the new status. + OnChange(fn func(prev, next int32)) Subscription +} + +// Subscription ends a Notifier registration. +type Subscription interface { + // Unsubscribe never blocks on a callback that is already running, so it may be called + // from inside one; a callback that has already been captured may still run once. + Unsubscribe() +} diff --git a/pkg/lb/lbtest/health.go b/pkg/lb/lbtest/health.go new file mode 100644 index 000000000..46390c326 --- /dev/null +++ b/pkg/lb/lbtest/health.go @@ -0,0 +1,85 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lbtest + +import ( + "slices" + "sync" + "sync/atomic" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Health is a settable lb.Health that announces its transitions the way lb.Notifier asks: +// synchronously, outside its own lock, with an Unsubscribe that never waits. +type Health struct { + status atomic.Int32 + mtx sync.Mutex + subs []*subscription +} + +type subscription struct { + health *Health + fn func(prev, next int32) + dead atomic.Bool +} + +// NewHealth returns a Health at the provided status. +func NewHealth(status int32) *Health { + h := &Health{} + h.status.Store(status) + return h +} + +// Get returns the current status. +func (h *Health) Get() int32 { + return h.status.Load() +} + +// Set stores a status and runs the registered callbacks when it is a change. +func (h *Health) Set(next int32) { + prev := h.status.Swap(next) + if prev == next { + return + } + h.mtx.Lock() + subs := h.subs + h.mtx.Unlock() + for _, s := range subs { + if !s.dead.Load() { + s.fn(prev, next) + } + } +} + +// OnChange registers fn to run after each change of status. +func (h *Health) OnChange(fn func(prev, next int32)) lb.Subscription { + s := &subscription{health: h, fn: fn} + h.mtx.Lock() + defer h.mtx.Unlock() + h.subs = append(slices.Clone(h.subs), s) + return s +} + +func (s *subscription) Unsubscribe() { + if s.dead.Swap(true) { + return + } + h := s.health + h.mtx.Lock() + defer h.mtx.Unlock() + h.subs = slices.DeleteFunc(slices.Clone(h.subs), func(o *subscription) bool { return o == s }) +} diff --git a/pkg/lb/lbtest/lbtest.go b/pkg/lb/lbtest/lbtest.go new file mode 100644 index 000000000..86cc0afe7 --- /dev/null +++ b/pkg/lb/lbtest/lbtest.go @@ -0,0 +1,523 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package lbtest is the conformance suite for lb.Selector implementations. A strategy that +// passes Run honors the contract the balancer and every plane adapter rely on. +package lbtest + +import ( + "math/rand/v2" + "slices" + "strconv" + "sync" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Options describes what a strategy promises beyond the common contract. +type Options struct { + // ExactWeights holds the strategy to exact apportionment: every run of total-weight + // consecutive picks against a stable pool gives each member exactly its weight. + ExactWeights bool +} + +const ( + passing int32 = 1 + failing int32 = -1 + floor = 1 +) + +// flows yields keyed flows from a fixed seed, so a run is reproducible +type flows struct{ r *rand.Rand } + +func newFlows() *flows { + return &flows{r: rand.New(rand.NewPCG(0x1b, 0x5eed))} // #nosec G404 -- reproducible test keys, not secrets +} + +func (f *flows) next() lb.Flow { + return lb.Flow{Key: f.r.Uint64(), HasKey: true} +} + +// Members returns one named, passing member per weight, with the Health that drives each. +func Members(weights ...int) ([]*lb.Member, []*Health) { + members := make([]*lb.Member, len(weights)) + healths := make([]*Health, len(weights)) + for i, w := range weights { + healths[i] = NewHealth(passing) + members[i] = lb.NewMember(lb.MemberOptions{ + Name: "member-" + strconv.Itoa(i), Weight: w, Health: healths[i], Value: i, + }) + } + return members, healths +} + +func newPool(t reporter, members []*lb.Member) *lb.Pool { + t.Helper() + p, err := lb.NewPool(members, floor) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return p +} + +func uniform(n int) []int { + weights := make([]int, n) + for i := range weights { + weights[i] = 1 + } + return weights +} + +// Run holds the selectors that newSelector returns to the lb.Selector contract. Each call of +// newSelector must return a new instance. +func Run(t *testing.T, newSelector func() lb.Selector, o Options) { + run(realT{t}, newSelector, o) +} + +func run(t suiteT, newSelector func() lb.Selector, o Options) { + t.Run("identity", func(t suiteT) { testIdentity(t, newSelector) }) + t.Run("no eligible member", func(t suiteT) { testNoEligibleMember(t, newSelector) }) + t.Run("picks only eligible members", func(t suiteT) { testOnlyEligible(t, newSelector) }) + t.Run("every member is reachable", func(t suiteT) { testReachable(t, newSelector) }) + t.Run("repick excludes the failed member", func(t suiteT) { testRepick(t, newSelector) }) + t.Run("in-flight accounting balances", func(t suiteT) { testInflight(t, newSelector) }) + t.Run("adversarial snapshots terminate", func(t suiteT) { testAdversarial(t, newSelector) }) + t.Run("zero allocations", func(t suiteT) { testZeroAlloc(t, newSelector) }) + t.Run("concurrent picks and swaps", func(t suiteT) { testConcurrent(t, newSelector) }) + if o.ExactWeights { + t.Run("exact weights", func(t suiteT) { testExactWeights(t, newSelector) }) + t.Run("exact weights after a live weight change", func(t suiteT) { + testExactAfterWeightChange(t, newSelector) + }) + } +} + +func testIdentity(t suiteT, newSelector func() lb.Selector) { + a, b := newSelector(), newSelector() + if a == nil || a.Name() == "" { + t.Fatal("a selector must have a name") + } + if first := a.Needs(); first != b.Needs() || first != a.Needs() { + t.Error("a strategy's needs must not vary") + } + if bal := lb.NewBalancer(a); bal.Needs() != a.Needs() || bal.Selector() != a { + t.Error("the balancer does not report its selector") + } +} + +func testNoEligibleMember(t suiteT, newSelector func() lb.Selector) { + f := newFlows() + b := lb.NewBalancer(newSelector()) + if _, ok := b.Pick(f.next()); ok { + t.Error("picked with no pool") + } + b.SetPool(newPool(t, nil)) + if _, ok := b.Pick(f.next()); ok { + t.Error("picked from an empty pool") + } + members, healths := Members(1, 1) + for _, h := range healths { + h.Set(failing) + } + b.SetPool(newPool(t, members)) + if _, ok := b.Pick(f.next()); ok { + t.Error("picked from a pool with no eligible member") + } + if _, ok := b.Repick(f.next(), nil); ok { + t.Error("repicked from a pool with no eligible member") + } + // a member that recovers is picked by the very next call + healths[1].Set(passing) + pk, ok := b.Pick(f.next()) + if !ok || pk.Member() != members[1] { + t.Error("the recovered member was not picked") + } + pk.Done(lb.OutcomeOK) + b.SetPool(nil) + if _, ok := b.Pick(f.next()); ok || b.Pool() != nil { + t.Error("picked after the pool was removed") + } +} + +func testOnlyEligible(t suiteT, newSelector func() lb.Selector) { + f := newFlows() + for _, weights := range [][]int{{1}, {1, 1}, {4}, {1, 3}, {2, 1, 5, 1, 1}} { + members, healths := Members(weights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + for down := range members { + healths[down].Set(failing) + for range 200 { + pk, ok := b.Pick(f.next()) + if len(members) == 1 { + if ok { + t.Fatalf("weights %v: picked the only member while it was failing", weights) + } + continue + } + if !ok { + t.Fatalf("weights %v: no pick with eligible members", weights) + } + if pk.Member() == members[down] || !slices.Contains(members, pk.Member()) { + t.Fatalf("weights %v: picked %q, which is not eligible", weights, pk.Member().Name()) + } + pk.Done(lb.OutcomeOK) + } + healths[down].Set(passing) + } + } +} + +func testReachable(t suiteT, newSelector func() lb.Selector) { + f := newFlows() + for _, weights := range [][]int{uniform(4), {1, 2, 3, 4}} { + members, _ := Members(weights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + seen := make(map[*lb.Member]int) + // some work is always in flight: a strategy that ranks members may, by design, give + // an idle pool's every flow to its best member + var held [32]lb.Pick + for i := range 4000 { + pk, ok := b.Pick(f.next()) + if !ok { + t.Fatal("no pick with eligible members") + } + seen[pk.Member()]++ + held[i%len(held)].Done(lb.OutcomeOK) + held[i%len(held)] = pk + } + for _, pk := range held { + pk.Done(lb.OutcomeOK) + } + for _, m := range members { + if seen[m] == 0 { + t.Errorf("weights %v: %s was never picked in 4000 flows", weights, m.Name()) + } + } + } +} + +func testRepick(t suiteT, newSelector func() lb.Selector) { + f := newFlows() + members, _ := Members(1, 3, 1) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + for range 300 { + pk, ok := b.Pick(f.next()) + if !ok { + t.Fatal("no pick") + } + pk.Done(lb.OutcomeConnectFailed) + again, ok := b.Repick(f.next(), pk.Member()) + if !ok || again.Member() == pk.Member() || !slices.Contains(members, again.Member()) { + t.Fatalf("repick after %s chose %v", pk.Member().Name(), again.Member()) + } + again.Done(lb.OutcomeOK) + } + only, _ := Members(1) + single := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, only)}) + if _, ok := single.Repick(f.next(), only[0]); ok { + t.Error("repicked the failed member of a pool of one") + } +} + +func testInflight(t suiteT, newSelector func() lb.Selector) { + f := newFlows() + members, _ := Members(1, 2, 1) + sel := newSelector() + b := lb.NewBalancer(sel, lb.BalancerOptions{Pool: newPool(t, members)}) + var open []lb.Pick + for range 64 { + pk, ok := b.Pick(f.next()) + if !ok { + t.Fatal("no pick") + } + pk.Established(time.Millisecond) + pk.FirstByte() + open = append(open, pk) + } + var inflight int64 + for _, m := range members { + inflight += m.Stats().Inflight() + } + want := int64(0) + if sel.Needs().Has(lb.NeedInflight) { + want = int64(len(open)) + } + if inflight != want { + t.Errorf("in-flight while %d picks are open = %d, want %d", len(open), inflight, want) + } + for i, pk := range open { + pk.Done(lb.Outcome(i % 4)) // #nosec G115 -- one of the four outcomes + } + for _, m := range members { + if got := m.Stats().Inflight(); got != 0 { + t.Errorf("%s holds %d in flight after every pick was done", m.Name(), got) + } + } + // the zero Pick, as returned with false, is safe to report on + var none lb.Pick + none.Done(lb.OutcomeOK) + if none.Member() != nil { + t.Error("the zero pick has a member") + } +} + +func testAdversarial(t suiteT, newSelector func() lb.Selector) { + done := make(chan struct{}) + var failure string + go func() { + defer close(done) + f := newFlows() + for _, weights := range [][]int{{1 << 30, 1}, {1, 1 << 30, 1 << 20, 1}, uniform(257)} { + members, healths := Members(weights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPoolQuiet(members)}) + // all but one ineligible, then all eligible + for _, h := range healths[1:] { + h.Set(failing) + } + for range 2000 { + pk, ok := b.Pick(f.next()) + if !ok || pk.Member() != members[0] { + failure = "did not pick the only eligible member" + return + } + pk.Done(lb.OutcomeOK) + } + for _, h := range healths[1:] { + h.Set(passing) + } + for range 20000 { + pk, ok := b.Pick(f.next()) + if !ok || !slices.Contains(members, pk.Member()) { + failure = "picked outside the snapshot" + return + } + pk.Done(lb.OutcomeOK) + } + b.Pool().Stop() + } + }() + select { + case <-done: + if failure != "" { + t.Error(failure) + } + case <-time.After(terminationLimit): + t.Fatal("selection did not terminate") + } +} + +// terminationLimit is how long the adversarial snapshots may take before selection is +// judged not to terminate +var terminationLimit = 30 * time.Second + +// newPoolQuiet builds a pool off the test goroutine, where t.Fatal is not allowed +func newPoolQuiet(members []*lb.Member) *lb.Pool { + p, err := lb.NewPool(members, floor) + if err != nil { + panic(err) + } + return p +} + +func testZeroAlloc(t suiteT, newSelector func() lb.Selector) { + for _, weights := range [][]int{uniform(6), {3, 1, 3, 1, 3, 1}} { + members, _ := Members(weights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + flow := newFlows().next() + // the first pick of a snapshot prepares it + if pk, ok := b.Pick(flow); ok { + pk.Done(lb.OutcomeOK) + } + allocs := testing.AllocsPerRun(1000, func() { + pk, ok := b.Pick(flow) + if !ok { + t.Fatal("no pick") + } + pk.Done(lb.OutcomeOK) + }) + if allocs != 0 { + t.Errorf("weights %v: a pick allocates %v times", weights, allocs) + } + } +} + +func testConcurrent(t suiteT, newSelector func() lb.Selector) { + members, healths := Members(1, 2, 1, 3, 1, 1) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + stop := make(chan struct{}) + var churn sync.WaitGroup + churn.Go(func() { + pools := []*lb.Pool{} + defer func() { + for _, p := range pools { + p.Stop() + } + }() + for i := 0; ; i++ { + select { + case <-stop: + return + default: + } + // member 0 stays up so a pick is always possible + h := healths[1+i%(len(healths)-1)] + h.Set(failing) + if i%8 == 0 { + p := newPoolQuiet(members[:2+i%(len(members)-1)]) + pools = append(pools, p) + b.SetPool(p) + } + h.Set(passing) + } + }) + var pickers sync.WaitGroup + for range 8 { + pickers.Go(func() { + f := newFlows() + for range 5000 { + pk, ok := b.Pick(f.next()) + if !ok { + t.Error("no pick although one member never fails") + return + } + if !slices.Contains(members, pk.Member()) { + t.Error("picked a member that was never in the pool") + return + } + pk.Done(lb.OutcomeOK) + } + }) + } + pickers.Wait() + close(stop) + churn.Wait() + for _, m := range members { + if got := m.Stats().Inflight(); got != 0 { + t.Errorf("%s holds %d in flight after every pick was done", m.Name(), got) + } + } +} + +func pickSequence(t reporter, b *lb.Balancer, n int) []*lb.Member { + t.Helper() + f := newFlows() + seq := make([]*lb.Member, n) + for i := range seq { + pk, ok := b.Pick(f.next()) + if !ok { + t.Fatal("no pick") + } + seq[i] = pk.Member() + pk.Done(lb.OutcomeOK) + } + return seq +} + +// assertWindows fails unless every run of total-weight consecutive picks is exact +func assertWindows(t reporter, seq, members []*lb.Member) { + t.Helper() + var total int + for _, m := range members { + total += m.Weight() + } + counts := make(map[*lb.Member]int, len(members)) + for i, m := range seq { + counts[m]++ + if i >= total { + counts[seq[i-total]]-- + } + if i < total-1 { + continue + } + for _, want := range members { + if counts[want] != want.Weight() { + t.Fatalf("picks %d-%d gave %s %d, want exactly its weight %d", + i-total+1, i, want.Name(), counts[want], want.Weight()) + } + } + } +} + +func testExactWeights(t suiteT, newSelector func() lb.Selector) { + // a pool large enough to leave any small-pool fast path a strategy may have + large := make([]int, 67) + for i := range large { + large[i] = 1 + i%4 + } + for _, weights := range [][]int{uniform(1), uniform(5), {4}, {1, 3, 2}, {7, 1, 1}, {2, 2}, {1, 1, 1, 9}, large} { + members, _ := Members(weights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + var total int + for _, w := range weights { + total += w + } + assertWindows(t, pickSequence(t, b, 9*total+3), members) + } +} + +func testExactAfterWeightChange(t suiteT, newSelector func() lb.Selector) { + members, healths := Members(1, 3, 2) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + pickSequence(t, b, 17) + // member 1 is rebuilt with a new weight, keeping its health and stats, in a new pool + reweighted := slices.Clone(members) + reweighted[1] = lb.NewMember(lb.MemberOptions{ + Name: members[1].Name(), Weight: 5, Health: healths[1], Stats: members[1].Stats(), + }) + b.SetPool(newPool(t, reweighted)) + assertWindows(t, pickSequence(t, b, 40), reweighted) + // a member that drops out leaves the rest exact among themselves + healths[0].Set(failing) + assertWindows(t, pickSequence(t, b, 40), reweighted[1:]) +} + +// Bench measures a pick at several pool sizes, uniform and weighted, from one goroutine and +// from many. A strategy should stay flat, or close to it, as the pool grows. +func Bench(b *testing.B, newSelector func() lb.Selector) { + for _, n := range []int{2, 8, 64, 512} { + for _, shape := range []string{"uniform", "weighted"} { + weights := uniform(n) + if shape == "weighted" { + for i := 0; i < n; i += 2 { + weights[i] = 3 + } + } + members, _ := Members(weights...) + bal := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(b, members)}) + name := shape + "/n=" + strconv.Itoa(n) + b.Run(name, func(b *testing.B) { + f := newFlows() + b.ReportAllocs() + for b.Loop() { + pk, _ := bal.Pick(f.next()) + pk.Done(lb.OutcomeOK) + } + }) + b.Run(name+"/parallel", func(b *testing.B) { + b.ReportAllocs() + b.RunParallel(func(pb *testing.PB) { + f := newFlows() + for pb.Next() { + pk, _ := bal.Pick(f.next()) + pk.Done(lb.OutcomeOK) + } + }) + }) + } + } +} diff --git a/pkg/lb/lbtest/lbtest_test.go b/pkg/lb/lbtest/lbtest_test.go new file mode 100644 index 000000000..c2e06eb88 --- /dev/null +++ b/pkg/lb/lbtest/lbtest_test.go @@ -0,0 +1,104 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lbtest_test + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" +) + +// byKey is a minimal strategy that declares every need, so the suite's accounting checks run +type byKey struct{} + +type byKeyPrepared struct{ members []*lb.Member } + +func newByKey() lb.Selector { return &byKey{} } + +func (*byKey) Name() string { return "by_key" } + +func (*byKey) Needs() lb.Needs { return lb.NeedKey | lb.NeedInflight | lb.NeedLatency } + +func (*byKey) Prepare(s *lb.Snapshot) lb.Prepared { return &byKeyPrepared{members: s.Members} } + +func (p *byKeyPrepared) Select(f lb.Flow) *lb.Member { + return p.members[f.Key%uint64(len(p.members))] +} + +func TestSuiteAcceptsAConformingSelector(t *testing.T) { + lbtest.Run(t, newByKey, lbtest.Options{}) +} + +// byWeight gives each member a span of the key space as wide as its weight +type byWeight struct{} + +type byWeightPrepared struct { + members []*lb.Member + total uint64 +} + +func (byWeight) Name() string { return "by_weight" } + +func (byWeight) Needs() lb.Needs { return lb.NeedKey } + +func (byWeight) Prepare(s *lb.Snapshot) lb.Prepared { + p := &byWeightPrepared{members: s.Members} + for _, m := range s.Members { + p.total += uint64(m.Weight()) + } + return p +} + +func (p *byWeightPrepared) Select(f lb.Flow) *lb.Member { + k := f.Key % p.total + for _, m := range p.members { + if w := uint64(m.Weight()); k < w { + return m + } else { + k -= w + } + } + return p.members[0] +} + +func TestWeightedSuiteAcceptsAStrategyThatHonorsWeights(t *testing.T) { + lbtest.RunWeighted(t, func() lb.Selector { return byWeight{} }, lbtest.WeightOptions{Keyed: true}) + lbtest.Run(t, func() lb.Selector { return byWeight{} }, lbtest.Options{}) +} + +func TestHealth(t *testing.T) { + h := lbtest.NewHealth(1) + var calls int + sub := h.OnChange(func(prev, next int32) { + calls++ + if prev != 1 || next != -1 { + t.Errorf("transition = %d to %d", prev, next) + } + }) + h.Set(1) + h.Set(-1) + sub.Unsubscribe() + sub.Unsubscribe() + h.Set(1) + if calls != 1 || h.Get() != 1 { + t.Errorf("calls = %d, status = %d", calls, h.Get()) + } +} + +func BenchmarkSuite(b *testing.B) { + lbtest.Bench(b, newByKey) +} diff --git a/pkg/lb/lbtest/rejection_test.go b/pkg/lb/lbtest/rejection_test.go new file mode 100644 index 000000000..09ecfb368 --- /dev/null +++ b/pkg/lb/lbtest/rejection_test.go @@ -0,0 +1,432 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lbtest + +import ( + "flag" + "fmt" + "slices" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// recorder is a suiteT that notes which checks failed instead of failing a real test, so the +// suite can be shown to reject a strategy that breaks its contract +type recorder struct { + mtx sync.Mutex + path string + failed map[string][]string + cleanups []func() +} + +// fatal unwinds a check the way t.Fatal ends a test +type fatal struct{} + +func newRecorder() *recorder { + return &recorder{failed: make(map[string][]string)} +} + +func (r *recorder) note(msg string) { + r.mtx.Lock() + defer r.mtx.Unlock() + r.failed[r.path] = append(r.failed[r.path], msg) +} + +func (r *recorder) Helper() {} + +func (r *recorder) Error(args ...any) { r.note(fmt.Sprint(args...)) } + +func (r *recorder) Errorf(format string, args ...any) { r.note(fmt.Sprintf(format, args...)) } + +func (r *recorder) Fatal(args ...any) { + r.note(fmt.Sprint(args...)) + panic(fatal{}) +} + +func (r *recorder) Fatalf(format string, args ...any) { + r.note(fmt.Sprintf(format, args...)) + panic(fatal{}) +} + +func (r *recorder) Cleanup(fn func()) { + r.mtx.Lock() + defer r.mtx.Unlock() + r.cleanups = append(r.cleanups, fn) +} + +func (r *recorder) Run(name string, check func(suiteT)) bool { + r.mtx.Lock() + r.path = name + r.mtx.Unlock() + func() { + defer func() { + if v := recover(); v != nil { + if _, ok := v.(fatal); !ok { + panic(v) + } + } + }() + check(r) + }() + r.mtx.Lock() + cleanups := r.cleanups + r.cleanups = nil + r.mtx.Unlock() + for _, fn := range slices.Backward(cleanups) { + fn() + } + r.mtx.Lock() + defer r.mtx.Unlock() + return len(r.failed[name]) == 0 +} + +func (r *recorder) failures(check string) []string { + r.mtx.Lock() + defer r.mtx.Unlock() + return slices.Clone(r.failed[check]) +} + +func (r *recorder) failedChecks() []string { + r.mtx.Lock() + defer r.mtx.Unlock() + var out []string + for name, msgs := range r.failed { + if len(msgs) > 0 { + out = append(out, name) + } + } + slices.Sort(out) + return out +} + +// broken is a strategy built from the ways one can break the contract. Left at its zero +// value it spreads flows by key, correctly. +type broken struct { + name string + needs lb.Needs + stale bool // keeps selecting from the first snapshot it ever saw + first bool // always selects the first member + declines bool // selects nothing + allocates bool // allocates on the selection path + hangs bool // never returns from a selection over more than one member + foreign bool // selects a member that belongs to no pool, once the pool has company + drifts bool // changes what it needs between calls + rotates bool // plain rotation, which is exact only when it ignores no weights + exact bool // weighted rotation over contiguous spans: exact apportionment + // keeps the weights it first saw for each member name, so it is exact until one changes + staleWeights bool + + calls atomic.Uint64 + members atomic.Pointer[[]*lb.Member] + weights sync.Map +} + +type brokenPrepared struct { + s *broken + members []*lb.Member + turns []*lb.Member +} + +var sink atomic.Pointer[[]byte] + +var stranger = lb.NewMember(lb.MemberOptions{Name: "stranger"}) + +func (s *broken) Name() string { return s.name } + +func (s *broken) Needs() lb.Needs { + if s.drifts && s.calls.Add(1)%2 == 0 { + return s.needs | lb.NeedKey + } + return s.needs +} + +func (s *broken) Prepare(snap *lb.Snapshot) lb.Prepared { + if s.stale { + s.members.CompareAndSwap(nil, &snap.Members) + return &brokenPrepared{s: s, members: *s.members.Load()} + } + p := &brokenPrepared{s: s, members: snap.Members} + if !s.exact && !s.staleWeights { + return p + } + var total int + for _, m := range snap.Members { + total += m.Weight() + } + if total > 1<<16 { + // too long a rotation to lay out turn by turn; such pools are not checked for exactness + return p + } + for _, m := range snap.Members { + w := m.Weight() + if s.staleWeights { + first, _ := s.weights.LoadOrStore(m.Name(), w) + w = first.(int) + } + // one entry per turn: a rotation over it gives each member exactly its weight + for range w { + p.turns = append(p.turns, m) + } + } + return p +} + +func (p *brokenPrepared) Select(f lb.Flow) *lb.Member { + s := p.s + switch { + case s.declines: + return nil + case s.hangs && len(p.members) > 1: + select {} + case s.foreign && len(p.members) > 1: + return stranger + case s.first: + return p.members[0] + case s.rotates: + return p.members[s.calls.Add(1)%uint64(len(p.members))] + case len(p.turns) > 0: + return p.turns[s.calls.Add(1)%uint64(len(p.turns))] + } + if s.allocates { + b := make([]byte, 64) + sink.Store(&b) + } + return p.members[f.Key%uint64(len(p.members))] +} + +func suiteOf(mutate func(*broken)) func() lb.Selector { + return func() lb.Selector { + s := &broken{name: "broken"} + if mutate != nil { + mutate(s) + } + return s + } +} + +// the recorder must not be the reason a check passes or fails +func TestRecorderAcceptsAConformingStrategy(t *testing.T) { + r := newRecorder() + run(r, suiteOf(nil), Options{}) + if failed := r.failedChecks(); len(failed) != 0 { + t.Fatalf("a conforming strategy failed %v: %v", failed, r.failures(failed[0])) + } + exact := newRecorder() + run(exact, suiteOf(func(s *broken) { s.exact = true }), Options{ExactWeights: true}) + if failed := exact.failedChecks(); len(failed) != 0 { + t.Fatalf("an exact strategy failed %v: %v", failed, exact.failures(failed[0])) + } + weighted := newRecorder() + runWeighted(weighted, suiteOf(func(s *broken) { s.exact = true }), WeightOptions{Tolerance: 0.001}) + if failed := weighted.failedChecks(); len(failed) != 0 { + t.Fatalf("an exact strategy failed %v: %v", failed, weighted.failures(failed[0])) + } + ran := false + r.Run("cleanup order", func(t suiteT) { + t.Cleanup(func() { ran = true }) + }) + if !ran { + t.Error("the recorder dropped a cleanup") + } +} + +// each broken strategy must be rejected, by the check that exists to catch it +func TestSuiteRejectsBrokenStrategies(t *testing.T) { + saved := terminationLimit + terminationLimit = 250 * time.Millisecond + t.Cleanup(func() { terminationLimit = saved }) + + for _, test := range []struct { + name string + break_ func(*broken) + opts Options + checks []string + says string + // hangs marks a strategy that only the check with a watchdog can be run against + hangs bool + }{ + { + name: "has no name", break_: func(s *broken) { s.name = "" }, + checks: []string{"identity"}, says: "must have a name", + }, + { + name: "needs vary", break_: func(s *broken) { s.drifts = true }, + checks: []string{"identity"}, says: "needs must not vary", + }, + { + name: "selects nothing", break_: func(s *broken) { s.declines = true }, + checks: []string{ + "no eligible member", "picks only eligible members", "every member is reachable", + "repick excludes the failed member", "in-flight accounting balances", + "adversarial snapshots terminate", "zero allocations", "concurrent picks and swaps", + }, + says: "was not picked", + }, + { + name: "selects outside the snapshot", break_: func(s *broken) { s.stale = true }, + checks: []string{"picks only eligible members", "repick excludes the failed member"}, + says: "not eligible", + }, + { + name: "selects a member of no pool", break_: func(s *broken) { s.foreign = true }, + checks: []string{ + "picks only eligible members", "adversarial snapshots terminate", "concurrent picks and swaps", + }, + says: "never in the pool", + }, + { + name: "starves members", break_: func(s *broken) { s.first = true }, + checks: []string{"every member is reachable"}, says: "was never picked", + }, + { + name: "allocates per pick", break_: func(s *broken) { s.allocates = true }, + checks: []string{"zero allocations"}, says: "allocates", + }, + { + name: "never returns", break_: func(s *broken) { s.hangs = true }, hangs: true, + checks: []string{"adversarial snapshots terminate"}, says: "did not terminate", + }, + { + name: "claims exact weights it does not keep", break_: func(s *broken) { s.rotates = true }, + opts: Options{ExactWeights: true}, + checks: []string{"exact weights", "exact weights after a live weight change"}, + says: "want exactly its weight", + }, + { + // exact over a stable pool, so only the live-change check can catch it + name: "misses a live weight change", break_: func(s *broken) { s.staleWeights = true }, + opts: Options{ExactWeights: true}, + checks: []string{"exact weights after a live weight change"}, + says: "want exactly its weight", + }, + } { + t.Run(test.name, func(t *testing.T) { + r := newRecorder() + if test.hangs { + r.Run("adversarial snapshots terminate", func(st suiteT) { testAdversarial(st, suiteOf(test.break_)) }) + } else { + run(r, suiteOf(test.break_), test.opts) + } + failed := r.failedChecks() + for _, check := range test.checks { + if !slices.Contains(failed, check) { + t.Errorf("%q passed a strategy that %s; checks that did fail: %v", check, test.name, failed) + } + } + var said bool + for _, check := range test.checks { + for _, msg := range r.failures(check) { + said = said || strings.Contains(msg, test.says) + } + } + if !said { + t.Errorf("no failure mentioned %q: %v", test.says, r.failed) + } + }) + } +} + +// a strategy that ignores weights must fail every form of the weighted suite +func TestWeightedSuiteRejectsAStrategyThatIgnoresWeights(t *testing.T) { + for name, o := range map[string]WeightOptions{ + "by load": {}, + "by key": {Keyed: true}, + "load only": {LoadOnly: true, Tolerance: 0.05}, + } { + t.Run(name, func(t *testing.T) { + r := newRecorder() + runWeighted(r, suiteOf(nil), o) + failed := r.failedChecks() + if len(failed) == 0 { + t.Fatal("the weighted suite passed a strategy that spreads flows evenly whatever the weights") + } + for _, check := range failed { + for _, msg := range r.failures(check) { + if !strings.Contains(msg, "picks, want") && !strings.Contains(msg, "flows behind") { + t.Errorf("%s: unexpected failure %q", check, msg) + } + } + } + if o.LoadOnly && slices.Contains(failed, "an idle pool is shared by weight") { + t.Error("a load-only strategy was held to the idle split") + } + }) + } + // and a strategy that lets one member fall behind is caught by its backlog + r := newRecorder() + runWeighted(r, suiteOf(func(s *broken) { s.first = true }), WeightOptions{LoadOnly: true}) + var backlog bool + for _, msg := range r.failures("sustained load is shared by weight") { + backlog = backlog || strings.Contains(msg, "flows behind") + } + if !backlog { + t.Errorf("a member buried under every flow was not reported as behind: %v", r.failed) + } +} + +// a strategy that selects nothing fails the weighted suite outright, in every form +func TestWeightedSuiteRejectsAStrategyThatDeclines(t *testing.T) { + for _, o := range []WeightOptions{{}, {Keyed: true}} { + r := newRecorder() + runWeighted(r, suiteOf(func(s *broken) { s.declines = true }), o) + failed := r.failedChecks() + want := 2 + if o.Keyed { + want = 1 + } + if len(failed) != want { + t.Errorf("keyed %v: %d checks failed, want %d: %v", o.Keyed, len(failed), want, r.failed) + } + for _, check := range failed { + if msgs := r.failures(check); len(msgs) != 1 || msgs[0] != "no pick" { + t.Errorf("%s: %v", check, msgs) + } + } + } + // the exact-weights checks stop at the first pick that is refused, too + r := newRecorder() + run(r, suiteOf(func(s *broken) { s.declines = true }), Options{ExactWeights: true}) + for _, check := range []string{"exact weights", "exact weights after a live weight change"} { + if msgs := r.failures(check); len(msgs) != 1 || msgs[0] != "no pick" { + t.Errorf("%s: %v", check, msgs) + } + } +} + +// Bench is part of the exported suite: it must run to completion over every shape it offers +func TestBenchRuns(t *testing.T) { + saved := flag.Lookup("test.benchtime").Value.String() + if err := flag.Set("test.benchtime", "1x"); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = flag.Set("test.benchtime", saved) }) + var ran atomic.Int64 + res := testing.Benchmark(func(b *testing.B) { + Bench(b, func() lb.Selector { + ran.Add(1) + return &broken{name: "benched"} + }) + }) + if ran.Load() != 8 { + t.Errorf("Bench built %d balancers, want one per size and shape", ran.Load()) + } + _ = res +} diff --git a/pkg/lb/lbtest/reporter.go b/pkg/lb/lbtest/reporter.go new file mode 100644 index 000000000..33db47afc --- /dev/null +++ b/pkg/lb/lbtest/reporter.go @@ -0,0 +1,43 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lbtest + +import "testing" + +// reporter is the part of testing.TB the suite's checks use; both *testing.T and *testing.B +// have it, and so does the recorder the suite's own tests drive it with +type reporter interface { + Helper() + Error(args ...any) + Errorf(format string, args ...any) + Fatal(args ...any) + Fatalf(format string, args ...any) + Cleanup(func()) +} + +// suiteT is a reporter that can run named checks. It stands between the suite and +// *testing.T so that the suite can be run against a strategy that is expected to fail it. +type suiteT interface { + reporter + Run(name string, check func(suiteT)) bool +} + +// realT adapts *testing.T, whose Run hands its function a *testing.T +type realT struct{ *testing.T } + +func (t realT) Run(name string, check func(suiteT)) bool { + return t.T.Run(name, func(t *testing.T) { check(realT{t}) }) +} diff --git a/pkg/lb/lbtest/weights.go b/pkg/lb/lbtest/weights.go new file mode 100644 index 000000000..ce588c5b7 --- /dev/null +++ b/pkg/lb/lbtest/weights.go @@ -0,0 +1,156 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lbtest + +import ( + "math" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// WeightOptions describes how a strategy is expected to honor member weights. +type WeightOptions struct { + // Tolerance is how far a member's share of picks may sit from its share of the pool's + // weight, as an absolute fraction; zero means 0.03. + Tolerance float64 + // Keyed marks a strategy whose weights divide the key space rather than the load: it is + // tested with many distinct keys instead of a load simulation. + Keyed bool + // LoadOnly marks a strategy that ranks members, so its weights bias the split only once + // work is in flight; idle, its best-ranked member rightly takes every flow. + LoadOnly bool +} + +// simulated weights and what one unit of weight can finish per tick +var ( + shareWeights = []int{1, 2, 5} + unitCapacity = 4 +) + +// RunWeighted holds a strategy to shares proportional to weight: over an idle pool, where +// nothing distinguishes the members but their weights, and under sustained load against +// members whose capacity is proportional to their weight. +func RunWeighted(t *testing.T, newSelector func() lb.Selector, o WeightOptions) { + runWeighted(realT{t}, newSelector, o) +} + +func runWeighted(t suiteT, newSelector func() lb.Selector, o WeightOptions) { + if o.Tolerance <= 0 { + o.Tolerance = 0.03 + } + if o.Keyed { + t.Run("key space is shared by weight", func(t suiteT) { testKeySpace(t, newSelector, o) }) + return + } + if !o.LoadOnly { + t.Run("an idle pool is shared by weight", func(t suiteT) { testIdleShares(t, newSelector, o) }) + } + t.Run("sustained load is shared by weight", func(t suiteT) { testLoadShares(t, newSelector, o) }) +} + +func assertShares(t reporter, members []*lb.Member, picks map[*lb.Member]int, tolerance float64) { + t.Helper() + var total, weight int + for _, m := range members { + total += picks[m] + weight += m.Weight() + } + for _, m := range members { + got := float64(picks[m]) / float64(total) + want := float64(m.Weight()) / float64(weight) + if math.Abs(got-want) > tolerance { + t.Errorf("%s (weight %d of %d) took %.3f of %d picks, want %.3f within %.3f", + m.Name(), m.Weight(), weight, got, total, want, tolerance) + } + } +} + +func testIdleShares(t suiteT, newSelector func() lb.Selector, o WeightOptions) { + members, _ := Members(shareWeights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + f := newFlows() + picks := make(map[*lb.Member]int) + for range 40000 { + pk, ok := b.Pick(f.next()) + if !ok { + t.Fatal("no pick") + } + picks[pk.Member()]++ + // every member answers alike, and nothing is ever left in flight + pk.Established(10 * time.Millisecond) + pk.Done(lb.OutcomeOK) + } + assertShares(t, members, picks, o.Tolerance) +} + +func testLoadShares(t suiteT, newSelector func() lb.Selector, o WeightOptions) { + members, _ := Members(shareWeights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + f := newFlows() + var capacity int + for _, w := range shareWeights { + capacity += w * unitCapacity + } + // arrivals run at nine tenths of what the pool can finish, so queues form and drain + arrivals := capacity * 9 / 10 + open := make(map[*lb.Member][]lb.Pick, len(members)) + picks := make(map[*lb.Member]int) + for range 3000 { + for _, m := range members { + done := min(len(open[m]), m.Weight()*unitCapacity) + for _, pk := range open[m][:done] { + pk.Done(lb.OutcomeOK) + } + open[m] = open[m][done:] + } + for range arrivals { + pk, ok := b.Pick(f.next()) + if !ok { + t.Fatal("no pick") + } + pk.Established(10 * time.Millisecond) + picks[pk.Member()]++ + open[pk.Member()] = append(open[pk.Member()], pk) + } + } + for _, m := range members { + for _, pk := range open[m] { + pk.Done(lb.OutcomeOK) + } + if backlog := len(open[m]); backlog > 4*m.Weight()*unitCapacity { + t.Errorf("%s was left %d flows behind, more than four ticks of its capacity", m.Name(), backlog) + } + } + assertShares(t, members, picks, o.Tolerance) +} + +func testKeySpace(t suiteT, newSelector func() lb.Selector, o WeightOptions) { + members, _ := Members(shareWeights...) + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: newPool(t, members)}) + f := newFlows() + picks := make(map[*lb.Member]int) + for range 200000 { + pk, ok := b.Pick(f.next()) + if !ok { + t.Fatal("no pick") + } + picks[pk.Member()]++ + pk.Done(lb.OutcomeOK) + } + assertShares(t, members, picks, o.Tolerance) +} diff --git a/pkg/lb/lc/lc.go b/pkg/lb/lc/lc.go new file mode 100644 index 000000000..980b43d2c --- /dev/null +++ b/pkg/lb/lc/lc.go @@ -0,0 +1,60 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package lc is the least connections strategy: each flow goes to the member with the fewest +// flows in flight for its weight. It suits long-lived work and small pools; it reads every +// member per pick. +package lc + +import ( + "sync/atomic" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Name is the strategy's name. +const Name = "least_connections" + +// New returns a least connections selector. +func New() lb.Selector { + return &selector{} +} + +type selector struct { + // rotation among members that tie, which is all of them while the pool is idle + pos atomic.Uint64 +} + +func (s *selector) Name() string { return Name } + +func (s *selector) Needs() lb.Needs { return lb.NeedInflight } + +func (s *selector) Prepare(snap *lb.Snapshot) lb.Prepared { + return &prepared{pos: &s.pos, members: snap.Members} +} + +type prepared struct { + pos *atomic.Uint64 + members []*lb.Member +} + +func (p *prepared) Select(lb.Flow) *lb.Member { + return lb.Least(p.members, p.pos, load) +} + +// load is the member's in-flight count per unit of weight; a weight is a capacity +func load(m *lb.Member) float64 { + return float64(m.Stats().Inflight()) / float64(m.Weight()) +} diff --git a/pkg/lb/lc/lc_test.go b/pkg/lb/lc/lc_test.go new file mode 100644 index 000000000..cd03f71c8 --- /dev/null +++ b/pkg/lb/lc/lc_test.go @@ -0,0 +1,96 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lc_test + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" + "github.com/trickstercache/trickster/v2/pkg/lb/lc" +) + +func TestConformance(t *testing.T) { + lbtest.Run(t, lc.New, lbtest.Options{}) +} + +func TestWeights(t *testing.T) { + lbtest.RunWeighted(t, lc.New, lbtest.WeightOptions{}) +} + +func BenchmarkSelect(b *testing.B) { + lbtest.Bench(b, lc.New) +} + +func newBalancer(t *testing.T, weights ...int) (*lb.Balancer, []*lb.Member) { + t.Helper() + members, _ := lbtest.Members(weights...) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return lb.NewBalancer(lc.New(), lb.BalancerOptions{Pool: p}), members +} + +func TestIdentity(t *testing.T) { + s := lc.New() + if s.Name() != lc.Name || s.Needs() != lb.NeedInflight { + t.Errorf("name %q needs %b", s.Name(), s.Needs()) + } +} + +// the member with the fewest flows in flight for its weight takes the next one +func TestPrefersTheLeastLoaded(t *testing.T) { + b, members := newBalancer(t, 1, 1, 1) + var held []lb.Pick + for range 9 { + pk, _ := b.Pick(lb.Flow{}) + held = append(held, pk) + } + for _, m := range members { + if got := m.Stats().Inflight(); got != 3 { + t.Fatalf("%s holds %d of 9 held flows, want 3", m.Name(), got) + } + } + // free one member entirely: it takes every new flow until it has caught up + var freed *lb.Member + for _, pk := range held { + if freed == nil { + freed = pk.Member() + } + if pk.Member() == freed { + pk.Done(lb.OutcomeOK) + } + } + for i := range 3 { + pk, _ := b.Pick(lb.Flow{}) + if pk.Member() != freed { + t.Fatalf("pick %d went to %s, not the idle member", i, pk.Member().Name()) + } + } +} + +// a weight is a capacity: a member three times the weight holds three times the flows +func TestWeightIsCapacity(t *testing.T) { + b, members := newBalancer(t, 1, 3) + for range 40 { + b.Pick(lb.Flow{}) + } + if light, heavy := members[0].Stats().Inflight(), members[1].Stats().Inflight(); light != 10 || heavy != 30 { + t.Errorf("held flows = %d and %d, want 10 and 30", light, heavy) + } +} diff --git a/pkg/lb/leaf.go b/pkg/lb/leaf.go new file mode 100644 index 000000000..83b581dde --- /dev/null +++ b/pkg/lb/leaf.go @@ -0,0 +1,198 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import "time" + +// MaxPickDepth is how many balancers a leaf pick may pass through: a pool, and the pools +// that its members may themselves be. A member that is balanced any deeper is refused. +const MaxPickDepth = 2 + +// LeafPick is a selection followed through members that are themselves balanced, down to a +// member that is not. It is returned by value and reports to every level it passed through. +type LeafPick struct { + levels [MaxPickDepth]Pick + depth int +} + +// FlowFunc supplies the flow for one level of a leaf pick. depth is how many levels are above +// it, level is the picker about to be asked, and via is the member whose payload offered it: +// nil for the outermost. It lets each level be keyed the way that level is configured; it must +// not retain its arguments. A retry may ask for the same depth more than once. +type FlowFunc func(depth int, level Picker, via *Member) Flow + +// Repicker is implemented by a Picker that can offer a flow its other members, as Balancer +// does, for a caller retrying work that some members could not take. +type Repicker interface { + // Alternatives returns the eligible members that skip does not report, the picker's choice + // for the flow first and the others after it. It commits the flow to none of them. + Alternatives(f Flow, skip func(*Member) bool) []*Member + // Commit commits a flow to a member that Alternatives returned, as Pick would have. + Commit(*Member) (Pick, bool) +} + +// PickLeaf commits a flow to a member that is not itself balanced, asking a member whose +// payload is a PickerProvider to pick again among its own. It returns false when any level has +// no eligible member or the members nest deeper than MaxPickDepth; the levels already +// committed are then released without prejudice to their members, whose share is refused +// rather than passed to a sibling. +func PickLeaf(p Picker, f Flow) (LeafPick, bool) { + return pickLeaf(p, func(int, Picker, *Member) Flow { return f }) +} + +// PickLeafFunc is PickLeaf with each level's flow supplied by flow. +func PickLeafFunc(p Picker, flow FlowFunc) (LeafPick, bool) { + return pickLeaf(p, flow) +} + +func pickLeaf(p Picker, flow FlowFunc) (LeafPick, bool) { + var lp LeafPick + var via *Member + for lp.depth < MaxPickDepth && p != nil { + pk, ok := p.Pick(flow(lp.depth, p, via)) + if !ok { + break + } + lp.levels[lp.depth] = pk + lp.depth++ + // a payload that offers no picker is a leaf, whether or not it could have offered one + via, p = pk.Member(), pickerOf(pk.Member()) + if p == nil { + return lp, true + } + } + lp.Done(OutcomeCanceled) + return LeafPick{}, false +} + +func pickerOf(m *Member) Picker { + if pp, ok := m.Value.(PickerProvider); ok { + return pp.Picker() + } + return nil +} + +// RepickLeafFunc is PickLeafFunc for a caller retrying work that the leaf members in failed +// could not take. No level that can offer alternatives commits to one of them, and a member +// whose own pool has no other leaf left is passed over for its siblings, so a leaf that can be +// reached is never given up on. It is not a selection path. +// +// The search is one traversal: each level is asked for its alternatives once, in its own +// order of preference, and each member is visited at most once, so the work is linear in the +// members searched however many of them turn out to have nothing to offer. +func RepickLeafFunc(p Picker, flow FlowFunc, failed ...*Member) (LeafPick, bool) { + tried := make(map[*Member]struct{}, len(failed)) + for _, m := range failed { + tried[m] = struct{}{} + } + var lp LeafPick + if !lp.repick(p, nil, flow, func(m *Member) bool { _, ok := tried[m]; return ok }) { + return LeafPick{}, false + } + return lp, true +} + +// repick extends lp from p down to a leaf, backing out of any member that leads to none. On +// failure lp is left as it was found. +func (lp *LeafPick) repick(p Picker, via *Member, flow FlowFunc, skip func(*Member) bool) bool { + if lp.depth >= MaxPickDepth { + return false + } + rp, can := p.(Repicker) + if !can { + // a picker with no alternatives to offer is asked for an ordinary pick + pk, ok := p.Pick(flow(lp.depth, p, via)) + return ok && lp.extend(pk, flow, skip) + } + for _, m := range rp.Alternatives(flow(lp.depth, p, via), skip) { + if pk, ok := rp.Commit(m); ok && lp.extend(pk, flow, skip) { + return true + } + } + return false +} + +// extend adds a committed pick to lp and follows it to a leaf, releasing it again, without +// prejudice to its member, when it leads to none +func (lp *LeafPick) extend(pk Pick, flow FlowFunc, skip func(*Member) bool) bool { + lp.levels[lp.depth] = pk + lp.depth++ + next := pickerOf(pk.Member()) + if next == nil || lp.repick(next, pk.Member(), flow, skip) { + return true + } + lp.depth-- + lp.levels[lp.depth] = Pick{} + pk.Done(OutcomeCanceled) + return false +} + +// Member returns the selected leaf member, or nil for the zero LeafPick. +func (lp LeafPick) Member() *Member { + if lp.depth == 0 { + return nil + } + return lp.levels[lp.depth-1].Member() +} + +// Depth returns how many balancers the pick passed through. +func (lp LeafPick) Depth() int { + return lp.depth +} + +// Level returns the pick made at one of the balancers passed through, outermost first, for a +// caller that reports to each level in its own terms; the zero Pick when out of range. +func (lp LeafPick) Level(i int) Pick { + if i < 0 || i >= lp.depth { + return Pick{} + } + return lp.levels[i] +} + +// Established reports to every level that the leaf was reached and how long that took. +func (lp LeafPick) Established(d time.Duration) { + for i := range lp.depth { + lp.levels[i].Established(d) + } +} + +// Reached reports that the leaf member was reached. Only a leaf is ever blamed for a failed +// connect, so only the leaf has a run of them to end. +func (lp LeafPick) Reached() { + if lp.depth > 0 { + lp.levels[lp.depth-1].Reached() + } +} + +// FirstByte reports the first sign of a response to every level. +func (lp LeafPick) FirstByte() { + for i := range lp.depth { + lp.levels[i].FirstByte() + } +} + +// Done reports the flow's end to every level. It must be called exactly once per LeafPick. +// A failure belongs to the leaf that failed: the levels above it, whose member is a whole +// pool, are released without prejudice, so one bad member does not condemn its pool. +func (lp LeafPick) Done(o Outcome) { + for i := range lp.depth { + if i < lp.depth-1 && o != OutcomeOK { + lp.levels[i].Done(OutcomeCanceled) + continue + } + lp.levels[i].Done(o) + } +} diff --git a/pkg/lb/leaf_test.go b/pkg/lb/leaf_test.go new file mode 100644 index 000000000..b9909c72d --- /dev/null +++ b/pkg/lb/leaf_test.go @@ -0,0 +1,570 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb_test + +import ( + "slices" + "strconv" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/hrw" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" + "github.com/trickstercache/trickster/v2/pkg/lb/lc" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" +) + +// nested is a member payload that is itself balanced +type nested struct{ picker lb.Picker } + +func (n nested) Picker() lb.Picker { return n.picker } + +func poolOf(t *testing.T, s lb.Selector, members ...*lb.Member) *lb.Balancer { + t.Helper() + p, err := lb.NewPool(members, 0) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return lb.NewBalancer(s, lb.BalancerOptions{Pool: p}) +} + +func leaf(name string, weight int) *lb.Member { + return lb.NewMember(lb.MemberOptions{Name: name, Weight: weight, Value: name}) +} + +func TestPickLeafOfAFlatPool(t *testing.T) { + a := leaf("a", 1) + b := poolOf(t, lc.New(), a) + lp, ok := lb.PickLeaf(b, lb.Flow{}) + if !ok || lp.Member() != a || lp.Depth() != 1 { + t.Fatalf("leaf = %v at depth %d, %v", lp.Member(), lp.Depth(), ok) + } + if a.Stats().Inflight() != 1 { + t.Errorf("in flight = %d", a.Stats().Inflight()) + } + lp.Established(time.Millisecond) + lp.FirstByte() + lp.Done(lb.OutcomeOK) + if a.Stats().Inflight() != 0 { + t.Errorf("in flight after Done = %d", a.Stats().Inflight()) + } + var none lb.LeafPick + none.Done(lb.OutcomeOK) + if none.Member() != nil || none.Depth() != 0 { + t.Error("the zero leaf pick has a member") + } + if _, ok := lb.PickLeaf(nil, lb.Flow{}); ok { + t.Error("picked a leaf from no picker") + } + if _, ok := lb.PickLeaf(lb.NewBalancer(rr.New()), lb.Flow{}); ok { + t.Error("picked a leaf from a balancer with no pool") + } +} + +// each level apportions by its own weights, and reports reach every level passed through +func TestPickLeafFollowsNestedPools(t *testing.T) { + a, b, c := leaf("a", 2), leaf("b", 1), leaf("c", 1) + inner1 := poolOf(t, lc.New(), a, b) + inner2 := poolOf(t, lc.New(), c) + m1 := lb.NewMember(lb.MemberOptions{Name: "inner1", Weight: 3, Value: nested{inner1}}) + m2 := lb.NewMember(lb.MemberOptions{Name: "inner2", Weight: 1, Value: nested{inner2}}) + outer := poolOf(t, lc.New(), m1, m2) + + var held []lb.LeafPick + for range 12 { + lp, ok := lb.PickLeaf(outer, lb.Flow{}) + if !ok || lp.Depth() != 2 { + t.Fatalf("pick at depth %d, %v", lp.Depth(), ok) + } + held = append(held, lp) + } + // 12 held flows: 9 and 3 at the outer level; 6 and 3 within inner1 + for m, want := range map[*lb.Member]int64{m1: 9, m2: 3, a: 6, b: 3, c: 3} { + if got := m.Stats().Inflight(); got != want { + t.Errorf("%s holds %d flows, want %d", m.Name(), got, want) + } + } + for _, lp := range held { + lp.Done(lb.OutcomeOK) + } + for _, m := range []*lb.Member{m1, m2, a, b, c} { + if got := m.Stats().Inflight(); got != 0 { + t.Errorf("%s holds %d flows after every pick was done", m.Name(), got) + } + } +} + +// an inner pool with no eligible member refuses its share; the outer level is released, and +// the share is not handed to a sibling +func TestPickLeafRefusesAnEmptyInnerPoolsShare(t *testing.T) { + members, healths := lbtest.Members(1) + healths[0].Set(-1) + down, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + defer down.Stop() + empty := lb.NewMember(lb.MemberOptions{Name: "empty", Value: nested{lb.NewBalancer(lc.New(), lb.BalancerOptions{Pool: down})}}) + live := lb.NewMember(lb.MemberOptions{Name: "live", Value: nested{poolOf(t, lc.New(), leaf("c", 1))}}) + outer := poolOf(t, rr.NewAt(0), empty, live) + var served, refused int + for range 8 { + lp, ok := lb.PickLeaf(outer, lb.Flow{}) + if !ok { + refused++ + continue + } + served++ + lp.Done(lb.OutcomeOK) + } + if served != 4 || refused != 4 { + t.Errorf("served %d, refused %d; want the empty inner pool's half refused", served, refused) + } + outerLC := poolOf(t, lc.New(), lb.NewMember(lb.MemberOptions{Name: "empty2", Value: nested{ + lb.NewBalancer(lc.New(), lb.BalancerOptions{Pool: down}), + }})) + if _, ok := lb.PickLeaf(outerLC, lb.Flow{}); ok { + t.Fatal("picked through an empty inner pool") + } + if got := outerLC.Pool().Configured()[0].Stats().Inflight(); got != 0 { + t.Errorf("the refused pick left %d in flight at the outer level", got) + } +} + +// members nested deeper than the bound are refused, and nothing is left in flight +func TestPickLeafDepthBound(t *testing.T) { + l2 := poolOf(t, lc.New(), leaf("z", 1)) + mid := lb.NewMember(lb.MemberOptions{Name: "l1-member", Value: nested{l2}}) + l1 := poolOf(t, lc.New(), mid) + top := lb.NewMember(lb.MemberOptions{Name: "l0-member", Value: nested{l1}}) + l0 := poolOf(t, lc.New(), top) + if _, ok := lb.PickLeaf(l0, lb.Flow{}); ok { + t.Fatal("a pool three deep was followed") + } + if top.Stats().Inflight() != 0 || mid.Stats().Inflight() != 0 { + t.Errorf("the refused pick left %d and %d in flight", top.Stats().Inflight(), mid.Stats().Inflight()) + } + // a payload that could offer a picker but has none is a leaf + plain := lb.NewMember(lb.MemberOptions{Name: "plain", Value: nested{}}) + if lp, ok := lb.PickLeaf(poolOf(t, lc.New(), plain), lb.Flow{}); !ok || lp.Member() != plain { + t.Error("a member with no picker of its own was not taken as the leaf") + } + if allocs := testing.AllocsPerRun(200, func() { + lp, _ := lb.PickLeaf(l1, lb.Flow{}) + lp.Done(lb.OutcomeOK) + }); allocs != 0 { + t.Errorf("a nested leaf pick allocates %v", allocs) + } +} + +// each level is asked with its own flow, and told which member led to it +func TestPickLeafFuncKeysEachLevel(t *testing.T) { + a, b := leaf("a", 1), leaf("b", 1) + inner := poolOf(t, &keyed{}, a, b) + via := lb.NewMember(lb.MemberOptions{Name: "inner", Value: nested{inner}}) + outer := poolOf(t, &keyed{}, via) + var asked []*lb.Member + lp, ok := lb.PickLeafFunc(outer, func(_ int, level lb.Picker, from *lb.Member) lb.Flow { + asked = append(asked, from) + if level == lb.Picker(inner) { + return lb.Flow{Key: 1, HasKey: true} + } + return lb.Flow{} + }) + if !ok || lp.Member() != b { + t.Fatalf("leaf = %v; the inner level's key selects b", lp.Member()) + } + if len(asked) != 2 || asked[0] != nil || asked[1] != via { + t.Errorf("levels were asked via %v", asked) + } + if lp.Level(0).Member() != via || lp.Level(1).Member() != b || lp.Level(2).Member() != nil || lp.Level(-1).Member() != nil { + t.Error("Level does not return each level's pick") + } + lp.Done(lb.OutcomeOK) +} + +// a failure belongs to the leaf: the pool above it is released, not blamed +func TestLeafPickBlamesOnlyTheLeaf(t *testing.T) { + a := leaf("a", 1) + inner := poolOf(t, lc.New(), a) + via := lb.NewMember(lb.MemberOptions{Name: "inner", Value: nested{inner}}) + outer := poolOf(t, lc.New(), via) + for _, o := range []lb.Outcome{lb.OutcomeConnectFailed, lb.OutcomeFailed} { + lp, _ := lb.PickLeaf(outer, lb.Flow{}) + lp.Done(o) + } + if a.Stats().Failures() != 2 || via.Stats().Failures() != 0 { + t.Errorf("failures: leaf %d, the pool above it %d", a.Stats().Failures(), via.Stats().Failures()) + } + if a.Stats().Inflight() != 0 || via.Stats().Inflight() != 0 { + t.Error("a failed leaf pick left something in flight") + } + lp, _ := lb.PickLeaf(outer, lb.Flow{}) + lp.Done(lb.OutcomeOK) + if a.Stats().Failures() != 0 { + t.Error("a success did not clear the leaf's failures") + } +} + +// a retry avoids the member that failed, at whichever level holds it +func TestRepickLeafAvoidsTheFailedMember(t *testing.T) { + a, b := leaf("a", 1), leaf("b", 1) + inner := poolOf(t, rr.NewAt(0), a, b) + outer := poolOf(t, rr.NewAt(0), lb.NewMember(lb.MemberOptions{Name: "inner", Value: nested{inner}})) + flow := func(int, lb.Picker, *lb.Member) lb.Flow { return lb.Flow{} } + for range 10 { + lp, ok := lb.RepickLeafFunc(outer, flow, a) + if !ok || lp.Member() != b { + t.Fatalf("retry chose %v", lp.Member()) + } + lp.Done(lb.OutcomeOK) + } + only := poolOf(t, rr.NewAt(0), a) + if _, ok := lb.RepickLeafFunc(only, flow, a); ok { + t.Error("retried onto the only member, which had failed") + } + // a picker that cannot avoid a member is asked for an ordinary pick + if lp, ok := lb.RepickLeafFunc(plainPicker{only}, flow, b); !ok || lp.Member() != a { + t.Errorf("plain picker retry = %v, %v", lp.Member(), ok) + } +} + +// plainPicker hides a balancer's Repick +type plainPicker struct{ b *lb.Balancer } + +func (p plainPicker) Needs() lb.Needs { return p.b.Needs() } + +func (p plainPicker) Pick(f lb.Flow) (lb.Pick, bool) { return p.b.Pick(f) } + +// a retry avoids every member the flow has failed on, and moves on from a pool that has no +// other member left to the pools beside it +func TestRepickLeafAvoidsEveryFailedMember(t *testing.T) { + a, b, c, d := leaf("a", 1), leaf("b", 1), leaf("c", 1), leaf("d", 1) + flow := func(int, lb.Picker, *lb.Member) lb.Flow { return lb.Flow{} } + flat := poolOf(t, lc.New(), a, b, c) + failed := []*lb.Member{a, b} + for range 10 { + lp, ok := lb.RepickLeafFunc(flat, flow, failed...) + if !ok || lp.Member() != c { + t.Fatalf("retry chose %v", lp.Member()) + } + lp.Done(lb.OutcomeOK) + } + if _, ok := lb.RepickLeafFunc(flat, flow, a, b, c); ok { + t.Error("retried although every member had failed") + } + if len(failed) != 2 || cap(failed) != 2 { + t.Error("the caller's list of failed members was modified") + } + + left := lb.NewMember(lb.MemberOptions{Name: "left", Value: nested{poolOf(t, lc.New(), a, b)}}) + right := lb.NewMember(lb.MemberOptions{Name: "right", Value: nested{poolOf(t, lc.New(), c, d)}}) + outer := poolOf(t, rr.NewAt(0), left, right) + for range 10 { + lp, ok := lb.RepickLeafFunc(outer, flow, a, b, c) + if !ok || lp.Member() != d { + t.Fatalf("nested retry chose %v, %v", lp.Member(), ok) + } + lp.Done(lb.OutcomeOK) + } + if _, ok := lb.RepickLeafFunc(outer, flow, a, b, c, d); ok { + t.Error("nested retry found a member although every leaf had failed") + } + for _, m := range []*lb.Member{a, b, c, d, left, right} { + if m.Stats().Inflight() != 0 { + t.Errorf("%s left %d in flight", m.Name(), m.Stats().Inflight()) + } + } +} + +// singletons builds n pools of one leaf each, as the members of an outer pool +func singletons(t *testing.T, n int) (pools, leaves []*lb.Member) { + t.Helper() + for i := range n { + m := leaf("leaf-"+strconv.Itoa(i), 1) + leaves = append(leaves, m) + pools = append(pools, lb.NewMember(lb.MemberOptions{ + Name: "pool-" + strconv.Itoa(i), Value: nested{poolOf(t, rr.NewAt(0), m)}, + })) + } + return pools, leaves +} + +// however many pools a retry finds spent on its way, it reaches a leaf that is still untried. +// An affinity strategy is the hard case: every retry starts at the same pool, and meets the +// spent ones in the same order, one per pass. +func TestRepickLeafReachesTheLastUntriedLeaf(t *testing.T) { + const n = 24 + pools, leaves := singletons(t, n) + outer := poolOf(t, hrw.New(), pools...) + flow := func(_ int, p lb.Picker, _ *lb.Member) lb.Flow { + if p.Needs().Has(lb.NeedKey) { + return lb.Flow{Key: lb.HashString("one-client"), HasKey: true} + } + return lb.Flow{} + } + // fail every leaf in the order the one client is offered them, down to the last + var failed []*lb.Member + for range n { + lp, ok := lb.RepickLeafFunc(outer, flow, failed...) + if !ok { + t.Fatalf("gave up with %d of %d leaves still untried", n-len(failed), n) + } + if slices.Contains(failed, lp.Member()) { + t.Fatalf("retried onto %s, which had failed", lp.Member().Name()) + } + failed = append(failed, lp.Member()) + lp.Done(lb.OutcomeConnectFailed) + } + if _, ok := lb.RepickLeafFunc(outer, flow, failed...); ok { + t.Error("found a member although every leaf had failed") + } + for _, m := range append(pools, leaves...) { + if m.Stats().Inflight() != 0 { + t.Errorf("%s left %d in flight", m.Name(), m.Stats().Inflight()) + } + } +} + +// a picker that cannot avoid members offers a spent pool again, which ends the search +func TestRepickLeafEndsWhenAPickerCannotAvoid(t *testing.T) { + pools, leaves := singletons(t, 3) + flow := func(int, lb.Picker, *lb.Member) lb.Flow { return lb.Flow{} } + outer := plainPicker{poolOf(t, rr.NewAt(0), pools...)} + if _, ok := lb.RepickLeafFunc(outer, flow, leaves...); ok { + t.Error("found a member although every leaf had failed") + } + // and one of its pools that still has a leaf to offer is found on the way round + if lp, ok := lb.RepickLeafFunc(outer, flow, leaves[0], leaves[1]); !ok || lp.Member() != leaves[2] { + t.Errorf("retry = %v, %v", lp.Member(), ok) + } else { + lp.Done(lb.OutcomeOK) + } +} + +// reaching the leaf ends its run of failed connects, under whichever pools it was picked through +func TestLeafPickReachedEndsTheLeafsRun(t *testing.T) { + a, b := leaf("a", 1), leaf("b", 1) + p, err := lb.NewPool([]*lb.Member{a, b}, 0) + if err != nil { + t.Fatal(err) + } + defer p.Stop() + inner := lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: p, Ejection: lb.EjectionOptions{Failures: 5}}) + outer := poolOf(t, rr.NewAt(0), lb.NewMember(lb.MemberOptions{Name: "inner", Value: nested{inner}})) + var failedOn *lb.Member + for range 2 { + lp, _ := lb.PickLeaf(outer, lb.Flow{}) + if failedOn == nil { + failedOn = lp.Member() + } + if lp.Member() == failedOn { + lp.Done(lb.OutcomeConnectFailed) + continue + } + lp.Done(lb.OutcomeCanceled) + } + if failedOn.Stats().ConnectFailures() != 1 { + t.Fatalf("connect failures = %d", failedOn.Stats().ConnectFailures()) + } + for range 2 { + lp, _ := lb.PickLeaf(outer, lb.Flow{}) + lp.Reached() + lp.Done(lb.OutcomeOK) + } + if failedOn.Stats().ConnectFailures() != 0 { + t.Error("reaching the leaf did not end its run of failed connects") + } +} + +// emptyPools builds n members that are each a pool with no member to offer, which an outer +// pool that is not told of their health goes on selecting +func emptyPools(t testing.TB, n int) []*lb.Member { + t.Helper() + out := make([]*lb.Member, n) + for i := range out { + p, err := lb.NewPool(nil, 0) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + b := lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: p}) + out[i] = lb.NewMember(lb.MemberOptions{Name: "empty-" + strconv.Itoa(i), Value: nested{b}}) + } + return out +} + +// lastSelector always chooses the last member, so a retry over pools that are all empty but +// the last meets every empty one before it, and counts how often it is prepared +type lastSelector struct{ prepares *int } + +type lastOf []*lb.Member + +func (s lastSelector) Name() string { return "last" } + +func (s lastSelector) Needs() lb.Needs { return 0 } + +func (s lastSelector) Prepare(snap *lb.Snapshot) lb.Prepared { + *s.prepares++ + return lastOf(snap.Members) +} + +func (l lastOf) Select(lb.Flow) *lb.Member { return l[len(l)-1] } + +// worstCase is an outer pool whose strategy prefers a member with nothing to offer, followed +// in its order of preference by every other empty member, and only then by the one live pool +func worstCase(t testing.TB, empties int) (outer *lb.Balancer, live *lb.Member, prepares *int) { + t.Helper() + live = leaf("live", 1) + lp, err := lb.NewPool([]*lb.Member{live}, 0) + if err != nil { + t.Fatal(err) + } + t.Cleanup(lp.Stop) + livePool := lb.NewMember(lb.MemberOptions{ + Name: "live-pool", Value: nested{lb.NewBalancer(rr.NewAt(0), lb.BalancerOptions{Pool: lp})}, + }) + // the choice is the last member; the search then runs on from the first, so the live pool, + // placed just ahead of the last, is the final member it comes to + members := emptyPools(t, empties) + members = append(members[:empties-1:empties-1], livePool, members[empties-1]) + op, err := lb.NewPool(members, 0) + if err != nil { + t.Fatal(err) + } + t.Cleanup(op.Stop) + prepares = new(int) + return lb.NewBalancer(lastSelector{prepares}, lb.BalancerOptions{Pool: op}), live, prepares +} + +// a retry that must pass many members with nothing to offer does work in proportion to their +// number: each level is asked once and prepared once, however many of its members are spent +func TestRepickLeafWorkIsLinearInSpentMembers(t *testing.T) { + for _, empties := range []int{10, 1000, 4000} { + outer, live, prepares := worstCase(t, empties) + var asked, skipped int + flow := func(int, lb.Picker, *lb.Member) lb.Flow { asked++; return lb.Flow{} } + failed := leaf("failed-elsewhere", 1) + lp, ok := lb.RepickLeafFunc(countingSkips{outer, &skipped}, flow, failed) + if !ok || lp.Member() != live || lp.Depth() != 2 { + t.Fatalf("%d empties: retry = %v, %v", empties, lp.Member(), ok) + } + lp.Done(lb.OutcomeOK) + // the outer level once, and once for each member looked into: the empties and the live pool + if want := empties + 2; asked != want { + t.Errorf("%d empties: %d levels asked for a flow, want %d", empties, asked, want) + } + if *prepares != 1 { + t.Errorf("%d empties: the outer strategy was prepared %d times, want once", empties, *prepares) + } + // every outer member is tested against the failed set once, never once per spent member + if want := empties + 1; skipped != want { + t.Errorf("%d empties: %d exclusion checks at the outer level, want %d", empties, skipped, want) + } + for _, m := range outer.Pool().Configured() { + if m.Stats().Inflight() != 0 { + t.Fatalf("%s left %d in flight", m.Name(), m.Stats().Inflight()) + } + } + } +} + +// countingSkips counts the exclusion checks made on behalf of one level +type countingSkips struct { + *lb.Balancer + n *int +} + +func (c countingSkips) Alternatives(f lb.Flow, skip func(*lb.Member) bool) []*lb.Member { + return c.Balancer.Alternatives(f, func(m *lb.Member) bool { *c.n++; return skip(m) }) +} + +func BenchmarkRepickLeafPastEmptyPools(b *testing.B) { + for _, empties := range []int{100, 1000, 10000} { + b.Run(strconv.Itoa(empties), func(b *testing.B) { + outer, _, _ := worstCase(b, empties) + flow := func(int, lb.Picker, *lb.Member) lb.Flow { return lb.Flow{} } + failed := leaf("failed-elsewhere", 1) + b.ReportAllocs() + for b.Loop() { + lp, ok := lb.RepickLeafFunc(outer, flow, failed) + if !ok { + b.Fatal("no leaf") + } + lp.Done(lb.OutcomeOK) + } + }) + } +} + +// a retry refuses what a first pick refuses: members nested deeper than a leaf pick follows, +// and a level whose strategy declines to choose +func TestRepickLeafRefusals(t *testing.T) { + flow := func(int, lb.Picker, *lb.Member) lb.Flow { return lb.Flow{} } + a := leaf("a", 1) + inner := lb.NewMember(lb.MemberOptions{Name: "inner", Value: nested{poolOf(t, rr.NewAt(0), a)}}) + middle := lb.NewMember(lb.MemberOptions{Name: "middle", Value: nested{poolOf(t, rr.NewAt(0), inner)}}) + outer := poolOf(t, rr.NewAt(0), middle) + if _, ok := lb.RepickLeafFunc(outer, flow, leaf("other", 1)); ok { + t.Error("a retry followed members nested deeper than a first pick may") + } + for _, m := range []*lb.Member{a, inner, middle} { + if m.Stats().Inflight() != 0 { + t.Errorf("%s left %d in flight", m.Name(), m.Stats().Inflight()) + } + } + declines := poolOf(t, &countingSelector{refuse: true}, leaf("b", 1), leaf("c", 1)) + if alts := declines.Alternatives(lb.Flow{}, nil); alts != nil { + t.Errorf("alternatives of a strategy that declines = %v", alts) + } + if _, ok := lb.RepickLeafFunc(declines, flow); ok { + t.Error("a retry chose for a strategy that declined to") + } + if alts := lb.NewBalancer(rr.NewAt(0)).Alternatives(lb.Flow{}, nil); alts != nil { + t.Error("a balancer with no pool has alternatives") + } +} + +// strayOf always chooses one member, whether or not it was offered it +type straySelector struct{ stray *lb.Member } + +type strayOf struct{ stray *lb.Member } + +func (s straySelector) Name() string { return "stray" } + +func (s straySelector) Needs() lb.Needs { return 0 } + +func (s straySelector) Prepare(*lb.Snapshot) lb.Prepared { return strayOf(s) } + +func (s strayOf) Select(lb.Flow) *lb.Member { return s.stray } + +// a strategy's choice comes first among the alternatives exactly as Pick would commit to it, +// so one that chooses outside what it was offered is no better hidden on a retry +func TestAlternativesHonorTheStrategysChoice(t *testing.T) { + a, b, stray := leaf("a", 1), leaf("b", 1), leaf("stray", 1) + alts := poolOf(t, straySelector{stray}, a, b).Alternatives(lb.Flow{}, nil) + if !slices.Equal(alts, []*lb.Member{stray, a, b}) { + t.Errorf("alternatives = %v", alts) + } + // the choice leads, and the others follow in pool order from it + c, d := leaf("c", 1), leaf("d", 1) + alts = poolOf(t, straySelector{c}, a, b, c, d).Alternatives(lb.Flow{}, func(m *lb.Member) bool { return m == b }) + if !slices.Equal(alts, []*lb.Member{c, d, a}) { + t.Errorf("alternatives = %v", alts) + } +} diff --git a/pkg/lb/least.go b/pkg/lb/least.go new file mode 100644 index 000000000..d11fc759b --- /dev/null +++ b/pkg/lb/least.go @@ -0,0 +1,60 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import "sync/atomic" + +// Least returns the member with the lowest score. Members that tie for it, which is every +// member of an idle pool, share the picks by weighted rotation over pos rather than the first +// of them taking every one. It neither locks nor allocates, and visits each member at most +// twice. score must be cheap and must not be negative; it may read live stats. +func Least(members []*Member, pos *atomic.Uint64, score func(*Member) float64) *Member { + if len(members) == 0 { + return nil + } + best := 0 + low := score(members[0]) + tiedWeight := uint64(members[0].weight) // #nosec G115 -- a member's weight is at least 1 + ties := 1 + for i := 1; i < len(members); i++ { + s := score(members[i]) + switch { + case s < low: + best, low, ties = i, s, 1 + tiedWeight = uint64(members[i].weight) // #nosec G115 -- at least 1 + case s == low: + ties++ + tiedWeight += uint64(members[i].weight) // #nosec G115 -- at least 1 + } + } + if ties == 1 { + return members[best] + } + // scores are live, so the tied set may have moved since the first pass: the walk is + // bounded by the slice and falls back to the first member found at the low score + k := pos.Add(1) % tiedWeight + for i := best; i < len(members); i++ { + if score(members[i]) != low { + continue + } + w := uint64(members[i].weight) // #nosec G115 -- at least 1 + if k < w { + return members[i] + } + k -= w + } + return members[best] +} diff --git a/pkg/lb/lt/lt.go b/pkg/lb/lt/lt.go new file mode 100644 index 000000000..26ce0fb4a --- /dev/null +++ b/pkg/lb/lt/lt.go @@ -0,0 +1,116 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package lt is the least time strategy: each flow goes to the member with the lowest +// latency average, scaled by its flows in flight and its weight. It reads every member per pick. +// +// A weight is a bias, not a share: it divides the member's score, so under load the split +// follows the weights, while an idle pool gives every flow to its best-scoring member. +// +// An average fades while its member is passed over, so a member ranked behind its peers on +// an old sample, or on a failure's penalty, is tried again rather than left there for good. +package lt + +import ( + "sync/atomic" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Name is the strategy's name. +const Name = "least_time" + +// Options tune the strategy. Zero values take the core's defaults. +type Options struct { + // Decay is the time constant with which a member's latency average yields to lower + // samples, and with which a failure's penalty fades while the member gets no work. + Decay time.Duration + // Penalty is the least latency a failed flow is recorded as. + Penalty time.Duration +} + +// New returns a least time selector. +func New(o Options) lb.Selector { + s := &selector{opts: o} + if s.opts.Decay <= 0 { + s.opts.Decay = lb.DefaultLatencyDecay + } + if s.opts.Penalty <= 0 { + s.opts.Penalty = lb.DefaultLatencyPenalty + } + return s +} + +type selector struct { + opts Options + // rotation among members that tie, which is all of them until one has a sample + pos atomic.Uint64 +} + +func (s *selector) Name() string { return Name } + +func (s *selector) Needs() lb.Needs { return lb.NeedInflight | lb.NeedLatency } + +// Latency tells the balancer how to average the samples it records for this strategy. +func (s *selector) Latency() lb.LatencyOptions { + return lb.LatencyOptions{Decay: s.opts.Decay, Penalty: s.opts.Penalty} +} + +func (s *selector) Prepare(snap *lb.Snapshot) lb.Prepared { + return &prepared{pos: &s.pos, members: snap.Members, decay: s.opts.Decay} +} + +type prepared struct { + pos *atomic.Uint64 + members []*lb.Member + decay time.Duration +} + +func (p *prepared) Select(lb.Flow) *lb.Member { + // a member with no sample yet is scored as its fastest healthy peer: it ties with the + // best rather than beating it, so it shares that member's flows until its own first + // sample ranks it, however idle the pool is. Its flows in flight bound the burst. + now, decay := time.Now(), p.decay + var best, bestAny float64 + for _, m := range p.members { + st := m.Stats() + l := float64(st.Faded(now, decay)) + if l <= 0 { + continue + } + if bestAny == 0 || l < bestAny { + bestAny = l + } + if st.Failures() == 0 && (best == 0 || l < best) { + best = l + } + } + cold := 1.0 + switch { + case best > 0: + cold = best + case bestAny > 0: + cold = bestAny + } + return lb.Least(p.members, p.pos, func(m *lb.Member) float64 { + st := m.Stats() + latency := float64(st.Faded(now, decay)) + if latency <= 0 { + latency = cold + } + return latency * float64(st.Inflight()+1) / float64(m.Weight()) + }) +} diff --git a/pkg/lb/lt/lt_test.go b/pkg/lb/lt/lt_test.go new file mode 100644 index 000000000..84e1b3324 --- /dev/null +++ b/pkg/lb/lt/lt_test.go @@ -0,0 +1,235 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lt_test + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" + "github.com/trickstercache/trickster/v2/pkg/lb/lt" +) + +func newSelector() lb.Selector { return lt.New(lt.Options{}) } + +func TestConformance(t *testing.T) { + lbtest.Run(t, newSelector, lbtest.Options{}) +} + +func TestWeights(t *testing.T) { + lbtest.RunWeighted(t, newSelector, lbtest.WeightOptions{LoadOnly: true}) +} + +func BenchmarkSelect(b *testing.B) { + lbtest.Bench(b, newSelector) +} + +func newBalancer(t *testing.T, o lt.Options, weights ...int) (*lb.Balancer, []*lb.Member) { + t.Helper() + members, _ := lbtest.Members(weights...) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return lb.NewBalancer(lt.New(o), lb.BalancerOptions{Pool: p}), members +} + +// serve runs n flows one at a time, each member answering in its own latency, and returns +// how many each member took +func serve(t *testing.T, b *lb.Balancer, n int, latency map[*lb.Member]time.Duration) map[*lb.Member]int { + t.Helper() + counts := make(map[*lb.Member]int) + for range n { + pk, ok := b.Pick(lb.Flow{}) + if !ok { + t.Fatal("no pick") + } + counts[pk.Member()]++ + pk.Established(latency[pk.Member()]) + pk.Done(lb.OutcomeOK) + } + return counts +} + +func TestIdentityAndTuning(t *testing.T) { + s := lt.New(lt.Options{}) + if s.Name() != lt.Name || s.Needs() != lb.NeedInflight|lb.NeedLatency { + t.Errorf("name %q needs %b", s.Name(), s.Needs()) + } + if got := s.(lb.LatencyTuner).Latency(); got.Decay != lb.DefaultLatencyDecay || got.Penalty != lb.DefaultLatencyPenalty { + t.Errorf("default tuning = %+v", got) + } + tuned := lt.New(lt.Options{Decay: time.Minute, Penalty: time.Second}).(lb.LatencyTuner).Latency() + if tuned.Decay != time.Minute || tuned.Penalty != time.Second { + t.Errorf("tuning = %+v", tuned) + } +} + +func TestPrefersTheFasterMember(t *testing.T) { + b, m := newBalancer(t, lt.Options{}, 1, 1, 1) + latency := map[*lb.Member]time.Duration{m[0]: 80 * time.Millisecond, m[1]: 5 * time.Millisecond, m[2]: 40 * time.Millisecond} + serve(t, b, 30, latency) + counts := serve(t, b, 300, latency) + // all but the odd flow that re-tries a member whose average has faded + if counts[m[1]] < 290 { + t.Errorf("with the pool idle the fastest member took %d of 300 flows", counts[m[1]]) + } +} + +// resolution is nanoseconds: a member answering in 200µs still ranks ahead of one at 900µs +func TestSubMillisecondMembersAreRanked(t *testing.T) { + b, m := newBalancer(t, lt.Options{}, 1, 1) + latency := map[*lb.Member]time.Duration{m[0]: 900 * time.Microsecond, m[1]: 200 * time.Microsecond} + serve(t, b, 20, latency) + if counts := serve(t, b, 100, latency); counts[m[1]] < 95 { + t.Errorf("the 200µs member took %d of 100 flows", counts[m[1]]) + } +} + +// in-flight work multiplies a member's score, so load spills to slower members instead of +// queueing on the fastest +func TestLoadSpillsToSlowerMembers(t *testing.T) { + b, m := newBalancer(t, lt.Options{}, 1, 1) + latency := map[*lb.Member]time.Duration{m[0]: 10 * time.Millisecond, m[1]: 30 * time.Millisecond} + serve(t, b, 20, latency) + for range 40 { + pk, _ := b.Pick(lb.Flow{}) + pk.Established(latency[pk.Member()]) + } + fast, slow := m[0].Stats().Inflight(), m[1].Stats().Inflight() + if slow == 0 || fast <= slow || fast > 4*slow { + t.Errorf("held flows = %d fast and %d slow; want about 3 to 1", fast, slow) + } +} + +// a member with no sample is scored at its peers' mean: it gets work, but not all of it +func TestColdMemberIsNeitherFloodedNorStarved(t *testing.T) { + members, healths := lbtest.Members(1, 1, 1) + healths[2].Set(-1) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + defer p.Stop() + b := lb.NewBalancer(newSelector(), lb.BalancerOptions{Pool: p}) + latency := map[*lb.Member]time.Duration{ + members[0]: 10 * time.Millisecond, members[1]: 30 * time.Millisecond, members[2]: 20 * time.Millisecond, + } + serve(t, b, 40, latency) + // the cold member joins while its peers are busy + healths[2].Set(1) + var held []lb.Pick + for range 60 { + pk, _ := b.Pick(lb.Flow{}) + held = append(held, pk) + } + cold := members[2].Stats().Inflight() + if cold == 0 { + t.Error("the cold member was starved") + } + if cold > 40 { + t.Errorf("the cold member was flooded with %d of 60 held flows", cold) + } + for _, pk := range held { + pk.Done(lb.OutcomeOK) + } +} + +// a member that fails in a millisecond must not win on speed +func TestFastFailingMemberLoses(t *testing.T) { + b, m := newBalancer(t, lt.Options{}, 1, 1) + healthy, failing := m[0], m[1] + counts := make(map[*lb.Member]int) + for range 400 { + pk, _ := b.Pick(lb.Flow{}) + counts[pk.Member()]++ + if pk.Member() == failing { + pk.Established(time.Millisecond) + pk.Done(lb.OutcomeFailed) + continue + } + pk.Established(50 * time.Millisecond) + pk.Done(lb.OutcomeOK) + } + if counts[failing] > 8 { + t.Errorf("the fast-failing member took %d of 400 flows", counts[failing]) + } + if failing.Stats().Latency() < lb.DefaultLatencyPenalty { + t.Errorf("a failure was recorded as %v", failing.Stats().Latency()) + } + _ = healthy +} + +// a penalty fades while the member is passed over, so it is tried again and, once it answers +// well, forgiven +func TestDecayForgives(t *testing.T) { + b, m := newBalancer(t, lt.Options{Decay: 20 * time.Millisecond, Penalty: 200 * time.Millisecond}, 1, 1) + good, recovering := m[0], m[1] + latency := map[*lb.Member]time.Duration{good: 10 * time.Millisecond, recovering: 2 * time.Millisecond} + for recovering.Stats().Failures() == 0 { + pk, _ := b.Pick(lb.Flow{}) + if pk.Member() == recovering { + pk.Done(lb.OutcomeFailed) + continue + } + pk.Established(latency[good]) + pk.Done(lb.OutcomeOK) + } + if counts := serve(t, b, 20, latency); counts[recovering] != 0 { + t.Fatalf("a freshly penalized member took %d of 20 flows", counts[recovering]) + } + // 200ms fades below the peer's 10ms after ln(20) decays, about 60ms + deadline := time.Now().Add(5 * time.Second) + for recovering.Stats().Failures() != 0 { + if time.Now().After(deadline) { + t.Fatal("the penalized member was never tried again") + } + time.Sleep(5 * time.Millisecond) + serve(t, b, 1, latency) + } + time.Sleep(100 * time.Millisecond) + if counts := serve(t, b, 50, latency); counts[recovering] < 45 { + t.Errorf("the recovered, faster member took %d of 50 flows", counts[recovering]) + } +} + +// with the pool idle, a member with no sample still gets a turn: it ties with the fastest +// member rather than waiting behind it for load that may never come +func TestColdMemberIsTriedWhenIdle(t *testing.T) { + b, m := newBalancer(t, lt.Options{}, 1, 1, 1) + latency := map[*lb.Member]time.Duration{m[0]: 5 * time.Millisecond, m[1]: 9 * time.Millisecond, m[2]: 300 * time.Millisecond} + counts := serve(t, b, 60, latency) + for _, member := range m { + if counts[member] == 0 { + t.Errorf("%s was never tried in 60 sequential flows: %v", member.Name(), counts) + } + } + if counts[m[2]] > 6 { + t.Errorf("the slow member took %d of 60 flows once it had been measured", counts[m[2]]) + } + // when the only sampled members are failing, a cold one is scored against them + failing, fresh := newBalancer(t, lt.Options{}, 1, 1) + pk, _ := failing.Pick(lb.Flow{}) + pk.Done(lb.OutcomeFailed) + next, _ := failing.Pick(lb.Flow{}) + if next.Member() == pk.Member() { + t.Errorf("a member with no sample lost to one that had just failed") + } + next.Done(lb.OutcomeOK) + _ = fresh +} diff --git a/pkg/lb/member.go b/pkg/lb/member.go new file mode 100644 index 000000000..613de05ac --- /dev/null +++ b/pkg/lb/member.go @@ -0,0 +1,113 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +// MemberOptions describes a Member to NewMember. +type MemberOptions struct { + // Name identifies the member. Named members must be unique within a pool. + Name string + // Group is the member's replica group; empty means the member's Name. + Group string + // Weight is the member's relative share; values below 1 mean 1. + Weight int + // Tier is the member's failover tier. A pool selects from the lowest tier that has an + // eligible member, so a higher tier stands by; values below 0 mean 0. + Tier int + // Health is the member's health source; nil means always eligible. + Health Health + // Stats carries runtime state over from a member this one replaces; nil starts fresh. + Stats *Stats + // Value is the owner's payload, returned untouched by Member.Value. + Value any +} + +// Member is one pool entry. It is immutable; its Stats and Health are live. +type Member struct { + name string + group string + weight int + tier int + hash uint64 + health Health + stats *Stats + // Value is the owner's payload: whatever it dispatches to. The core never reads it. + Value any +} + +// NewMember returns a Member for the provided options. +func NewMember(o MemberOptions) *Member { + m := &Member{ + name: o.Name, + group: o.Group, + weight: max(o.Weight, 1), + tier: max(o.Tier, 0), + hash: hashString(o.Name), + health: o.Health, + stats: o.Stats, + Value: o.Value, + } + if m.group == "" { + m.group = m.name + } + if m.stats == nil { + m.stats = &Stats{} + } + return m +} + +// Name returns the member's name. +func (m *Member) Name() string { return m.name } + +// Group returns the member's replica group, which is its name unless one was set. +func (m *Member) Group() string { return m.group } + +// Weight returns the member's relative share, always at least 1. +func (m *Member) Weight() int { return m.weight } + +// Tier returns the member's failover tier, 0 being the first selected from. +func (m *Member) Tier() int { return m.tier } + +// Hash returns a hash of the member's name that is stable across processes and restarts. +func (m *Member) Hash() uint64 { return m.hash } + +// Health returns the member's health source, or nil when it has none. +func (m *Member) Health() Health { return m.health } + +// Stats returns the member's runtime state; never nil. +func (m *Member) Stats() *Stats { return m.stats } + +// eligible reports whether the member's current status meets floor and it is not ejected +func (m *Member) eligible(floor int32, nowNano int64) bool { + if m.stats.ejectedUntil.Load() > nowNano { + return false + } + return m.health == nil || m.health.Get() >= floor +} + +const ( + fnvOffset64 uint64 = 14695981039346656037 + fnvPrime64 uint64 = 1099511628211 +) + +// hashString is FNV-1a: unseeded, so every process agrees on a name's hash +func hashString(s string) uint64 { + h := fnvOffset64 + for i := range len(s) { + h ^= uint64(s[i]) + h *= fnvPrime64 + } + return h +} diff --git a/pkg/lb/observer.go b/pkg/lb/observer.go new file mode 100644 index 000000000..f0eb7a2a3 --- /dev/null +++ b/pkg/lb/observer.go @@ -0,0 +1,49 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +// EventKind identifies what an Event reports. +type EventKind uint8 + +const ( + // EventSnapshot reports that a pool published a new snapshot. + EventSnapshot EventKind = iota + 1 + // EventPanic reports a panic recovered while rebuilding a snapshot. + EventPanic + // EventEjected reports a member taken out of selection for repeated connect failures. + EventEjected +) + +// Event is one notification to an Observer. Fields beyond Kind are set as the kind requires. +type Event struct { + Kind EventKind + // Gen, Eligible, Configured and Tier describe the snapshot of an EventSnapshot + Gen uint64 + Eligible int + Configured int + Tier int + // Panic and Stack carry the recovered value and stack of an EventPanic + Panic any + Stack []byte + // Member names the member of an EventEjected + Member string +} + +// Observer receives events from the core, which itself neither logs nor meters. Observe is +// never called on a selection path, and must not call back into the pool that invoked it. +type Observer interface { + Observe(Event) +} diff --git a/pkg/lb/p2c/p2c.go b/pkg/lb/p2c/p2c.go new file mode 100644 index 000000000..c0080d05e --- /dev/null +++ b/pkg/lb/p2c/p2c.go @@ -0,0 +1,134 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package p2c is the power of two choices strategy: two members are drawn at random and the +// flow goes to the less loaded of them. It approximates least connections at a cost that does +// not grow with the pool. +package p2c + +import ( + "math/rand/v2" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Name is the strategy's name. +const Name = "power_of_two_choices" + +// New returns a power of two choices selector. +func New() lb.Selector { + return selector{} +} + +type selector struct{} + +func (selector) Name() string { return Name } + +func (selector) Needs() lb.Needs { return lb.NeedInflight } + +func (selector) Prepare(snap *lb.Snapshot) lb.Prepared { + members := snap.Members + p := &prepared{members: members, n: uint64(len(members)), uniform: true} + for _, m := range members { + if m.Weight() != members[0].Weight() { + p.uniform = false + break + } + } + if p.uniform { + return p + } + // member i owns the span [cumulative[i-1], cumulative[i]) of the sampling space + p.cumulative = make([]uint64, len(members)) + for i, m := range members { + p.total += uint64(m.Weight()) // #nosec G115 -- a member's weight is at least 1 + p.cumulative[i] = p.total + } + return p +} + +type prepared struct { + members []*lb.Member + n uint64 + uniform bool + cumulative []uint64 + total uint64 +} + +func (p *prepared) Select(lb.Flow) *lb.Member { + if p.n == 1 { + return p.members[0] + } + r := rand.Uint64() // #nosec G404 -- load spreading, not a secret + hi, lo := r>>32, r&0xffffffff + var a, b *lb.Member + if p.uniform { + i := hi * p.n >> 32 + j := lo * (p.n - 1) >> 32 + if j >= i { + j++ + } + a, b = p.members[i], p.members[j] + } else { + a, b = p.weightedPair(hi, lo) + } + // the lower in-flight count per unit of weight wins, compared without dividing; a tie + // goes to the first draw, which keeps an idle pool's split proportional to weight + loadA := a.Stats().Inflight() * int64(b.Weight()) + loadB := b.Stats().Inflight() * int64(a.Weight()) + if loadB < loadA { + return b + } + return a +} + +// weightedPair draws two distinct members with probability proportional to weight: the +// second draw is made over the sampling space with the first member's span cut out +func (p *prepared) weightedPair(hi, lo uint64) (first, second *lb.Member) { + i := p.find(scale(hi, p.total)) + var start uint64 + if i > 0 { + start = p.cumulative[i-1] + } + width := p.cumulative[i] - start + k := scale(lo, p.total-width) + if k >= start { + k += width + } + return p.members[i], p.members[p.find(k)] +} + +// scale maps 32 random bits onto [0, n); exact to within one part in 2^32 for n below 2^32, +// and still in range above that +func scale(r32, n uint64) uint64 { + if n>>32 == 0 { + return r32 * n >> 32 + } + return (r32<<32 | r32) % n +} + +// find returns the member whose span holds k, which must be below total +func (p *prepared) find(k uint64) int { + lo, hi := 0, len(p.cumulative)-1 + for lo < hi { + mid := int(uint(lo+hi) >> 1) // #nosec G115 -- the sum of two slice indexes + if p.cumulative[mid] > k { + hi = mid + } else { + lo = mid + 1 + } + } + return lo +} diff --git a/pkg/lb/p2c/p2c_test.go b/pkg/lb/p2c/p2c_test.go new file mode 100644 index 000000000..f30713dff --- /dev/null +++ b/pkg/lb/p2c/p2c_test.go @@ -0,0 +1,110 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package p2c_test + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" + "github.com/trickstercache/trickster/v2/pkg/lb/p2c" +) + +func TestConformance(t *testing.T) { + lbtest.Run(t, p2c.New, lbtest.Options{}) +} + +func TestWeights(t *testing.T) { + lbtest.RunWeighted(t, p2c.New, lbtest.WeightOptions{}) +} + +func BenchmarkSelect(b *testing.B) { + lbtest.Bench(b, p2c.New) +} + +func newBalancer(t *testing.T, weights ...int) (*lb.Balancer, []*lb.Member) { + t.Helper() + members, _ := lbtest.Members(weights...) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return lb.NewBalancer(p2c.New(), lb.BalancerOptions{Pool: p}), members +} + +func TestIdentity(t *testing.T) { + s := p2c.New() + if s.Name() != p2c.Name || s.Needs() != lb.NeedInflight { + t.Errorf("name %q needs %b", s.Name(), s.Needs()) + } +} + +// with flows held open, the spread stays far tighter than random placement would leave it +func TestBalancesHeldFlows(t *testing.T) { + b, members := newBalancer(t, 1, 1, 1, 1, 1, 1, 1, 1) + for range 8000 { + b.Pick(lb.Flow{}) + } + for _, m := range members { + if got := m.Stats().Inflight(); got < 960 || got > 1040 { + t.Errorf("%s holds %d of 8000 held flows, want 1000 within 40", m.Name(), got) + } + } +} + +// a member already carrying a load is passed over for an idle one whenever the two are drawn +func TestPrefersTheIdleMember(t *testing.T) { + b, members := newBalancer(t, 1, 1) + var busy *lb.Member + for range 50 { + pk, _ := b.Pick(lb.Flow{}) + if busy == nil { + busy = pk.Member() + } + if pk.Member() != busy { + pk.Done(lb.OutcomeOK) + } + } + before := busy.Stats().Inflight() + for range 200 { + pk, _ := b.Pick(lb.Flow{}) + if pk.Member() == busy { + t.Fatalf("picked the member holding %d flows over an idle one", before) + } + pk.Done(lb.OutcomeOK) + } + _ = members +} + +// weights beyond 32 bits of total still sample in range +func TestHugeWeights(t *testing.T) { + b, members := newBalancer(t, 1<<31, 1<<31, 1<<31, 7) + seen := make(map[*lb.Member]bool) + for range 2000 { + pk, ok := b.Pick(lb.Flow{}) + if !ok { + t.Fatal("no pick") + } + seen[pk.Member()] = true + pk.Done(lb.OutcomeOK) + } + for _, m := range members[:3] { + if !seen[m] { + t.Errorf("%s was never drawn", m.Name()) + } + } +} diff --git a/pkg/lb/pool.go b/pkg/lb/pool.go new file mode 100644 index 000000000..f0db84260 --- /dev/null +++ b/pkg/lb/pool.go @@ -0,0 +1,212 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import ( + "errors" + "fmt" + "runtime/debug" + "slices" + "sync" + "sync/atomic" + "time" +) + +var ( + // ErrNilMember is returned by NewPool for a nil member. + ErrNilMember = errors.New("lb: nil pool member") + // ErrDuplicateMember is returned by NewPool when two members share a name. + ErrDuplicateMember = errors.New("lb: duplicate pool member name") +) + +// Snapshot is an immutable view of a pool's eligible members. Holders may retain it, and must +// not modify it. +type Snapshot struct { + // Members holds the eligible members of the lowest tier that has any, in pool order + Members []*Member + // Tier is the failover tier the members belong to; 0 when there are none + Tier int + // Gen increases with each snapshot a pool publishes, starting at 1 + Gen uint64 +} + +// PoolOptions are the optional settings of a Pool. +type PoolOptions struct { + // Observer receives the pool's events; nil discards them. + Observer Observer +} + +// Pool is a fixed membership whose eligible subset is republished as member health changes. +// Reading the subset costs one atomic load. +type Pool struct { + members []*Member + floor int32 + observer Observer + snap atomic.Pointer[Snapshot] + // held across read-build-publish so snapshots are published in order + mtx sync.Mutex + gen uint64 + stopped bool + subs []Subscription +} + +// NewPool returns a started Pool over members, whose snapshots hold the members with a health +// status of at least floor. Membership is fixed: a change of membership is a new Pool. +func NewPool(members []*Member, floor int, opts ...PoolOptions) (*Pool, error) { + names := make(map[string]struct{}, len(members)) + for _, m := range members { + if m == nil { + return nil, ErrNilMember + } + if m.name == "" { + continue + } + if _, dup := names[m.name]; dup { + return nil, fmt.Errorf("%w: %s", ErrDuplicateMember, m.name) + } + names[m.name] = struct{}{} + } + p := &Pool{members: slices.Clone(members), floor: clampFloor(floor)} + if len(opts) > 0 { + p.observer = opts[0].Observer + } + p.snap.Store(&Snapshot{}) + // subscribe before the first build, so a transition between the two forces a rebuild + // that queues behind the build rather than being lost + subs := make([]Subscription, 0, len(p.members)) + for _, m := range p.members { + if n, ok := m.health.(Notifier); ok { + subs = append(subs, n.OnChange(p.onChange)) + } + } + p.mtx.Lock() + p.subs = subs + p.mtx.Unlock() + p.Refresh() + return p, nil +} + +func clampFloor(floor int) int32 { + const lo, hi = -1 << 31, 1<<31 - 1 + return int32(min(max(floor, lo), hi)) // #nosec G115 -- clamped to the int32 range +} + +// Snapshot returns the pool's current eligible members; never nil. +func (p *Pool) Snapshot() *Snapshot { + return p.snap.Load() +} + +// Configured returns every member of the pool, eligible or not, in pool order. +func (p *Pool) Configured() []*Member { + return slices.Clone(p.members) +} + +// Len returns the number of configured members. +func (p *Pool) Len() int { + return len(p.members) +} + +// Floor returns the minimum health status of an eligible member. +func (p *Pool) Floor() int { + return int(p.floor) +} + +// onChange rebuilds only when a transition carries the member across the floor +func (p *Pool) onChange(prev, next int32) { + if (prev >= p.floor) == (next >= p.floor) { + return + } + p.Refresh() +} + +// Refresh rebuilds and publishes the snapshot from the members' current statuses. A pool +// whose members announce their transitions does this on its own. No-op once stopped. +func (p *Pool) Refresh() { + if ev, ok := p.rebuild(); ok && p.observer != nil { + p.observer.Observe(ev) + } +} + +func (p *Pool) rebuild() (ev Event, ok bool) { + p.mtx.Lock() + defer p.mtx.Unlock() + defer func() { + if r := recover(); r != nil { + ev, ok = Event{Kind: EventPanic, Panic: r, Stack: debug.Stack()}, true + } + }() + if p.stopped { + return Event{}, false + } + // statuses are read here, not taken from a callback's arguments, which may be stale + nowNano := time.Now().UnixNano() + eligible := make([]*Member, 0, len(p.members)) + tier := 0 + for _, m := range p.members { + if !m.eligible(p.floor, nowNano) || (len(eligible) > 0 && m.tier > tier) { + continue + } + if m.tier < tier { + // a lower tier has a live member after all; the standbys collected so far stand down + eligible = eligible[:0] + } + tier = m.tier + eligible = append(eligible, m) + } + p.gen++ + p.snap.Store(&Snapshot{Members: eligible, Tier: tier, Gen: p.gen}) + return Event{ + Kind: EventSnapshot, Gen: p.gen, Eligible: len(eligible), Configured: len(p.members), Tier: tier, + }, true +} + +// eject takes m out of selection until the provided time, unless that would leave the pool +// without a member or put more than maxPercent of its members out at once. It reports whether +// the member was ejected; the caller then refreshes the pool. +func (p *Pool) eject(m *Member, until time.Time, maxPercent int) bool { + p.mtx.Lock() + defer p.mtx.Unlock() + if p.stopped || !slices.Contains(p.members, m) { + return false + } + nowNano := time.Now().UnixNano() + var live, out int + for _, o := range p.members { + switch { + case o.stats.ejectedUntil.Load() > nowNano: + out++ + case o.eligible(p.floor, nowNano): + live++ + } + } + if m.stats.ejectedUntil.Load() > nowNano || live <= 1 || (out+1)*100 > maxPercent*len(p.members) { + return false + } + m.stats.ejectedUntil.Store(until.UnixNano()) + return true +} + +// Stop ends the pool's subscriptions. Its last snapshot stays readable and is never replaced. +func (p *Pool) Stop() { + p.mtx.Lock() + subs := p.subs + p.subs = nil + p.stopped = true + p.mtx.Unlock() + for _, s := range subs { + s.Unsubscribe() + } +} diff --git a/pkg/lb/pool_test.go b/pkg/lb/pool_test.go new file mode 100644 index 000000000..66c07d3cc --- /dev/null +++ b/pkg/lb/pool_test.go @@ -0,0 +1,411 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import ( + "errors" + "slices" + "sync" + "sync/atomic" + "testing" +) + +// fakeHealth is a Notifier with the contract the core asks for: callbacks run synchronously, +// outside its lock, and Unsubscribe never waits +type fakeHealth struct { + status atomic.Int32 + mtx sync.Mutex + subs []*fakeSub +} + +type fakeSub struct { + h *fakeHealth + fn func(prev, next int32) + dead atomic.Bool +} + +func (s *fakeSub) Unsubscribe() { + s.dead.Store(true) + s.h.mtx.Lock() + defer s.h.mtx.Unlock() + s.h.subs = slices.DeleteFunc(slices.Clone(s.h.subs), func(o *fakeSub) bool { return o == s }) +} + +func newHealth(status int32) *fakeHealth { + h := &fakeHealth{} + h.status.Store(status) + return h +} + +func (h *fakeHealth) Get() int32 { return h.status.Load() } + +func (h *fakeHealth) OnChange(fn func(prev, next int32)) Subscription { + s := &fakeSub{h: h, fn: fn} + h.mtx.Lock() + defer h.mtx.Unlock() + h.subs = append(slices.Clone(h.subs), s) + return s +} + +func (h *fakeHealth) set(next int32) { + prev := h.status.Swap(next) + h.mtx.Lock() + subs := h.subs + h.mtx.Unlock() + for _, s := range subs { + if !s.dead.Load() { + s.fn(prev, next) + } + } +} + +func (h *fakeHealth) subscribers() int { + h.mtx.Lock() + defer h.mtx.Unlock() + return len(h.subs) +} + +// plainHealth has no Notifier, so its pool follows it only through Refresh +type plainHealth struct{ status atomic.Int32 } + +func (h *plainHealth) Get() int32 { return h.status.Load() } + +type recordingObserver struct { + mtx sync.Mutex + events []Event +} + +func (o *recordingObserver) Observe(ev Event) { + o.mtx.Lock() + defer o.mtx.Unlock() + o.events = append(o.events, ev) +} + +func (o *recordingObserver) kinds(k EventKind) []Event { + o.mtx.Lock() + defer o.mtx.Unlock() + var out []Event + for _, ev := range o.events { + if ev.Kind == k { + out = append(out, ev) + } + } + return out +} + +func names(s *Snapshot) []string { + out := make([]string, len(s.Members)) + for i, m := range s.Members { + out[i] = m.Name() + } + return out +} + +func mustPool(t *testing.T, members []*Member, floor int, opts ...PoolOptions) *Pool { + t.Helper() + p, err := NewPool(members, floor, opts...) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return p +} + +func TestNewMember(t *testing.T) { + h := newHealth(1) + payload := &struct{}{} + m := NewMember(MemberOptions{Name: "a", Weight: 3, Health: h, Value: payload}) + if m.Name() != "a" || m.Group() != "a" || m.Weight() != 3 || m.Health() != h || m.Value != payload { + t.Errorf("unexpected member: %+v", m) + } + if m.Stats() == nil { + t.Fatal("a member always has stats") + } + st := m.Stats() + if st.Inflight() != 0 || st.Latency() != 0 || st.Failures() != 0 || !st.LastSample().IsZero() { + t.Error("fresh stats are not zero") + } + st.stamp.Store(42) + if st.LastSample().UnixNano() != 42 { + t.Error("last sample time was not reported") + } + grouped := NewMember(MemberOptions{Name: "a2", Group: "shard", Weight: -4, Stats: st}) + if grouped.Group() != "shard" || grouped.Weight() != 1 || grouped.Stats() != st { + t.Errorf("unexpected member: %+v", grouped) + } + // the hash is a fixed function of the name: every process must agree on it + if got := NewMember(MemberOptions{Name: "a"}).Hash(); got != 0xaf63dc4c8601ec8c || got != m.Hash() { + t.Errorf("hash of %q = %#x", "a", got) + } + if grouped.Hash() == m.Hash() { + t.Error("different names share a hash") + } +} + +func TestNewPoolRefusesBadMembership(t *testing.T) { + a := NewMember(MemberOptions{Name: "a"}) + if _, err := NewPool([]*Member{a, nil}, 0); !errors.Is(err, ErrNilMember) { + t.Errorf("nil member: %v", err) + } + if _, err := NewPool([]*Member{a, NewMember(MemberOptions{Name: "a"})}, 0); !errors.Is(err, ErrDuplicateMember) { + t.Errorf("duplicate member: %v", err) + } + // unnamed members are anonymous, not duplicates of one another + p := mustPool(t, []*Member{NewMember(MemberOptions{}), NewMember(MemberOptions{})}, 0) + if len(p.Snapshot().Members) != 2 { + t.Error("anonymous members were refused") + } + empty := mustPool(t, nil, 0) + if s := empty.Snapshot(); s == nil || len(s.Members) != 0 || s.Gen != 1 { + t.Errorf("empty pool snapshot = %+v", s) + } +} + +func TestPoolSnapshotFollowsHealth(t *testing.T) { + ha, hb, hc := newHealth(1), newHealth(0), newHealth(-1) + members := []*Member{ + NewMember(MemberOptions{Name: "a", Health: ha}), + NewMember(MemberOptions{Name: "b", Health: hb}), + NewMember(MemberOptions{Name: "c", Health: hc}), + NewMember(MemberOptions{Name: "always"}), + } + obs := &recordingObserver{} + p := mustPool(t, members, 1, PoolOptions{Observer: obs}) + if p.Len() != 4 || p.Floor() != 1 || len(p.Configured()) != 4 { + t.Errorf("len %d floor %d", p.Len(), p.Floor()) + } + p.Configured()[0] = nil + if p.Configured()[0] != members[0] { + t.Error("Configured exposed the pool's own slice") + } + first := p.Snapshot() + if got := names(first); !slices.Equal(got, []string{"a", "always"}) || first.Gen != 1 { + t.Fatalf("initial snapshot = %v gen %d", got, first.Gen) + } + + // a transition across the floor is visible to the very next Snapshot, in pool order + hc.set(1) + if got := names(p.Snapshot()); !slices.Equal(got, []string{"a", "c", "always"}) { + t.Fatalf("after c passes = %v", got) + } + if got := names(first); !slices.Equal(got, []string{"a", "always"}) { + t.Errorf("a published snapshot was modified: %v", got) + } + + // a transition that stays on one side of the floor publishes nothing + before := p.Snapshot() + hb.set(-1) + hb.set(-2) + hb.set(0) + ha.set(1) + if p.Snapshot() != before { + t.Error("a transition that did not cross the floor republished") + } + ha.set(0) + if got := p.Snapshot(); !slices.Equal(names(got), []string{"c", "always"}) || got.Gen != 3 { + t.Errorf("after a falls = %v gen %d", names(got), got.Gen) + } + snaps := obs.kinds(EventSnapshot) + if len(snaps) != 3 || snaps[2].Gen != 3 || snaps[2].Eligible != 2 || snaps[2].Configured != 4 { + t.Errorf("snapshot events = %+v", snaps) + } +} + +func TestPoolWithoutNotifierFollowsRefresh(t *testing.T) { + h := &plainHealth{} + p := mustPool(t, []*Member{NewMember(MemberOptions{Name: "a", Health: h})}, 1) + h.status.Store(1) + if len(p.Snapshot().Members) != 0 { + t.Fatal("a pool without a Notifier cannot have seen the change") + } + p.Refresh() + if len(p.Snapshot().Members) != 1 { + t.Error("Refresh did not pick up the change") + } +} + +func TestPoolStop(t *testing.T) { + h := newHealth(1) + p := mustPool(t, []*Member{NewMember(MemberOptions{Name: "a", Health: h})}, 1) + if h.subscribers() != 1 { + t.Fatalf("subscribers = %d", h.subscribers()) + } + last := p.Snapshot() + p.Stop() + p.Stop() + if h.subscribers() != 0 { + t.Error("Stop left a subscription behind") + } + h.set(-1) + p.Refresh() + // a callback captured before Stop may still arrive; it must publish nothing + p.onChange(1, -1) + if p.Snapshot() != last { + t.Error("a stopped pool republished") + } +} + +// a callback may stop the pool that is calling it +func TestPoolStopFromInsideACallback(t *testing.T) { + h := newHealth(1) + p := mustPool(t, []*Member{NewMember(MemberOptions{Name: "a", Health: h})}, 1) + h.OnChange(func(_, _ int32) { p.Stop() }) + done := make(chan struct{}) + go func() { + defer close(done) + h.set(-1) + h.set(1) + }() + <-done + if h.subscribers() != 1 { + t.Errorf("subscribers = %d, want only the test's own", h.subscribers()) + } +} + +// many goroutines flipping many members never publish out of order, and the last snapshot +// matches the final statuses +func TestPoolTransitionStorm(t *testing.T) { + const n = 12 + healths := make([]*fakeHealth, n) + members := make([]*Member, n) + for i := range n { + healths[i] = newHealth(-1) + members[i] = NewMember(MemberOptions{Name: string(rune('a' + i)), Health: healths[i]}) + } + p := mustPool(t, members, 1) + stop := make(chan struct{}) + var watcher sync.WaitGroup + watcher.Go(func() { + var last uint64 + for { + select { + case <-stop: + return + default: + } + if gen := p.Snapshot().Gen; gen < last { + t.Errorf("generation went backwards: %d after %d", gen, last) + return + } else { + last = gen + } + } + }) + var wg sync.WaitGroup + for i := range n { + wg.Go(func() { + for range 300 { + healths[i].set(1) + healths[i].set(-1) + } + if i%3 == 0 { + healths[i].set(1) + } + }) + } + wg.Wait() + close(stop) + watcher.Wait() + var want []string + for i := 0; i < n; i += 3 { + want = append(want, members[i].Name()) + } + if got := names(p.Snapshot()); !slices.Equal(got, want) { + t.Errorf("final snapshot = %v, want %v", got, want) + } +} + +// a transition that lands between subscription and the first build is never lost +func TestPoolTransitionRacingConstruction(t *testing.T) { + for range 500 { + h := newHealth(-1) + var wg sync.WaitGroup + wg.Go(func() { h.set(1) }) + p, err := NewPool([]*Member{NewMember(MemberOptions{Name: "a", Health: h})}, 1) + if err != nil { + t.Fatal(err) + } + wg.Wait() + got := len(p.Snapshot().Members) + p.Stop() + if got != 1 { + t.Fatalf("a transition racing construction was lost: %d members", got) + } + } +} + +type panickyHealth struct{ armed atomic.Bool } + +func (h *panickyHealth) Get() int32 { + if h.armed.Load() { + panic("health source blew up") + } + return 1 +} + +// a panic while rebuilding is recovered and reported, leaves the last snapshot in place, and +// does not wedge the pool +func TestPoolRecoversRebuildPanic(t *testing.T) { + h := &panickyHealth{} + obs := &recordingObserver{} + p := mustPool(t, []*Member{NewMember(MemberOptions{Name: "a", Health: h})}, 1, PoolOptions{Observer: obs}) + last := p.Snapshot() + h.armed.Store(true) + p.Refresh() + if p.Snapshot() != last { + t.Error("a failed rebuild replaced the snapshot") + } + panics := obs.kinds(EventPanic) + if len(panics) != 1 || panics[0].Panic != "health source blew up" || len(panics[0].Stack) == 0 { + t.Fatalf("panic events = %+v", panics) + } + h.armed.Store(false) + p.Refresh() + if got := p.Snapshot(); got == last || got.Gen != last.Gen+1 { + t.Error("the pool did not rebuild after a recovered panic") + } + // a panic with no observer is still recovered + quiet := mustPool(t, []*Member{NewMember(MemberOptions{Name: "a", Health: h})}, 1) + h.armed.Store(true) + quiet.Refresh() +} + +func TestSnapshotZeroAlloc(t *testing.T) { + p := mustPool(t, []*Member{NewMember(MemberOptions{Name: "a"})}, 0) + if allocs := testing.AllocsPerRun(1000, func() { _ = p.Snapshot() }); allocs != 0 { + t.Errorf("Snapshot allocates %v", allocs) + } +} + +func BenchmarkPoolSnapshot(b *testing.B) { + members := make([]*Member, 8) + for i := range members { + members[i] = NewMember(MemberOptions{Name: string(rune('a' + i)), Health: newHealth(1)}) + } + p, err := NewPool(members, 1) + if err != nil { + b.Fatal(err) + } + defer p.Stop() + b.ReportAllocs() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + if len(p.Snapshot().Members) != len(members) { + b.Fatal("short snapshot") + } + } + }) +} diff --git a/pkg/lb/rr/rr.go b/pkg/lb/rr/rr.go new file mode 100644 index 000000000..1a572567d --- /dev/null +++ b/pkg/lb/rr/rr.go @@ -0,0 +1,204 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package rr is the round robin strategy: a lock-free rotation that gives each member exactly +// its weight of every total-weight consecutive selections, with a heavier member's turns +// spread through the rotation rather than taken back to back. +package rr + +import ( + "math/bits" + "math/rand/v2" + "slices" + "sync/atomic" + + "github.com/trickstercache/trickster/v2/pkg/lb" +) + +// Name is the strategy's name. +const Name = "round_robin" + +const ( + // scheduleMax is the largest total weight whose rotation is laid out as a table, one + // entry per turn; larger totals are walked by stride instead + scheduleMax = 4096 + // linearScanMax is the pool size up to which scanning the spans beats a binary search + linearScanMax = 16 +) + +// New returns a round robin selector whose rotation starts at a random turn, so that +// replicas started together do not all send their first flows to the same member. +func New() lb.Selector { + return NewAt(rand.Uint64()) // #nosec G404 -- a starting offset, not a secret +} + +// NewAt returns a round robin selector whose rotation starts after the provided turn, for a +// caller that needs a reproducible sequence. +func NewAt(start uint64) lb.Selector { + s := &selector{} + s.pos.Store(start) + return s +} + +type selector struct { + // the rotation outlives snapshots, so a change of membership does not restart it + pos atomic.Uint64 +} + +func (s *selector) Name() string { return Name } + +func (s *selector) Needs() lb.Needs { return 0 } + +func (s *selector) Prepare(snap *lb.Snapshot) lb.Prepared { + members := snap.Members + var total uint64 + uniform := true + for _, m := range members { + total += uint64(m.Weight()) // #nosec G115 -- a member's weight is at least 1 + if m.Weight() != members[0].Weight() { + uniform = false + } + } + if uniform { + return &rotation{pos: &s.pos, members: members, n: uint64(len(members))} + } + if total <= scheduleMax { + return &schedule{pos: &s.pos, turns: layout(members, total), total: total} + } + // member i owns the span [cumulative[i-1], cumulative[i]) of the rotation + cumulative := make([]uint64, len(members)) + var sum uint64 + for i, m := range members { + sum += uint64(m.Weight()) // #nosec G115 -- at least 1 + cumulative[i] = sum + } + return &strided{pos: &s.pos, members: members, cumulative: cumulative, total: total, stride: stride(total)} +} + +// rotation serves members of one weight in turn +type rotation struct { + pos *atomic.Uint64 + members []*lb.Member + n uint64 +} + +func (r *rotation) Select(lb.Flow) *lb.Member { + return r.members[r.pos.Add(1)%r.n] +} + +// schedule is one full rotation laid out turn by turn. Any total consecutive turns are one +// pass over it from some offset, so each member is selected exactly its weight of them. +type schedule struct { + pos *atomic.Uint64 + turns []*lb.Member + total uint64 +} + +func (s *schedule) Select(lb.Flow) *lb.Member { + return s.turns[s.pos.Add(1)%s.total] +} + +// layout spaces each member's turns evenly through the rotation, and staggers the members so +// that those of one weight do not all come due together: member i of n has its j-th turn due +// at (j + (i + 1/2) / n) * total / weight, and turns are taken in order of when they are due +func layout(members []*lb.Member, total uint64) []*lb.Member { + // a turn is due at numerator / (2 * n * weight), kept as the fraction's parts so that + // two turns compare by cross-multiplying rather than by dividing + type turn struct { + numerator uint64 + weight uint64 + member int + } + turns := make([]turn, 0, total) + n := uint64(len(members)) + for i, m := range members { + w := uint64(m.Weight()) // #nosec G115 -- at least 1 + for j := range w { + phase := 2*(j*n+uint64(i)) + 1 // #nosec G115 -- a slice index + turns = append(turns, turn{numerator: phase * total, weight: w, member: i}) + } + } + // the products are at most about 4 * scheduleMax^4, inside 64 bits + slices.SortStableFunc(turns, func(a, b turn) int { + l, r := a.numerator*b.weight, b.numerator*a.weight + switch { + case l < r: + return -1 + case l > r: + return 1 + } + return a.member - b.member + }) + out := make([]*lb.Member, len(turns)) + for i, t := range turns { + out[i] = members[t.member] + } + return out +} + +// strided walks a rotation too long to lay out: turn c lands on position c*stride mod total. +// The stride shares no factor with total, so total consecutive turns land on every position +// once, which keeps the apportionment exact; near total/phi, it also scatters them evenly. +type strided struct { + pos *atomic.Uint64 + members []*lb.Member + cumulative []uint64 + total uint64 + stride uint64 +} + +func (r *strided) Select(lb.Flow) *lb.Member { + hi, lo := bits.Mul64(r.pos.Add(1)%r.total, r.stride) + _, k := bits.Div64(hi, lo, r.total) + if len(r.cumulative) <= linearScanMax { + for i, end := range r.cumulative { + if end > k { + return r.members[i] + } + } + } + // binary search for the first member whose span ends beyond k; k < total, so one does + low, high := 0, len(r.cumulative)-1 + for low < high { + mid := int(uint(low+high) >> 1) // #nosec G115 -- the sum of two slice indexes + if r.cumulative[mid] > k { + high = mid + } else { + low = mid + 1 + } + } + return r.members[low] +} + +// stride returns the number nearest total/phi that shares no factor with total. The search +// ends by the time it has widened to 1, which shares a factor with nothing. +func stride(total uint64) uint64 { + const invPhi = 0.6180339887498949 + ideal := uint64(float64(total) * invPhi) + for d := uint64(0); ; d++ { + for _, s := range [2]uint64{ideal + d, ideal - d} { + if s >= 1 && s < total && gcd(s, total) == 1 { + return s + } + } + } +} + +func gcd(a, b uint64) uint64 { + for b != 0 { + a, b = b, a%b + } + return a +} diff --git a/pkg/lb/rr/rr_test.go b/pkg/lb/rr/rr_test.go new file mode 100644 index 000000000..e96b25477 --- /dev/null +++ b/pkg/lb/rr/rr_test.go @@ -0,0 +1,189 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package rr_test + +import ( + "slices" + "testing" + + "github.com/trickstercache/trickster/v2/pkg/lb" + "github.com/trickstercache/trickster/v2/pkg/lb/lbtest" + "github.com/trickstercache/trickster/v2/pkg/lb/rr" +) + +func TestConformance(t *testing.T) { + lbtest.Run(t, rr.New, lbtest.Options{ExactWeights: true}) +} + +func TestWeights(t *testing.T) { + lbtest.RunWeighted(t, rr.New, lbtest.WeightOptions{Tolerance: 0.001}) +} + +func BenchmarkSelect(b *testing.B) { + lbtest.Bench(b, rr.New) +} + +func balancerAt(t *testing.T, start uint64, weights ...int) *lb.Balancer { + t.Helper() + members, _ := lbtest.Members(weights...) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + t.Cleanup(p.Stop) + return lb.NewBalancer(rr.NewAt(start), lb.BalancerOptions{Pool: p}) +} + +func sequence(t *testing.T, b *lb.Balancer, n int) []int { + t.Helper() + seq := make([]int, n) + for i := range seq { + pk, ok := b.Pick(lb.Flow{}) + if !ok { + t.Fatal("no pick") + } + seq[i] = pk.Member().Value.(int) + } + return seq +} + +func TestIdentity(t *testing.T) { + s := rr.New() + if s.Name() != rr.Name || s.Needs() != 0 { + t.Errorf("name %q needs %d", s.Name(), s.Needs()) + } +} + +// a heavier member's turns are spread through the rotation, not taken back to back +func TestSequence(t *testing.T) { + for _, test := range []struct { + weights []int + want []int + }{ + {[]int{1, 1, 1}, []int{1, 2, 0, 1, 2, 0}}, + {[]int{3, 1}, []int{0, 0, 1, 0, 0, 0, 1, 0}}, + {[]int{1, 3, 2}, []int{1, 2, 1, 1, 2, 0, 1, 2, 1, 1, 2, 0}}, + {[]int{2, 2}, []int{1, 0, 1, 0}}, + } { + if got := sequence(t, balancerAt(t, 0, test.weights...), len(test.want)); !slices.Equal(got, test.want) { + t.Errorf("weights %v: sequence = %v, want %v", test.weights, got, test.want) + } + } +} + +// replicas started together must not all begin on the same member +func TestRandomStart(t *testing.T) { + first := make(map[int]bool) + for range 64 { + members, _ := lbtest.Members(1, 1, 1, 1) + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + b := lb.NewBalancer(rr.New(), lb.BalancerOptions{Pool: p}) + first[sequence(t, b, 1)[0]] = true + p.Stop() + } + if len(first) < 3 { + t.Errorf("64 new selectors began on only %d of 4 members", len(first)) + } +} + +// no member is handed more than its fair share, plus a turn or two, of any run of turns: +// the opposite of serving a weight-9 member nine times running +func TestSpread(t *testing.T) { + big := make([]int, 40) + for i := range big { + big[i] = 1 + i%7 + } + heavy := make([]int, 40) + for i := range heavy { + heavy[i] = 150 + i + } + for _, weights := range [][]int{ + {3, 1}, {1, 3, 2}, {9, 1}, {1, 1, 1, 9}, {7, 1, 1}, {5, 3, 2}, {100, 10, 1}, {2, 3}, big, + // beyond the laid-out schedule: walked by stride, over few members and over many + {4000, 3000, 2000, 1000}, {50000, 1}, {6000, 1, 1, 1}, heavy, + } { + var total int + for _, w := range weights { + total += w + } + seq := sequence(t, balancerAt(t, 0, weights...), 2*total) + worst := 0 + for m, w := range weights { + // prefix[i] is how many of the first i turns went to m + prefix := make([]int, len(seq)+1) + for i, got := range seq { + prefix[i+1] = prefix[i] + if got == m { + prefix[i+1]++ + } + } + // every run length of a short rotation; a sample of them for a long one + step := 1 + total/64 + for k := 1; k <= total; k += step { + fair := (w*k + total - 1) / total + for start := 0; start+k <= len(seq); start++ { + worst = max(worst, prefix[start+k]-prefix[start]-fair) + } + } + } + if limit := spreadLimit(total); worst > limit { + t.Errorf("weights %v: a member took %d turns more than its fair share of a run, limit %d", + weights, worst, limit) + } + } +} + +// a laid-out rotation keeps every member within one turn of fair; a strided one drifts a +// little further, as any fixed stride must +func spreadLimit(total int) int { + if total <= 4096 { + return 1 + } + return 4 +} + +// the rotation belongs to the selector: a pool swapped for one of the same membership +// continues it rather than restarting it +func TestRotationSurvivesPoolSwap(t *testing.T) { + members, _ := lbtest.Members(2, 1, 3) + b := lb.NewBalancer(rr.NewAt(0)) + var got []int + for range 6 { + p, err := lb.NewPool(members, 1) + if err != nil { + t.Fatal(err) + } + defer p.Stop() + b.SetPool(p) + got = append(got, sequence(t, b, 5)...) + } + period := got[:6] + for i, m := range got { + if m != period[i%6] { + t.Fatalf("pick %d = member %d; the rotation restarted across a swap: %v", i, m, got) + } + } + counts := make(map[int]int) + for _, m := range period { + counts[m]++ + } + if counts[0] != 2 || counts[1] != 1 || counts[2] != 3 { + t.Errorf("one rotation = %v", period) + } +} diff --git a/pkg/lb/selector.go b/pkg/lb/selector.go new file mode 100644 index 000000000..fb3a8e0b0 --- /dev/null +++ b/pkg/lb/selector.go @@ -0,0 +1,90 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +// Needs declares what a selector consumes, so that no caller does work for a selector that +// would ignore it: a selector with no needs costs a pick nothing beyond its own Select. +type Needs uint8 + +const ( + // NeedKey asks for Flow.Key, as affinity strategies do. + NeedKey Needs = 1 << iota + // NeedInflight asks for each member's in-flight count to be kept. + NeedInflight + // NeedLatency asks for each member's latency average to be kept. + NeedLatency +) + +// Has reports whether every need in want is declared. +func (n Needs) Has(want Needs) bool { + return n&want == want +} + +// Flow is the unit of work being balanced: a request, a connection, a session. It is passed +// by value, and is a struct so it can grow without changing Selector. +type Flow struct { + // Key is a hash of whatever identifies the flow for affinity; stable across processes + Key uint64 + // HasKey is false when the caller had nothing to derive a key from + HasKey bool +} + +// Selector is a load-balancing strategy. An instance holds the strategy's own state, such as +// a rotation counter, and belongs to one Balancer. +type Selector interface { + // Name identifies the strategy. + Name() string + // Needs declares what the strategy consumes. + Needs() Needs + // Prepare runs off the selection path, once per snapshot, and returns the strategy's + // precomputed form of it. It is only called with a snapshot that has members. + Prepare(*Snapshot) Prepared +} + +// Prepared is a strategy bound to one immutable snapshot. +type Prepared interface { + // Select returns a member of the snapshot it was prepared from. It is the selection + // path: it must not lock, block or allocate, and must finish in a bounded number of steps. + Select(Flow) *Member +} + +// Picker is what every pick-one mechanism presents to every plane. +type Picker interface { + // Needs is the Needs of the strategy behind the picker. + Needs() Needs + // Pick commits one flow to a member, or returns false when no member is eligible. + Pick(Flow) (Pick, bool) +} + +// PickerProvider is implemented by a member payload that is itself balanced, such as a pool +// whose members are pools. It returns nil when the payload does not pick one member per flow. +type PickerProvider interface { + Picker() Picker +} + +// Outcome is how a committed flow ended, as far as its member is concerned. +type Outcome uint8 + +const ( + // OutcomeOK means the member did the work. + OutcomeOK Outcome = iota + // OutcomeFailed means the member was reached and then failed or answered badly. + OutcomeFailed + // OutcomeConnectFailed means the member could not be reached at all. + OutcomeConnectFailed + // OutcomeCanceled means the caller gave up first, which says nothing about the member. + OutcomeCanceled +) diff --git a/pkg/lb/stats.go b/pkg/lb/stats.go new file mode 100644 index 000000000..07493da75 --- /dev/null +++ b/pkg/lb/stats.go @@ -0,0 +1,101 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import ( + "math" + "sync/atomic" + "time" +) + +// Stats is a member's runtime state. It is held by pointer so it can outlive the Member that +// carries it, such as when a member is rebuilt with a new weight. +type Stats struct { + inflight atomic.Int64 + // float64 bits of an average in nanoseconds; last writer wins + latency atomic.Uint64 + // unix nanosecond time of the last latency sample + stamp atomic.Int64 + // consecutive failed outcomes, of either kind + fails atomic.Int32 + // consecutive failures to reach the member, which alone count toward ejection + connectFails atomic.Int32 + // unix nanosecond time until which the member is ejected; it outlives a change of pool + ejectedUntil atomic.Int64 +} + +// Ejected reports whether the member is out of selection at now for repeated connect failures. +func (s *Stats) Ejected(now time.Time) bool { + return s.ejectedUntil.Load() > now.UnixNano() +} + +// Inflight returns the units of work currently committed to the member. +func (s *Stats) Inflight() int64 { + return s.inflight.Load() +} + +// Latency returns the member's latency average, or 0 when it has no sample. +func (s *Stats) Latency() time.Duration { + return time.Duration(math.Float64frombits(s.latency.Load())) +} + +// LastSample returns when the latency average was last updated; zero when never. +func (s *Stats) LastSample() time.Time { + ns := s.stamp.Load() + if ns == 0 { + return time.Time{} + } + return time.Unix(0, ns) +} + +// Failures returns the member's count of consecutive failed outcomes. +func (s *Stats) Failures() int32 { + return s.fails.Load() +} + +// ConnectFailures returns the member's count of consecutive failures to reach it. +func (s *Stats) ConnectFailures() int32 { + return s.connectFails.Load() +} + +// Faded returns the latency average as it stands at now, having faded toward zero since the +// last sample: close to exponentially with time constant decay, at the cost of one division. +// A member that is given no work gets no fresh sample, so this is how one ranked behind its +// peers, on an old sample or a penalty, is eventually tried again. +func (s *Stats) Faded(now time.Time, decay time.Duration) time.Duration { + avg := math.Float64frombits(s.latency.Load()) + elapsed := now.UnixNano() - s.stamp.Load() + if avg == 0 || elapsed <= 0 || decay <= 0 { + return time.Duration(avg) + } + // the reciprocal of e^x's cubic expansion: positive, falling and continuous for x >= 0 + x := float64(elapsed) / float64(decay) + return time.Duration(avg / (1 + x*(1+x*(0.5+x/6)))) +} + +// observe folds a latency sample, in nanoseconds, into a peak average: a sample at or above +// the average replaces it at once; one below pulls it down, further the older the average is. +// Concurrent samples may overwrite one another, which a load signal tolerates. +func (s *Stats) observe(sample float64, now int64, decay float64) { + stamp := s.stamp.Load() + avg := math.Float64frombits(s.latency.Load()) + if stamp != 0 && sample < avg { + elapsed := float64(max(now-stamp, 0)) + sample += (avg - sample) * math.Exp(-elapsed/decay) + } + s.latency.Store(math.Float64bits(sample)) + s.stamp.Store(now) +} diff --git a/pkg/lb/stats_test.go b/pkg/lb/stats_test.go new file mode 100644 index 000000000..12a06fe65 --- /dev/null +++ b/pkg/lb/stats_test.go @@ -0,0 +1,198 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package lb + +import ( + "math" + "net/netip" + "sync/atomic" + "testing" + "time" +) + +func TestStatsObservePeakAndDecay(t *testing.T) { + const decay = float64(10 * time.Second) + ms := float64(time.Millisecond) + s := &Stats{} + t0 := int64(1_000_000_000) + s.observe(40*ms, t0, decay) + if s.Latency() != 40*time.Millisecond || s.LastSample().UnixNano() != t0 { + t.Fatalf("first sample: %v at %v", s.Latency(), s.LastSample()) + } + // a higher sample replaces the average at once + s.observe(90*ms, t0+1, decay) + if s.Latency() != 90*time.Millisecond { + t.Errorf("peak = %v", s.Latency()) + } + // a lower sample right away barely moves it; one decay later it has moved 1-1/e of the way + s.observe(10*ms, t0+2, decay) + if got := s.Latency(); got < 89*time.Millisecond { + t.Errorf("an immediate low sample moved the average to %v", got) + } + s.observe(10*ms, t0+2+int64(decay), decay) + got := float64(s.Latency()) + if expect := 10*ms + 80*ms/math.E; math.Abs(got-expect) > ms { + t.Errorf("after one decay = %v, want about %v", time.Duration(got), time.Duration(expect)) + } + // sub-millisecond samples keep their resolution + fast := &Stats{} + fast.observe(250_000, t0, decay) + if fast.Latency() != 250*time.Microsecond { + t.Errorf("sub-millisecond sample = %v", fast.Latency()) + } + // a clock that steps backwards does not inflate the average + fast.observe(100_000, t0-5, decay) + if fast.Latency() > 250*time.Microsecond { + t.Errorf("after a backwards step = %v", fast.Latency()) + } +} + +func TestStatsFaded(t *testing.T) { + s := &Stats{} + now := time.Unix(100, 0) + if s.Faded(now, time.Second) != 0 { + t.Error("no sample fades to something") + } + s.observe(float64(8*time.Second), now.UnixNano(), float64(time.Second)) + for name, got := range map[string]time.Duration{ + "at the sample time": s.Faded(now, 10*time.Second), + "before the sample time": s.Faded(now.Add(-time.Hour), 10*time.Second), + "with no decay": s.Faded(now.Add(time.Minute), 0), + } { + if got != 8*time.Second { + t.Errorf("%s = %v", name, got) + } + } + // one decay later it is near 1/e; it keeps falling, and never turns negative + got := s.Faded(now.Add(10*time.Second), 10*time.Second) + eightSeconds := float64(8 * time.Second) + if want := time.Duration(eightSeconds / math.E); got < want || got > want+200*time.Millisecond { + t.Errorf("one decay later = %v, want a little above %v", got, want) + } + last := got + for _, after := range []time.Duration{20 * time.Second, time.Minute, time.Hour, 1000 * time.Hour} { + next := s.Faded(now.Add(after), 10*time.Second) + if next >= last || next < 0 { + t.Errorf("after %v = %v, not below %v", after, next, last) + } + last = next + } + if tenth := s.Faded(now.Add(80*time.Second), 10*time.Second); tenth > 80*time.Millisecond { + t.Errorf("a 8s penalty is still %v after eight decays", tenth) + } +} + +func TestHashes(t *testing.T) { + // golden values: a change here moves every client's affinity on every deployed replica + if got := HashString("tenant-42"); got != 0x5d8d50c585b0dfd7 { + t.Errorf("HashString = %#x", got) + } + if HashBytes([]byte("tenant-42")) != HashString("tenant-42") { + t.Error("HashBytes and HashString disagree") + } + if HashFold("API.Example.COM") != HashString("api.example.com") || HashFold("a") == HashFold("b") { + t.Error("HashFold does not fold ASCII case") + } + if Mix(1) == Mix(2) || Mix(0) != 0 { + t.Error("unexpected mix") + } + v4 := netip.MustParseAddr("192.0.2.7") + if HashAddr(v4, 64) != HashAddr(netip.MustParseAddr("::ffff:192.0.2.7"), 64) { + t.Error("an IPv4-mapped address keys differently from its IPv4 form") + } + if HashAddr(v4, 64) == HashAddr(netip.MustParseAddr("192.0.2.8"), 64) { + t.Error("distinct IPv4 addresses share a key") + } + a := netip.MustParseAddr("2001:db8:1:2:aaaa:bbbb:cccc:dddd") + b := netip.MustParseAddr("2001:db8:1:2:1111:2222:3333:4444") + c := netip.MustParseAddr("2001:db8:1:3:aaaa:bbbb:cccc:dddd") + if HashAddr(a, 64) != HashAddr(b, 64) { + t.Error("two addresses in one /64 key differently") + } + if HashAddr(a, 64) == HashAddr(c, 64) { + t.Error("two /64s share a key") + } + for _, whole := range []int{0, 128, 200, -1} { + if HashAddr(a, whole) == HashAddr(b, whole) { + t.Errorf("prefix %d did not keep the whole address", whole) + } + } + if allocs := testing.AllocsPerRun(100, func() { _ = HashAddr(a, 64) + HashFold("Host") }); allocs != 0 { + t.Errorf("hashing allocates %v", allocs) + } +} + +func FuzzHashBytes(f *testing.F) { + f.Add([]byte("client")) + f.Add([]byte{}) + f.Fuzz(func(t *testing.T, b []byte) { + if HashBytes(b) != HashString(string(b)) { + t.Errorf("HashBytes and HashString disagree on %q", b) + } + }) +} + +func TestLeast(t *testing.T) { + var pos atomic.Uint64 + if Least(nil, &pos, func(*Member) float64 { return 0 }) != nil { + t.Error("the least of nothing is something") + } + members := []*Member{ + NewMember(MemberOptions{Name: "a", Weight: 1}), + NewMember(MemberOptions{Name: "b", Weight: 3}), + NewMember(MemberOptions{Name: "c", Weight: 1}), + } + scores := map[*Member]float64{members[0]: 5, members[1]: 2, members[2]: 9} + score := func(m *Member) float64 { return scores[m] } + for range 5 { + if got := Least(members, &pos, score); got != members[1] { + t.Fatalf("least = %s", got.Name()) + } + } + // ties share the picks by weight, however many rounds are played + scores[members[2]] = 2 + counts := map[*Member]int{} + for range 40 { + counts[Least(members, &pos, score)]++ + } + if counts[members[0]] != 0 || counts[members[1]] != 30 || counts[members[2]] != 10 { + t.Errorf("tied picks = a:%d b:%d c:%d, want 0, 30, 10", + counts[members[0]], counts[members[1]], counts[members[2]]) + } + // a tied set that moves between the two passes still yields a member at the low score + var calls int + shifty := func(m *Member) float64 { + calls++ + if calls > len(members) && m != members[0] { + return 7 + } + return 1 + } + if got := Least(members, &pos, shifty); got != members[0] { + t.Errorf("after the tie dissolved = %s", got.Name()) + } + calls = 0 + vanishing := func(*Member) float64 { + calls++ + if calls > len(members) { + return 7 + } + return 1 + } + if got := Least(members, &pos, vanishing); got != members[0] { + t.Errorf("after every tie vanished = %s", got.Name()) + } +} diff --git a/pkg/lb/tier_test.go b/pkg/lb/tier_test.go new file mode 100644 index 000000000..77080b469 --- /dev/null +++ b/pkg/lb/tier_test.go @@ -0,0 +1,93 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package lb + +import ( + "slices" + "testing" +) + +func TestPoolSelectsFromTheLowestLiveTier(t *testing.T) { + health := map[string]*fakeHealth{} + member := func(name string, tier int) *Member { + health[name] = newHealth(1) + return NewMember(MemberOptions{Name: name, Tier: tier, Health: health[name]}) + } + obs := &recordingObserver{} + // a standby listed ahead of the members it stands by for is a standby all the same + p := mustPool(t, []*Member{ + member("standby", 1), member("a", 0), member("last-resort", 2), member("b", -3), + }, 1, PoolOptions{Observer: obs}) + expect := func(tier int, want ...string) { + t.Helper() + snap := p.Snapshot() + if got := names(snap); !slices.Equal(got, want) || snap.Tier != tier { + t.Fatalf("snapshot = tier %d %v, want tier %d %v", snap.Tier, got, tier, want) + } + events := obs.kinds(EventSnapshot) + if last := events[len(events)-1]; last.Tier != tier || last.Eligible != len(want) { + t.Fatalf("event = %+v", last) + } + } + if got := p.Configured()[3].Tier(); got != 0 { + t.Errorf("a negative tier = %d, want 0", got) + } + expect(0, "a", "b") + health["a"].set(-1) + expect(0, "b") + health["b"].set(-1) + expect(1, "standby") + health["standby"].set(-1) + expect(2, "last-resort") + health["last-resort"].set(-1) + expect(0) + // one primary returning takes every flow back from the standbys + health["standby"].set(1) + health["last-resort"].set(1) + expect(1, "standby") + health["b"].set(1) + expect(0, "b") +} + +type headSelector struct{} + +type headOf []*Member + +func (headSelector) Name() string { return "head" } + +func (headSelector) Needs() Needs { return 0 } + +func (headSelector) Prepare(snap *Snapshot) Prepared { return headOf(snap.Members) } + +func (h headOf) Select(Flow) *Member { return h[0] } + +func TestEjectingTheOnlyPrimaryFailsOver(t *testing.T) { + primary := NewMember(MemberOptions{Name: "primary"}) + standby := NewMember(MemberOptions{Name: "standby", Tier: 1}) + p := mustPool(t, []*Member{primary, standby}, 0) + b := NewBalancer(headSelector{}, BalancerOptions{ + Pool: p, Ejection: EjectionOptions{Failures: 1, MaxPercent: 50}, + }) + pk, ok := b.Pick(Flow{}) + if !ok || pk.Member() != primary { + t.Fatalf("pick = %v, %v", pk.Member(), ok) + } + pk.Done(OutcomeConnectFailed) + if got := names(p.Snapshot()); !slices.Equal(got, []string{"standby"}) { + t.Fatalf("after ejection = %v", got) + } +} diff --git a/pkg/observability/metrics/metrics.go b/pkg/observability/metrics/metrics.go index a7c2a735f..82c1aebb5 100644 --- a/pkg/observability/metrics/metrics.go +++ b/pkg/observability/metrics/metrics.go @@ -225,6 +225,42 @@ var ( []string{keys.Listener_Name, keys.Reason}, ) + // ProxyStreamMemberConnections counts the connections and UDP sessions a stream listener + // committed to a load balancer pool member, by how they went + ProxyStreamMemberConnections = prometheus.NewCounterVec( + prometheus.CounterOpts{ + Namespace: metricNamespace, + Subsystem: proxySubsystem, + Name: "stream_member_connections_total", + Help: "Count of connections and UDP sessions committed to a load balancer pool member, by result", + }, + []string{keys.Listener_Name, keys.Protocol, keys.Backend_Name, keys.Result}, + ) + + // ProxyStreamMemberActiveConnections gauges the connections and UDP sessions open to a + // load balancer pool member + ProxyStreamMemberActiveConnections = prometheus.NewGaugeVec( + prometheus.GaugeOpts{ + Namespace: metricNamespace, + Subsystem: proxySubsystem, + Name: "stream_member_active_connections", + Help: "Number of connections and UDP sessions open to a load balancer pool member", + }, + []string{keys.Listener_Name, keys.Protocol, keys.Backend_Name}, + ) + + // ProxyStreamMemberConnectDuration observes how long connecting to a pool member took + ProxyStreamMemberConnectDuration = prometheus.NewHistogramVec( + prometheus.HistogramOpts{ + Namespace: metricNamespace, + Subsystem: proxySubsystem, + Name: "stream_member_connect_duration_seconds", + Help: "Time taken to connect to a load balancer pool member", + Buckets: defaultBuckets, + }, + []string{keys.Listener_Name, keys.Protocol, keys.Backend_Name}, + ) + // ProxyStreamBytes counts the bytes stream listeners relayed, in from clients and out to them ProxyStreamBytes = prometheus.NewCounterVec( prometheus.CounterOpts{ @@ -820,16 +856,26 @@ var ( []string{keys.Mechanism, keys.Variant}, ) - // ALBPoolRefreshPanicRecovered counts recovered panics in ALB pool refresh - // worker goroutines (checkHealth, listenStatusUpdates). A dead worker leaves - // the healthy-target snapshot stale; the per-call re-filter in Targets() - // still produces correct dispatch, but operator-visible gauges drift. + // ALBMemberEjections counts pool members taken out of selection by passive health, which + // acts on repeated failures to connect rather than on a health check + ALBMemberEjections = prometheus.NewCounterVec( + prometheus.CounterOpts{ + Namespace: metricNamespace, + Subsystem: albSubsystem, + Name: "member_ejections_total", + Help: "Count of ALB pool members ejected by passive health after repeated connect failures.", + }, + []string{keys.ALB_Name, keys.Member}, + ) + + // ALBPoolRefreshPanicRecovered counts panics recovered while an ALB pool rebuilt its + // healthy-member snapshot; the pool keeps serving the snapshot it last published. ALBPoolRefreshPanicRecovered = prometheus.NewCounterVec( prometheus.CounterOpts{ Namespace: metricNamespace, Subsystem: albSubsystem, Name: "pool_refresh_panic_recovered_total", - Help: "Count of recovered panics in ALB pool refresh worker goroutines, by worker.", + Help: "Count of panics recovered while an ALB pool rebuilt its healthy-member snapshot, by worker.", }, []string{keys.Worker}, ) @@ -1021,6 +1067,18 @@ var ( }, []string{keys.Backend_Name}, ) + + // ALBPoolOnBackup flags ALB pools that have backup members and are dispatching to them + // because no other member is available. + ALBPoolOnBackup = prometheus.NewGaugeVec( + prometheus.GaugeOpts{ + Namespace: metricNamespace, + Subsystem: albSubsystem, + Name: "pool_on_backup", + Help: "1 while an ALB pool dispatches to its backup members because no other member is available; 0 otherwise.", + }, + []string{keys.Backend_Name}, + ) ) func init() { @@ -1031,6 +1089,10 @@ func init() { prometheus.MustRegister(ProxyStreamConnections) prometheus.MustRegister(ProxyStreamActiveConnections) prometheus.MustRegister(ProxyStreamBytes) + prometheus.MustRegister(ProxyStreamMemberConnections) + prometheus.MustRegister(ProxyStreamMemberActiveConnections) + prometheus.MustRegister(ProxyStreamMemberConnectDuration) + prometheus.MustRegister(ALBMemberEjections) prometheus.MustRegister(ProxyStreamDroppedDatagrams) prometheus.MustRegister(FrontendRequestStatus) prometheus.MustRegister(FrontendRequestDuration) @@ -1057,6 +1119,7 @@ func init() { prometheus.MustRegister(HealthcheckStatusNotifyPanicRecovered) prometheus.MustRegister(ALBPoolAdmitsFailing) prometheus.MustRegister(ALBPoolFloorReset) + prometheus.MustRegister(ALBPoolOnBackup) prometheus.MustRegister(CacheObjectOperations) prometheus.MustRegister(CacheObjectOperationDuration) prometheus.MustRegister(CacheByteOperations) @@ -1149,6 +1212,10 @@ var backendSeriesVecs = []partialDeleter{ HealthcheckStatusNotifyPanicRecovered, ALBPoolAdmitsFailing, ALBPoolFloorReset, + ALBPoolOnBackup, + ProxyStreamMemberConnections, + ProxyStreamMemberActiveConnections, + ProxyStreamMemberConnectDuration, } // DeleteBackendSeries removes every metric series labeled with the provided @@ -1162,6 +1229,8 @@ func DeleteBackendSeries(backendName string) { for _, v := range backendSeriesVecs { v.DeletePartialMatch(labels) } + // a pool member is named by the member label where the backend_name is its ALB's + ALBMemberEjections.DeletePartialMatch(prometheus.Labels{keys.Member: backendName}) } // ALB Autodiscovery metrics diff --git a/pkg/proxy/handlers/health/health_status_test.go b/pkg/proxy/handlers/health/health_status_test.go index 9f2d6e98b..d48eda564 100644 --- a/pkg/proxy/handlers/health/health_status_test.go +++ b/pkg/proxy/handlers/health/health_status_test.go @@ -422,7 +422,7 @@ func TestUpdateStatusTextEdgeCases(t *testing.T) { albOpts.Provider = providers.ALB albOpts.ALBOptions = ao.New() albOpts.ALBOptions.MechanismName = names.MechanismRR - albOpts.ALBOptions.Pool = ao.Members("down-only", "down-only") + albOpts.ALBOptions.Pool = ao.Members("down-only") downOnly := healthcheck.NewStatus("down-only", providers.Prometheus, "", healthcheck.StatusFailing, now().Add(-time.Minute), nil) diff --git a/pkg/proxy/hostnames/hostnames.go b/pkg/proxy/hostnames/hostnames.go index 944a1c554..2e189020d 100644 --- a/pkg/proxy/hostnames/hostnames.go +++ b/pkg/proxy/hostnames/hostnames.go @@ -22,6 +22,7 @@ package hostnames import ( "errors" + "net" "strings" ) @@ -107,3 +108,17 @@ func ToAnyDepth(h string) string { } return h } + +// reservedTLD is the top-level domain reserved never to resolve +const reservedTLD = ".invalid" + +// Reserved reports whether a host, or the host of a host:port address, is under the reserved +// .invalid domain and so can never be resolved or connected to. +func Reserved(addr string) bool { + host, _, err := net.SplitHostPort(addr) + if err != nil { + host = addr + } + host = strings.TrimSuffix(strings.ToLower(host), ".") + return strings.HasSuffix(host, reservedTLD) +} diff --git a/pkg/proxy/hostnames/hostnames_test.go b/pkg/proxy/hostnames/hostnames_test.go index af6d4e33c..17b857ee0 100644 --- a/pkg/proxy/hostnames/hostnames_test.go +++ b/pkg/proxy/hostnames/hostnames_test.go @@ -72,3 +72,16 @@ func TestClassification(t *testing.T) { require.Equal(t, "**.example.com", ToAnyDepth("**.example.com")) require.Equal(t, "example.com", ToAnyDepth("example.com")) } + +func TestReserved(t *testing.T) { + for _, addr := range []string{"unresolved.kgw.invalid:1", "x.INVALID.:9", "x.invalid", "[x.invalid]:5"} { + if !Reserved(addr) { + t.Errorf("Reserved(%q) = false", addr) + } + } + for _, addr := range []string{"invalid.example.com:1", "10.0.0.1:1", "", "invalid", "[::1]:80"} { + if Reserved(addr) { + t.Errorf("Reserved(%q) = true", addr) + } + } +} diff --git a/pkg/proxy/l4/feedback_test.go b/pkg/proxy/l4/feedback_test.go new file mode 100644 index 000000000..b3341b124 --- /dev/null +++ b/pkg/proxy/l4/feedback_test.go @@ -0,0 +1,360 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package l4 + +import ( + "crypto/tls" + "go/parser" + "go/token" + "net" + "net/netip" + "os" + "slices" + "strconv" + "strings" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" +) + +// the relay tells a route how its dial went and, if it connected, when the connection ended +func TestServerReportsToItsRoute(t *testing.T) { + up := rotate(echoServer(t, "echo:", nil)) + _, addr := startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + conn := dialTCP(t, addr) + if got := exchange(t, conn, "hi"); got != "echo:hi" { + t.Fatalf("reply = %q", got) + } + flows, routes := up.seen() + if len(flows) != 1 || len(routes) != 1 { + t.Fatalf("%d flows, %d routes", len(flows), len(routes)) + } + local := netip.MustParseAddrPort(conn.LocalAddr().String()) + if flows[0].Protocol != ProtocolTCP || flows[0].Client != local || flows[0].ServerName != "" { + t.Errorf("flow = %+v, want the client %v", flows[0], local) + } + r := routes[0] + if r.dialed.Load() != 1 || r.failed() || r.dialTook.Load() <= 0 { + t.Errorf("dial report: %d calls, failed %v, took %d", r.dialed.Load(), r.failed(), r.dialTook.Load()) + } + if r.closed.Load() != 0 { + t.Error("the route was closed while its connection is open") + } + _ = conn.Close() + waitFor(t, func() bool { return r.closed.Load() == 1 }) +} + +func TestServerReportsAFailedDialAndNothingAfter(t *testing.T) { + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + dead := ln.Addr().String() + _ = ln.Close() + up := rotate(dead) + srv, addr := startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + expectClosed(t, dialTCP(t, addr)) + waitFor(t, func() bool { return srv.ActiveConnections() == 0 }) + _, routes := up.seen() + if len(routes) != 1 || routes[0].dialed.Load() != 1 || !routes[0].failed() { + t.Fatalf("a failed dial was not reported: %d routes", len(routes)) + } + if routes[0].closed.Load() != 0 { + t.Error("a route whose dial failed was also closed") + } +} + +func TestServerOffersTheServerNameToItsUpstream(t *testing.T) { + cert := selfSigned(t, "shop.example.com") + up := rotate(echoServer(t, "tls:", &tls.Config{Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12})) + _, addr := startServer(t, ProtocolTLS, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + conn, err := tls.DialWithDialer(&net.Dialer{Timeout: 3 * time.Second}, "tcp", addr, + &tls.Config{ServerName: "shop.example.com", InsecureSkipVerify: true}) // #nosec G402 -- a test certificate + if err != nil { + t.Fatal(err) + } + defer conn.Close() + if got := exchange(t, conn, "hi"); got != "tls:hi" { + t.Fatalf("reply = %q", got) + } + flows, _ := up.seen() + if len(flows) != 1 || flows[0].Protocol != ProtocolTLS || flows[0].ServerName != "shop.example.com" { + t.Errorf("flows = %+v", flows) + } +} + +func TestPacketServerReportsToItsRoute(t *testing.T) { + up := rotate(udpEcho(t, "echo:")) + srv, addr, _ := startPacketServer(t, &Config{ + Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{IdleTimeout: timeconv.Duration(150 * time.Millisecond)}, + }) + client := udpClient(t, addr) + if got := datagram(t, client, "one"); got != "echo:one" { + t.Fatalf("reply = %q", got) + } + _ = datagram(t, client, "two") + flows, routes := up.seen() + if len(flows) != 1 || len(routes) != 1 { + t.Fatalf("a session of two datagrams made %d picks", len(flows)) + } + local := netip.MustParseAddrPort(client.LocalAddr().String()) + if flows[0].Protocol != ProtocolUDP || flows[0].Client.Port() != local.Port() { + t.Errorf("flow = %+v, want the client port %d", flows[0], local.Port()) + } + if routes[0].dialed.Load() != 1 || routes[0].failed() { + t.Error("the session's dial was not reported as a success") + } + waitFor(t, func() bool { return srv.ActiveSessions() == 0 }) + waitFor(t, func() bool { return routes[0].closed.Load() == 1 }) +} + +// the relay engine knows nothing of backends, load balancing or metrics: what it needs from +// them it declares as Upstream, Route and Observer, and they adapt to it +func TestImportBoundary(t *testing.T) { + const module = "github.com/trickstercache/trickster/v2/pkg/" + forbidden := []string{module + "backends", module + "lb"} + // tests may log; the relay itself may not + forbiddenOutsideTests := []string{module + "observability"} + entries, err := os.ReadDir(".") + if err != nil { + t.Fatal(err) + } + fset := token.NewFileSet() + var files int + for _, e := range entries { + if e.IsDir() || !strings.HasSuffix(e.Name(), ".go") { + continue + } + f, err := parser.ParseFile(fset, e.Name(), nil, parser.ImportsOnly) + if err != nil { + t.Fatal(err) + } + files++ + for _, imp := range f.Imports { + name, _ := strconv.Unquote(imp.Path.Value) + banned := forbidden + if !strings.HasSuffix(e.Name(), "_test.go") { + banned = append(slices.Clone(forbidden), forbiddenOutsideTests...) + } + for _, bad := range banned { + if name == bad || strings.HasPrefix(name, bad+"/") { + t.Errorf("%s imports %s: the relay must not depend on backends, the balancer or the metrics; "+ + "they adapt to it through Upstream and Observer", e.Name(), name) + } + } + } + } + if files == 0 { + t.Fatal("no source files were checked") + } +} + +// the relay reports its own events to its listener's observer, and to nothing else +func TestRelayReportsToItsObserver(t *testing.T) { + counts := &countingObserver{} + srv, addr := startServer(t, ProtocolTCP, &Config{ + Observer: counts, Table: tableOf(t, map[string]Upstream{"": rotate("", echoServer(t, "echo:", nil))}), + }) + conn := dialTCP(t, addr) + if got := exchange(t, conn, "hello"); got != "echo:hello" { + t.Fatalf("reply = %q", got) + } + if active, _, _ := counts.snapshot(); active != 1 { + t.Errorf("active while a connection is open = %d", active) + } + _ = conn.Close() + expectClosed(t, dialTCP(t, addr)) + waitFor(t, func() bool { return srv.ActiveConnections() == 0 }) + waitFor(t, func() bool { active, _, _ := counts.snapshot(); return active == 0 }) + _, results, bytes := counts.snapshot() + if results[ResultProxied] != 1 || results[ResultNoUpstream] != 1 { + t.Errorf("results = %v", results) + } + if bytes[DirectionIn] != int64(len("hello\n")) || bytes[DirectionOut] != int64(len("echo:hello\n")) { + t.Errorf("bytes = %v", bytes) + } + + udpCounts := &countingObserver{} + _, udpAddr, _ := startPacketServer(t, &Config{ + Observer: udpCounts, Table: tableOf(t, map[string]Upstream{"": rotate(udpEcho(t, "echo:"))}), + Options: &options.Options{IdleTimeout: timeconv.Duration(100 * time.Millisecond)}, + }) + if got := datagram(t, udpClient(t, udpAddr), "ping"); got != "echo:ping" { + t.Fatalf("reply = %q", got) + } + waitFor(t, func() bool { active, _, _ := udpCounts.snapshot(); return active == 0 }) + _, results, bytes = udpCounts.snapshot() + if results[ResultProxied] != 1 || bytes[DirectionIn] != 4 || bytes[DirectionOut] != 9 { + t.Errorf("udp results = %v, bytes = %v", results, bytes) + } + // with no observer the relay still works; that is what every other test here runs with + var none *Config + none.observer().Opened() + none.observer().Ended() + none.observer().Result(ResultProxied) + none.observer().Bytes(DirectionIn, 1) + none.observer().Dropped(DropQueueFull) +} + +func deadAddr(t *testing.T) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + return ln.Addr().String() +} + +// the first byte the upstream sends is reported once, however much follows +func TestServerReportsTheFirstUpstreamByte(t *testing.T) { + up := rotate(echoServer(t, "echo:", nil)) + _, addr := startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + conn := dialTCP(t, addr) + _, routes := func() ([]Flow, []*recordedRoute) { + waitFor(t, func() bool { _, r := up.seen(); return len(r) == 1 }) + return up.seen() + }() + if routes[0].firstByte.Load() != 0 { + t.Error("a first byte was reported before the upstream sent one") + } + for _, line := range []string{"one", "two", "three"} { + if got := exchange(t, conn, line); got != "echo:"+line { + t.Fatalf("reply = %q", got) + } + } + if got := routes[0].firstByte.Load(); got != 1 { + t.Errorf("first byte reported %d times", got) + } + flows, _ := up.seen() + if flows[0].Listener != "test" { + t.Errorf("flow listener = %q", flows[0].Listener) + } +} + +// a failed dial moves to the next route the upstream offers, inside one connect timeout, and +// every route tried hears how its dial went +func TestServerRetriesAFailedDial(t *testing.T) { + dead := deadAddr(t) + live := echoServer(t, "echo:", nil) + up := retrying{rotate(live, dead, dead)} + up.retries = 2 + _, addr := startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + if got := exchange(t, dialTCP(t, addr), "hi"); got != "echo:hi" { + t.Fatalf("reply = %q", got) + } + _, routes := up.seen() + if len(routes) != 3 || !routes[0].failed() || !routes[1].failed() || routes[2].failed() { + t.Fatalf("%d routes tried; want two failed dials and then a success", len(routes)) + } + if routes[0].closed.Load() != 0 || routes[1].closed.Load() != 0 { + t.Error("a route whose dial failed was also closed") + } + + // retries are bounded by what the upstream offers + none := retrying{rotate(dead)} + _, addr = startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": none})}) + expectClosed(t, dialTCP(t, addr)) + waitFor(t, func() bool { _, r := none.seen(); return len(r) == 1 && r[0].failed() }) + + // a final route is never retried, whatever the upstream could offer + final := retrying{rotate(live, dead)} + final.retries, final.final = 5, true + _, addr = startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": final})}) + expectClosed(t, dialTCP(t, addr)) + waitFor(t, func() bool { _, r := final.seen(); return len(r) == 1 && r[0].failed() }) + time.Sleep(50 * time.Millisecond) + if _, r := final.seen(); len(r) != 1 { + t.Errorf("a final route was retried: %d routes", len(r)) + } + + // an upstream that runs out of routes mid-retry ends the connection + refusing := retrying{rotate("", dead)} + refusing.retries = 3 + _, addr = startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": refusing})}) + expectClosed(t, dialTCP(t, addr)) +} + +// the connect timeout covers every attempt together, not each one +func TestServerRetriesShareOneConnectTimeout(t *testing.T) { + // an address that accepts nothing and refuses nothing: the dial hangs until it times out + blackhole := "192.0.2.1:9" + up := retrying{rotate(blackhole, blackhole, blackhole, blackhole)} + up.retries = 10 + _, addr := startServer(t, ProtocolTCP, &Config{ + Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{ConnectTimeout: timeconv.Duration(150 * time.Millisecond)}, + }) + began := time.Now() + expectClosed(t, dialTCP(t, addr)) + if took := time.Since(began); took > 2*time.Second { + t.Errorf("retries ran for %v against a 150ms connect timeout", took) + } + if _, r := up.seen(); len(r) > 2 { + t.Errorf("%d routes were dialed inside one connect timeout", len(r)) + } +} + +// a udp upstream with nothing listening answers a datagram with a port-unreachable, which is +// the only sign it is down; the session's route hears of it when the session ends +func TestPacketServerReportsARefusingUpstream(t *testing.T) { + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + gone := pc.LocalAddr().String() + _ = pc.Close() + up := rotate(gone) + srv, addr, _ := startPacketServer(t, &Config{ + Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{IdleTimeout: timeconv.Duration(300 * time.Millisecond)}, + }) + client := udpClient(t, addr) + for range 3 { + _, _ = client.Write([]byte("anyone?")) + time.Sleep(20 * time.Millisecond) + } + waitFor(t, func() bool { return srv.ActiveSessions() == 0 }) + // a session that ends on a refusal is gone, so a later datagram may have opened another + _, routes := up.seen() + if len(routes) == 0 || routes[0].closed.Load() != 1 { + t.Fatalf("%d routes", len(routes)) + } + if routes[0].closeErr.Load() == nil { + t.Skip("this platform did not report the refused datagram on the connected socket") + } + if routes[0].firstByte.Load() != 0 { + t.Error("a first reply was reported from an upstream that never answered") + } +} + +func TestPacketServerReportsTheFirstReply(t *testing.T) { + up := rotate(udpEcho(t, "echo:")) + _, addr, _ := startPacketServer(t, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + client := udpClient(t, addr) + for _, msg := range []string{"one", "two", "three"} { + if got := datagram(t, client, msg); got != "echo:"+msg { + t.Fatalf("reply = %q", got) + } + } + _, routes := up.seen() + if len(routes) != 1 || routes[0].firstByte.Load() != 1 { + t.Errorf("first reply reported %d times over a session of three", routes[0].firstByte.Load()) + } +} diff --git a/pkg/proxy/l4/observe/observe.go b/pkg/proxy/l4/observe/observe.go new file mode 100644 index 000000000..630456acb --- /dev/null +++ b/pkg/proxy/l4/observe/observe.go @@ -0,0 +1,62 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +// Package observe binds the stream relay's events to Trickster's metrics. +package observe + +import ( + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" + + "github.com/prometheus/client_golang/prometheus" +) + +// listener meters one stream listener. The series touched per connection and per datagram +// are resolved once, here, so the relay's hot path is a plain add with no label lookup. +type listener struct { + name, protocol string + active prometheus.Gauge + in, out prometheus.Counter +} + +// Listener returns the observer for the named stream listener. +func Listener(name, protocol string) l4.Observer { + return &listener{ + name: name, protocol: protocol, + active: metrics.ProxyStreamActiveConnections.WithLabelValues(name, protocol), + in: metrics.ProxyStreamBytes.WithLabelValues(name, protocol, l4.DirectionIn), + out: metrics.ProxyStreamBytes.WithLabelValues(name, protocol, l4.DirectionOut), + } +} + +func (l *listener) Opened() { l.active.Inc() } + +func (l *listener) Ended() { l.active.Dec() } + +func (l *listener) Result(result string) { + metrics.ProxyStreamConnections.WithLabelValues(l.name, l.protocol, result).Inc() +} + +func (l *listener) Bytes(direction string, n int64) { + if direction == l4.DirectionIn { + l.in.Add(float64(n)) + return + } + l.out.Add(float64(n)) +} + +func (l *listener) Dropped(reason string) { + metrics.ProxyStreamDroppedDatagrams.WithLabelValues(l.name, reason).Inc() +} diff --git a/pkg/proxy/l4/observe/observe_test.go b/pkg/proxy/l4/observe/observe_test.go new file mode 100644 index 000000000..95ba8bee5 --- /dev/null +++ b/pkg/proxy/l4/observe/observe_test.go @@ -0,0 +1,55 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package observe + +import ( + "testing" + + "github.com/trickstercache/trickster/v2/pkg/observability/metrics" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4" + + "github.com/prometheus/client_golang/prometheus/testutil" +) + +func TestListenerMetersTheRelay(t *testing.T) { + const name = "observe-test" + o := Listener(name, l4.ProtocolUDP) + active := metrics.ProxyStreamActiveConnections.WithLabelValues(name, l4.ProtocolUDP) + in := metrics.ProxyStreamBytes.WithLabelValues(name, l4.ProtocolUDP, l4.DirectionIn) + out := metrics.ProxyStreamBytes.WithLabelValues(name, l4.ProtocolUDP, l4.DirectionOut) + proxied := metrics.ProxyStreamConnections.WithLabelValues(name, l4.ProtocolUDP, l4.ResultProxied) + dropped := metrics.ProxyStreamDroppedDatagrams.WithLabelValues(name, l4.DropQueueFull) + + o.Opened() + o.Opened() + o.Ended() + o.Result(l4.ResultProxied) + o.Bytes(l4.DirectionIn, 100) + o.Bytes(l4.DirectionIn, 20) + o.Bytes(l4.DirectionOut, 7) + o.Dropped(l4.DropQueueFull) + for series, want := range map[string][2]float64{ + "active": {testutil.ToFloat64(active), 1}, + "in": {testutil.ToFloat64(in), 120}, + "out": {testutil.ToFloat64(out), 7}, + "proxied": {testutil.ToFloat64(proxied), 1}, + "dropped": {testutil.ToFloat64(dropped), 1}, + } { + if want[0] != want[1] { + t.Errorf("%s = %v, want %v", series, want[0], want[1]) + } + } +} diff --git a/pkg/proxy/l4/observer.go b/pkg/proxy/l4/observer.go new file mode 100644 index 000000000..bce8c8e30 --- /dev/null +++ b/pkg/proxy/l4/observer.go @@ -0,0 +1,43 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package l4 + +// Observer receives one listener's relay events. The relay itself neither logs nor meters; +// an Observer is bound to its listener, so it can resolve whatever it counts with ahead of +// time and leave the per-datagram path a plain add. +type Observer interface { + // Opened reports a connection or session admitted; Ended reports that it is over. + Opened() + Ended() + // Result reports how a connection or session was disposed of: one of the Result consts. + Result(result string) + // Bytes reports payload relayed in a direction: DirectionIn or DirectionOut. + Bytes(direction string, n int64) + // Dropped reports a datagram dropped, by reason: one of the Drop consts. + Dropped(reason string) +} + +type nopObserver struct{} + +func (nopObserver) Opened() {} + +func (nopObserver) Ended() {} + +func (nopObserver) Result(string) {} + +func (nopObserver) Bytes(string, int64) {} + +func (nopObserver) Dropped(string) {} diff --git a/pkg/proxy/l4/proxyheader_test.go b/pkg/proxy/l4/proxyheader_test.go new file mode 100644 index 000000000..0c0446b6f --- /dev/null +++ b/pkg/proxy/l4/proxyheader_test.go @@ -0,0 +1,88 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package l4 + +import ( + "crypto/tls" + "net" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" +) + +// headerListener accepts connections that arrived behind a PROXY protocol header, as the +// daemon's listener hands them over +type headerListener struct{ net.Listener } + +type headerConn struct{ net.Conn } + +func (l headerListener) Accept() (net.Conn, error) { + c, err := l.Listener.Accept() + if err != nil { + return nil, err + } + return headerConn{c}, nil +} + +func (headerConn) ProxyTLV(typ byte) ([]byte, bool) { + return []byte{typ, typ}, typ == 0xEA +} + +func TestFlowCarriesTheProxyHeader(t *testing.T) { + cert := selfSigned(t, "shop.example.com") + tlsConf := &tls.Config{Certificates: []tls.Certificate{cert}, MinVersion: tls.VersionTLS12} + for protocol, upstreamTLS := range map[string]*tls.Config{ProtocolTCP: nil, ProtocolTLS: tlsConf} { + up := rotate(echoServer(t, "echo:", upstreamTLS)) + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + srv := NewServer("test", protocol, &Config{ + Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{ConnectTimeout: timeconv.Duration(time.Second)}, + }) + go func() { _ = srv.Serve(headerListener{ln}) }() + t.Cleanup(func() { _ = srv.Close() }) + conn := dialTCP(t, ln.Addr().String()) + if upstreamTLS != nil { + conn = tls.Client(conn, &tls.Config{ServerName: "shop.example.com", InsecureSkipVerify: true}) // #nosec G402 + } + if got := exchange(t, conn, "hi"); got != "echo:hi" { + t.Fatalf("%s: reply = %q", protocol, got) + } + up.mu.Lock() + flow := up.flows[0] + up.mu.Unlock() + if flow.Proxy == nil { + t.Fatalf("%s: the flow lost the connection's PROXY header", protocol) + } + if v, ok := flow.Proxy.ProxyTLV(0xEA); !ok || len(v) != 2 { + t.Errorf("%s: TLV = %v, %v", protocol, v, ok) + } + } + // a plain connection has no header to offer + up := rotate(echoServer(t, "echo:", nil)) + _, addr := startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + exchange(t, dialTCP(t, addr), "hi") + up.mu.Lock() + defer up.mu.Unlock() + if up.flows[0].Proxy != nil { + t.Error("a connection without a PROXY header produced one") + } +} diff --git a/pkg/proxy/l4/reload_characterization_test.go b/pkg/proxy/l4/reload_characterization_test.go new file mode 100644 index 000000000..bae4d2132 --- /dev/null +++ b/pkg/proxy/l4/reload_characterization_test.go @@ -0,0 +1,68 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package l4 + +import ( + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" +) + +// a reload routes new connections only: one already relayed stays on the member it committed to +func TestServerUpdateRoutesNewConnectionsOnly(t *testing.T) { + before, after := echoServer(t, "before:", nil), echoServer(t, "after:", nil) + srv, addr := startServer(t, ProtocolTCP, + &Config{Table: tableOf(t, map[string]Upstream{"": Static(before)})}) + held := dialTCP(t, addr) + if got := exchange(t, held, "one"); got != "before:one" { + t.Fatalf("reply = %q", got) + } + srv.Update(&Config{Table: tableOf(t, map[string]Upstream{"": Static(after)})}) + if got := exchange(t, dialTCP(t, addr), "new"); got != "after:new" { + t.Errorf("a connection accepted after the update = %q", got) + } + if got := exchange(t, held, "two"); got != "before:two" { + t.Errorf("a connection relayed before the update = %q", got) + } + // an update that routes nothing refuses new connections and still leaves the held one alone + srv.Update(nil) + expectClosed(t, dialTCP(t, addr)) + if got := exchange(t, held, "three"); got != "before:three" { + t.Errorf("a held connection after an empty update = %q", got) + } +} + +func TestPacketServerUpdateRoutesNewSessionsOnly(t *testing.T) { + before, after := udpEcho(t, "before:"), udpEcho(t, "after:") + idle := &options.Options{IdleTimeout: timeconv.Duration(5 * time.Second)} + srv, addr, _ := startPacketServer(t, &Config{ + Table: tableOf(t, map[string]Upstream{"": Static(before)}), Options: idle, + }) + held := udpClient(t, addr) + if got := datagram(t, held, "one"); got != "before:one" { + t.Fatalf("reply = %q", got) + } + srv.Update(&Config{Table: tableOf(t, map[string]Upstream{"": Static(after)}), Options: idle}) + if got := datagram(t, udpClient(t, addr), "new"); got != "after:new" { + t.Errorf("a session opened after the update = %q", got) + } + if got := datagram(t, held, "two"); got != "before:two" { + t.Errorf("a session opened before the update = %q", got) + } +} diff --git a/pkg/proxy/l4/routes_test.go b/pkg/proxy/l4/routes_test.go new file mode 100644 index 000000000..6915c4ac7 --- /dev/null +++ b/pkg/proxy/l4/routes_test.go @@ -0,0 +1,174 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package l4 + +import ( + "maps" + "sync" + "sync/atomic" + "time" +) + +// rotating is a test upstream that commits flows to its addresses in turn and keeps what the +// relay reports of each. An empty or undialable address holds its turn and refuses it. +type rotating struct { + addrs []string + pos atomic.Uint64 + // retries is how many more routes a flow may be offered after a failed dial + retries int + final bool + + mu sync.Mutex + flows []Flow + routes []*recordedRoute +} + +type recordedRoute struct { + addr string + final bool + attempt int + dialed atomic.Int32 + dialErr atomic.Pointer[error] + dialTook atomic.Int64 + firstByte atomic.Int32 + closed atomic.Int32 + closeErr atomic.Pointer[error] +} + +func rotate(addrs ...string) *rotating { + return &rotating{addrs: addrs} +} + +func (u *rotating) Pick(f Flow) (Route, bool) { + u.mu.Lock() + u.flows = append(u.flows, f) + u.mu.Unlock() + if len(u.addrs) == 0 { + return nil, false + } + addr := u.addrs[u.pos.Add(1)%uint64(len(u.addrs))] + if addr == "" || Refusing(addr) { + return nil, false + } + r := &recordedRoute{addr: addr, final: u.final} + u.mu.Lock() + u.routes = append(u.routes, r) + u.mu.Unlock() + return r, true +} + +// retrying is a rotating upstream that offers the next address when a dial fails +type retrying struct{ *rotating } + +func (u retrying) Retry(f Flow, failed Route) (Route, bool) { + prev := failed.(*recordedRoute) + if prev.attempt >= u.retries { + return nil, false + } + next, ok := u.Pick(f) + if ok { + next.(*recordedRoute).attempt = prev.attempt + 1 + } + return next, ok +} + +func (u *rotating) seen() ([]Flow, []*recordedRoute) { + u.mu.Lock() + defer u.mu.Unlock() + return append([]Flow(nil), u.flows...), append([]*recordedRoute(nil), u.routes...) +} + +func (r *recordedRoute) Addr() string { return r.addr } + +func (r *recordedRoute) Final() bool { return r.final } + +func (r *recordedRoute) Dialed(d time.Duration, err error) { + r.dialed.Add(1) + r.dialTook.Store(int64(d)) + if err != nil { + r.dialErr.Store(&err) + } +} + +func (r *recordedRoute) FirstByte() { r.firstByte.Add(1) } + +func (r *recordedRoute) Closed(err error) { + r.closed.Add(1) + if err != nil { + r.closeErr.Store(&err) + } +} + +func (r *recordedRoute) failed() bool { return r.dialErr.Load() != nil } + +// countingObserver keeps what a relay reports of itself +type countingObserver struct { + mu sync.Mutex + active int + results map[string]int + drops map[string]int + bytes map[string]int64 +} + +func (o *countingObserver) add(m *map[string]int, key string) { + o.mu.Lock() + defer o.mu.Unlock() + if *m == nil { + *m = make(map[string]int) + } + (*m)[key]++ +} + +func (o *countingObserver) Opened() { + o.mu.Lock() + o.active++ + o.mu.Unlock() +} + +func (o *countingObserver) Ended() { + o.mu.Lock() + o.active-- + o.mu.Unlock() +} + +func (o *countingObserver) Result(r string) { o.add(&o.results, r) } + +func (o *countingObserver) Dropped(reason string) { o.add(&o.drops, reason) } + +func (o *countingObserver) Bytes(direction string, n int64) { + o.mu.Lock() + defer o.mu.Unlock() + if o.bytes == nil { + o.bytes = make(map[string]int64) + } + o.bytes[direction] += n +} + +func (o *countingObserver) dropped(reason string) float64 { + o.mu.Lock() + defer o.mu.Unlock() + return float64(o.drops[reason]) +} + +func (o *countingObserver) snapshot() (active int, results map[string]int, bytes map[string]int64) { + o.mu.Lock() + defer o.mu.Unlock() + results = make(map[string]int, len(o.results)) + maps.Copy(results, o.results) + bytes = make(map[string]int64, len(o.bytes)) + maps.Copy(bytes, o.bytes) + return o.active, results, bytes +} diff --git a/pkg/proxy/l4/server.go b/pkg/proxy/l4/server.go index cb9a1188c..0272a102c 100644 --- a/pkg/proxy/l4/server.go +++ b/pkg/proxy/l4/server.go @@ -25,7 +25,6 @@ import ( "sync/atomic" "time" - "github.com/trickstercache/trickster/v2/pkg/observability/metrics" "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" ) @@ -67,6 +66,15 @@ type Config struct { // MaxConnections bounds the connections relayed at once, or the UDP sessions open at once; // zero applies no bound to connections and the default bound to sessions MaxConnections int + // Observer receives the listener's events; nil discards them + Observer Observer +} + +func (c *Config) observer() Observer { + if c == nil || c.Observer == nil { + return nopObserver{} + } + return c.Observer } func (c *Config) table() *Table { @@ -194,7 +202,7 @@ func (s *Server) track(conn net.Conn) bool { } s.conns[conn] = struct{}{} s.workers.Add(1) - metrics.ProxyStreamActiveConnections.WithLabelValues(s.name, s.protocol).Inc() + s.cfg.Load().observer().Opened() return true } @@ -204,7 +212,7 @@ func (s *Server) untrack(conn net.Conn) { delete(s.conns, conn) s.slot.Broadcast() s.mu.Unlock() - metrics.ProxyStreamActiveConnections.WithLabelValues(s.name, s.protocol).Dec() + s.cfg.Load().observer().Ended() s.workers.Done() } @@ -228,13 +236,15 @@ func (s *Server) untrackUpstream(conn net.Conn) { } func (s *Server) result(r string) { - metrics.ProxyStreamConnections.WithLabelValues(s.name, s.protocol, r).Inc() + s.cfg.Load().observer().Result(r) } func (s *Server) handle(client net.Conn) { defer s.untrack(client) cfg := s.cfg.Load() opts := cfg.options() + // read before the connection is wrapped to replay what was peeked of it + proxy, _ := client.(ProxyHeader) var host string if s.protocol == ProtocolTLS { var err error @@ -249,17 +259,30 @@ func (s *Server) handle(client net.Conn) { s.result(ResultNoRoute) return } - addr, ok := up.Addr() - if !ok { - s.result(ResultNoUpstream) - return + flow := flowOf(s.name, s.protocol, client.RemoteAddr(), host) + flow.Proxy = proxy + var upstream net.Conn + var route Route + if racer, races := up.(Racer); races { + routes := racer.Race(flow) + if len(routes) == 0 { + s.result(ResultNoUpstream) + return + } + upstream, route = s.race(routes, opts.Connect()) + } else { + picked, ok := up.Pick(flow) + if !ok { + s.result(ResultNoUpstream) + return + } + upstream, route = s.dial(up, flow, picked, opts.Connect()) } - dialer := net.Dialer{Timeout: opts.Connect()} - upstream, err := dialer.DialContext(s.ctx, "tcp", addr) - if err != nil { + if upstream == nil { s.result(ResultDialFailed) return } + defer route.Closed(nil) if !s.trackUpstream(upstream) { _ = upstream.Close() s.result(ResultRefused) @@ -267,9 +290,94 @@ func (s *Server) handle(client net.Conn) { } defer s.untrackUpstream(upstream) s.result(ResultProxied) - in, out := relay(client, upstream, opts.Idle()) - metrics.ProxyStreamBytes.WithLabelValues(s.name, s.protocol, DirectionIn).Add(float64(in)) - metrics.ProxyStreamBytes.WithLabelValues(s.name, s.protocol, DirectionOut).Add(float64(out)) + in, out := relay(client, upstream, opts.Idle(), route.FirstByte) + obs := cfg.observer() + obs.Bytes(DirectionIn, in) + obs.Bytes(DirectionOut, out) +} + +// dial connects to the route, moving to another that the upstream offers when a dial fails, +// all within one connect timeout. It returns the connection with the route it was made on, +// or nil once no route is left; every route it tried has been told how its dial went. +func (s *Server) dial(up Upstream, flow Flow, route Route, timeout time.Duration) (net.Conn, Route) { + ctx, cancel := context.WithTimeout(s.ctx, timeout) + defer cancel() + var dialer net.Dialer + for { + began := time.Now() + conn, err := dialer.DialContext(ctx, "tcp", route.Addr()) + route.Dialed(time.Since(began), err) + if err == nil { + return conn, route + } + retrier, can := up.(Retrier) + if !can || route.Final() || ctx.Err() != nil { + return nil, nil + } + next, ok := retrier.Retry(flow, route) + if !ok { + return nil, nil + } + route = next + } +} + +// raceState is what the dials of one race share: the first to connect takes it +type raceState struct { + mu sync.Mutex + settled sync.Cond + pending int + conn net.Conn + route Route + took time.Duration +} + +// race dials every route at once within the connect timeout and returns the first connection +// made with its route, or nil when none connects. Every route is told how its dial went; one +// that lost the race, or was cut short by it, was abandoned. +func (s *Server) race(routes []Route, timeout time.Duration) (net.Conn, Route) { + ctx, cancel := context.WithTimeout(s.ctx, timeout) + // the winner ends the dials still in progress + defer cancel() + st := &raceState{pending: len(routes)} + st.settled.L = &st.mu + for _, route := range routes { + go func() { + var dialer net.Dialer + began := time.Now() + conn, err := dialer.DialContext(ctx, "tcp", route.Addr()) + took := time.Since(began) + st.mu.Lock() + won := err == nil && st.conn == nil + lost := st.conn != nil + if won { + st.conn, st.route, st.took = conn, route, took + } + st.pending-- + st.settled.Signal() + st.mu.Unlock() + switch { + case won: + // reported by the flow's own worker, ahead of anything else it reports + case lost: + if conn != nil { + _ = conn.Close() + } + route.Dialed(0, ErrAbandoned) + default: + route.Dialed(took, err) + } + }() + } + st.mu.Lock() + defer st.mu.Unlock() + for st.conn == nil && st.pending > 0 { + st.settled.Wait() + } + if st.conn != nil { + st.route.Dialed(st.took, nil) + } + return st.conn, st.route } // Shutdown stops accepting and waits for relayed connections to end until ctx is done. @@ -325,13 +433,13 @@ func (s *Server) ActiveConnections() int { // and to it; a side that ends cleanly half-closes its peer, and a side that fails closes both. // The idle timeout is the connection's: a direction that has read nothing for that long ends // the relay only when the other has moved nothing either. -func relay(client, upstream net.Conn, idle time.Duration) (int64, int64) { +func relay(client, upstream net.Conn, idle time.Duration, firstByte func()) (int64, int64) { var in, out int64 var last atomic.Int64 last.Store(time.Now().UnixNano()) var wg sync.WaitGroup - wg.Go(func() { in = pipe(upstream, client, idle, &last) }) - wg.Go(func() { out = pipe(client, upstream, idle, &last) }) + wg.Go(func() { in = pipe(upstream, client, idle, &last, nil) }) + wg.Go(func() { out = pipe(client, upstream, idle, &last, firstByte) }) wg.Wait() return in, out } @@ -341,7 +449,8 @@ var bufPool = sync.Pool{New: func() any { return &b }} -func pipe(dst, src net.Conn, idle time.Duration, last *atomic.Int64) int64 { +// pipe copies src to dst until src ends; first, when set, is called once, on the first read +func pipe(dst, src net.Conn, idle time.Duration, last *atomic.Int64, first func()) int64 { bp := bufPool.Get().(*[]byte) defer bufPool.Put(bp) buf := *bp @@ -352,6 +461,10 @@ func pipe(dst, src net.Conn, idle time.Duration, last *atomic.Int64) int64 { } nr, rerr := src.Read(buf) if nr > 0 { + if first != nil { + first() + first = nil + } now := time.Now() last.Store(now.UnixNano()) if idle > 0 { diff --git a/pkg/proxy/l4/server_test.go b/pkg/proxy/l4/server_test.go index 6ebeec697..53e9a5322 100644 --- a/pkg/proxy/l4/server_test.go +++ b/pkg/proxy/l4/server_test.go @@ -29,7 +29,6 @@ import ( "testing" "time" - "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" ) @@ -165,11 +164,8 @@ func TestServerRotatesAcrossPoolAndRefusesADeadMembersShare(t *testing.T) { } dead := deadLn.Addr().String() _ = deadLn.Close() - pl := pooledOf(t, "alb", - member(originBackend(t, "dead", dead), 1, healthcheck.StatusPassing), - member(originBackend(t, "a", a), 1, healthcheck.StatusPassing), - member(originBackend(t, "b", b), 1, healthcheck.StatusPassing)) - _, addr := startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": FromBackend(pl)})}) + up := rotate(dead, a, b) + _, addr := startServer(t, ProtocolTCP, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) seen := make(map[string]int) var refused int for range 6 { @@ -201,7 +197,7 @@ func TestServerRefusesWhatItCannotRoute(t *testing.T) { _ = deadLn.Close() for name, tbl := range map[string]*Table{ "empty": NewTable(), - "no_member": tableOf(t, map[string]Upstream{"": FromBackend(pooledOf(t, "none"))}), + "no_member": tableOf(t, map[string]Upstream{"": rotate()}), "dead": tableOf(t, map[string]Upstream{"": Static(dead)}), } { t.Run(name, func(t *testing.T) { @@ -384,7 +380,7 @@ func TestPipeClosesBothOnWriteFailure(t *testing.T) { _ = dstPeer.Close() go func() { _, _ = srcPeer.Write([]byte("data")) }() var last atomic.Int64 - if n := pipe(dst, src, 0, &last); n != 0 { + if n := pipe(dst, src, 0, &last, nil); n != 0 { t.Errorf("bytes written to a closed destination = %d", n) } if _, err := srcPeer.Write([]byte("more")); err == nil { diff --git a/pkg/proxy/l4/spread_test.go b/pkg/proxy/l4/spread_test.go new file mode 100644 index 000000000..04837f979 --- /dev/null +++ b/pkg/proxy/l4/spread_test.go @@ -0,0 +1,323 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package l4 + +import ( + "context" + "errors" + "net" + "strings" + "sync" + "testing" + "time" + + "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" + "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" +) + +// spreading is an upstream that commits each flow to every one of its addresses at once +type spreading struct { + addrs []string + + mu sync.Mutex + routes []*recordedRoute +} + +func (u *spreading) route(addr string) *recordedRoute { + r := &recordedRoute{addr: addr, final: true} + u.mu.Lock() + u.routes = append(u.routes, r) + u.mu.Unlock() + return r +} + +func (u *spreading) to() []string { + u.mu.Lock() + defer u.mu.Unlock() + return u.addrs +} + +func (u *spreading) empty() { + u.mu.Lock() + u.addrs = nil + u.mu.Unlock() +} + +func (u *spreading) Pick(Flow) (Route, bool) { + addrs := u.to() + if len(addrs) == 0 { + return nil, false + } + return u.route(addrs[0]), true +} + +func (u *spreading) Race(Flow) []Route { + addrs := u.to() + routes := make([]Route, len(addrs)) + for i, addr := range addrs { + routes[i] = u.route(addr) + } + return routes +} + +func (u *spreading) Mirror(Flow, Route) []Route { + addrs := u.to() + routes := make([]Route, 0, len(addrs)) + for _, addr := range addrs[1:] { + routes = append(routes, u.route(addr)) + } + return routes +} + +func (u *spreading) seen() []*recordedRoute { + u.mu.Lock() + defer u.mu.Unlock() + return append([]*recordedRoute(nil), u.routes...) +} + +// refusedAddr is an address nothing listens at +func refusedAddr(t *testing.T) string { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + addr := ln.Addr().String() + _ = ln.Close() + return addr +} + +func abandoned(r *recordedRoute) bool { + err := r.dialErr.Load() + return err != nil && errors.Is(*err, ErrAbandoned) +} + +func TestServerRacesItsRoutes(t *testing.T) { + counts := &countingObserver{} + up := &spreading{addrs: []string{refusedAddr(t), echoServer(t, "a:", nil), echoServer(t, "b:", nil)}} + srv, addr := startServer(t, ProtocolTCP, &Config{ + Observer: counts, Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{ConnectTimeout: timeconv.Duration(2 * time.Second)}, + }) + conn := dialTCP(t, addr) + reply := exchange(t, conn, "hi") + if reply != "a:hi" && reply != "b:hi" { + t.Fatalf("reply = %q", reply) + } + _ = conn.Close() + waitFor(t, func() bool { return srv.ActiveConnections() == 0 }) + routes := up.seen() + if len(routes) != 3 { + t.Fatalf("%d routes", len(routes)) + } + waitFor(t, func() bool { + return routes[0].dialed.Load() == 1 && routes[1].dialed.Load() == 1 && routes[2].dialed.Load() == 1 + }) + var won, lost int + for _, r := range routes[1:] { + switch { + case !r.failed(): + won++ + if r.closed.Load() != 1 || !strings.HasPrefix(reply, map[string]string{up.addrs[1]: "a:", up.addrs[2]: "b:"}[r.addr]) { + t.Errorf("the winner %s closed %d times behind reply %q", r.addr, r.closed.Load(), reply) + } + case abandoned(r): + lost++ + if r.closed.Load() != 0 { + t.Error("a route that lost the race was reported closed") + } + } + } + if won != 1 || lost != 1 { + t.Errorf("%d winners and %d abandoned of two live routes", won, lost) + } + // the route nothing listens at failed on its own account, or was cut short by the winner + if !routes[0].failed() || routes[0].closed.Load() != 0 { + t.Error("a refused route was not reported failed") + } + if _, results, _ := counts.snapshot(); results[ResultProxied] != 1 { + t.Errorf("results = %v", results) + } +} + +func TestServerRaceWithNoWinner(t *testing.T) { + counts := &countingObserver{} + up := &spreading{addrs: []string{refusedAddr(t), refusedAddr(t)}} + _, addr := startServer(t, ProtocolTCP, &Config{ + Observer: counts, Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{ConnectTimeout: timeconv.Duration(time.Second)}, + }) + expectClosed(t, dialTCP(t, addr)) + for _, r := range up.seen() { + if r.dialed.Load() != 1 || !r.failed() || abandoned(r) { + t.Errorf("%s: dialed %d times, failed %v", r.addr, r.dialed.Load(), r.failed()) + } + } + // an upstream with nothing to race refuses the flow + up.empty() + expectClosed(t, dialTCP(t, addr)) + waitFor(t, func() bool { + _, results, _ := counts.snapshot() + return results[ResultDialFailed] == 1 && results[ResultNoUpstream] == 1 + }) +} + +// udpSink records the datagrams it receives and answers each one +func udpSink(t *testing.T, prefix string) (string, func() []string) { + t.Helper() + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = pc.Close() }) + var mu sync.Mutex + var got []string + go func() { + buf := make([]byte, 2048) + for { + n, addr, err := pc.ReadFrom(buf) + if err != nil { + return + } + mu.Lock() + got = append(got, string(buf[:n])) + mu.Unlock() + _, _ = pc.WriteTo(append([]byte(prefix), buf[:n]...), addr) + } + }() + return pc.LocalAddr().String(), func() []string { + mu.Lock() + defer mu.Unlock() + return append([]string(nil), got...) + } +} + +func TestPacketServerMirrorsDatagrams(t *testing.T) { + primary, primaryGot := udpSink(t, "primary:") + second, secondGot := udpSink(t, "second:") + third, thirdGot := udpSink(t, "third:") + const unreachable = "192.0.2.1:9" + up := &spreading{addrs: []string{primary, second, unreachable, third}} + srv, addr, _ := startPacketServerWith(t, &Config{ + Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{IdleTimeout: timeconv.Duration(200 * time.Millisecond)}, + }, func(ctx context.Context, a string) (net.Conn, error) { + if a == unreachable { + return nil, errors.New("no route to host") + } + var d net.Dialer + return d.DialContext(ctx, "udp", a) + }) + client := udpClient(t, addr) + for _, msg := range []string{"one", "two", "three"} { + if got := datagram(t, client, msg); got != "primary:"+msg { + t.Fatalf("reply = %q", got) + } + } + // every mirror got every datagram, in order, and nothing a mirror answered came back + for name, got := range map[string]func() []string{"primary": primaryGot, "second": secondGot, "third": thirdGot} { + waitFor(t, func() bool { return len(got()) == 3 }) + if strings.Join(got(), ",") != "one,two,three" { + t.Errorf("%s received %v", name, got()) + } + } + _ = client.SetReadDeadline(time.Now().Add(100 * time.Millisecond)) + if n, err := client.Read(make([]byte, 64)); err == nil { + t.Errorf("a mirror's reply reached the client: %d bytes", n) + } + waitFor(t, func() bool { return srv.ActiveSessions() == 0 }) + routes := up.seen() + if len(routes) != 4 { + t.Fatalf("%d routes", len(routes)) + } + for i, r := range routes { + wantFailed := r.addr == unreachable + wantClosed := int32(1) + if wantFailed { + wantClosed = 0 + } + if r.dialed.Load() != 1 || r.failed() != wantFailed || r.closed.Load() != wantClosed { + t.Errorf("route %d (%s): dialed %d, failed %v, closed %d", i, r.addr, r.dialed.Load(), r.failed(), r.closed.Load()) + } + } + if routes[0].firstByte.Load() != 1 || routes[1].firstByte.Load() != 0 { + t.Error("only the answering route has a first reply to report") + } +} + +func TestPacketServerBoundsItsMirrors(t *testing.T) { + primary, _ := udpSink(t, "primary:") + up := &spreading{addrs: []string{primary}} + for range MaxUDPMirrors + 2 { + sink, _ := udpSink(t, "mirror:") + up.addrs = append(up.addrs, sink) + } + srv, addr, _ := startPacketServer(t, &Config{Table: tableOf(t, map[string]Upstream{"": up})}) + if got := datagram(t, udpClient(t, addr), "x"); got != "primary:x" { + t.Fatalf("reply = %q", got) + } + var open, left int + for _, r := range up.seen()[1:] { + if abandoned(r) { + left++ + } else if !r.failed() { + open++ + } + } + if open != MaxUDPMirrors || left != 2 { + t.Errorf("%d mirrors open and %d abandoned", open, left) + } + _ = srv.Close() + for i, r := range up.seen() { + if !r.failed() && r.closed.Load() != 1 { + t.Errorf("route %d was closed %d times", i, r.closed.Load()) + } + } +} + +// a flow closed while it was still dialing closes the mirrors it dialed, too +func TestCloseDoesNotPublishLateMirrors(t *testing.T) { + primary, _ := udpSink(t, "primary:") + mirror, _ := udpSink(t, "mirror:") + up := &spreading{addrs: []string{primary, mirror}} + dialing := make(chan struct{}, 2) + srv, addr, _ := startPacketServerWith(t, &Config{ + Table: tableOf(t, map[string]Upstream{"": up}), + Options: &options.Options{ConnectTimeout: timeconv.Duration(10 * time.Second)}, + }, func(ctx context.Context, a string) (net.Conn, error) { + if a == primary { + dialing <- struct{}{} + <-ctx.Done() + } + // the dial ignores the cancellation and hands back a usable socket anyway + return net.Dial("udp", a) + }) + if _, err := udpClient(t, addr).Write([]byte("x")); err != nil { + t.Fatal(err) + } + <-dialing + _ = srv.Close() + routes := up.seen() + if len(routes) != 2 { + t.Fatalf("%d routes", len(routes)) + } + for i, r := range routes { + if r.failed() || r.closed.Load() != 1 { + t.Errorf("route %d: failed %v, closed %d", i, r.failed(), r.closed.Load()) + } + } +} diff --git a/pkg/proxy/l4/table.go b/pkg/proxy/l4/table.go index 4ac6c6a4d..e17b53f20 100644 --- a/pkg/proxy/l4/table.go +++ b/pkg/proxy/l4/table.go @@ -21,21 +21,12 @@ package l4 import ( "errors" "fmt" - "net" "slices" "strings" - "sync" - "sync/atomic" - "github.com/trickstercache/trickster/v2/pkg/backends" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" "github.com/trickstercache/trickster/v2/pkg/proxy/hostnames" ) -// maxPoolDepth bounds how far a pool of pools is followed when selecting an address: a pool -// whose members are themselves pools, which is what a weighted rule over discovered endpoints is. -const maxPoolDepth = 2 - var ( // ErrDuplicateHost indicates a host already routed by the table. ErrDuplicateHost = errors.New("host is already routed") @@ -43,159 +34,6 @@ var ( ErrDuplicateCatchAll = errors.New("the listener already has a catch-all upstream") ) -// Upstream chooses the address a connection is relayed to. -type Upstream interface { - // Addr returns the host:port to dial for the next connection, or false to refuse it. - Addr() (string, bool) -} - -// reservedTLD is the top-level domain reserved never to resolve; an address under it refuses -// connections without a lookup, which is how a member that must refuse its share is expressed. -const reservedTLD = ".invalid" - -// Refusing reports whether an address can never be dialed: one whose host is under the reserved -// .invalid domain. -func Refusing(addr string) bool { - host, _, err := net.SplitHostPort(addr) - if err != nil { - host = addr - } - host = strings.TrimSuffix(strings.ToLower(host), ".") - return strings.HasSuffix(host, reservedTLD) -} - -// Static returns an upstream with one fixed address, or one that refuses every connection when -// the address can never be dialed. -func Static(addr string) Upstream { - if Refusing(addr) { - return refusingUpstream{} - } - return staticUpstream{addr: addr} -} - -type staticUpstream struct { - addr string -} - -func (s staticUpstream) Addr() (string, bool) { - return s.addr, true -} - -// refusingUpstream refuses every connection, without a lookup or a dial -type refusingUpstream struct{} - -func (refusingUpstream) Addr() (string, bool) { - return "", false -} - -// pooled is implemented by a backend whose members are a load balancer pool. -type pooled interface { - Pool() pool.Pool -} - -// FromBackend returns an upstream over a backend: a pool holder commits each connection to one -// healthy member by weighted rotation, a member that is itself a pool choosing again the same -// way, and any other backend is dialed at its origin host. A member that cannot be dialed refuses -// its share rather than passing it to a sibling, as the HTTP load balancer does. -func FromBackend(b backends.Backend) Upstream { - return fromBackend(b, 0) -} - -func fromBackend(b backends.Backend, depth int) Upstream { - if b == nil { - return nil - } - if p, ok := b.(pooled); ok { - return &poolUpstream{pool: p.Pool, depth: depth} - } - cfg := b.Configuration() - if cfg == nil || cfg.Host == "" { - return nil - } - return Static(cfg.Host) -} - -// poolUpstream rotates over a pool's healthy members with the members' weights, so consecutive -// connections are apportioned as the round robin mechanism apportions requests. -type poolUpstream struct { - pool func() pool.Pool - depth int - pos atomic.Uint64 - // nested keeps one rotation per member that is itself a pool, keyed by the member's - // backend, which outlives the pools swapped beneath it as membership changes - nested sync.Map -} - -func (p *poolUpstream) Addr() (string, bool) { - pl := p.pool() - if pl == nil { - return "", false - } - targets := pl.Targets() - if len(targets) == 0 { - return "", false - } - t := targets[p.start(targets)] - if t == nil { - return "", false - } - return p.resolve(t) -} - -func (p *poolUpstream) start(targets pool.Targets) int { - // each member owns a weight-sized span of the rotation, as the round robin mechanism does - var total int - weighted := false - for _, t := range targets { - if t == nil { - continue - } - if t.Weight() != 1 { - weighted = true - } - total += t.Weight() - } - if total == 0 { - return 0 - } - k := int(p.pos.Add(1) % uint64(total)) // #nosec G115 -- the value is below total, an int sum - if !weighted { - return k - } - for i, t := range targets { - if t == nil { - continue - } - k -= t.Weight() - if k < 0 { - return i - } - } - return len(targets) - 1 -} - -func (p *poolUpstream) resolve(t *pool.Target) (string, bool) { - b := t.Backend() - if b == nil { - return "", false - } - if _, ok := b.(pooled); ok { - if p.depth+1 >= maxPoolDepth { - return "", false - } - v, _ := p.nested.LoadOrStore(b, fromBackend(b, p.depth+1)) - if up, ok := v.(Upstream); ok && up != nil { - return up.Addr() - } - return "", false - } - cfg := b.Configuration() - if cfg == nil || cfg.Host == "" || Refusing(cfg.Host) { - return "", false - } - return cfg.Host, true -} - // Table routes a connection to an upstream by the server name it offered: an exact host first, // then the longest wildcard suffix, then the catch-all for a connection naming no routed host. type Table struct { diff --git a/pkg/proxy/l4/table_test.go b/pkg/proxy/l4/table_test.go index 16112d0b9..fc3828478 100644 --- a/pkg/proxy/l4/table_test.go +++ b/pkg/proxy/l4/table_test.go @@ -18,58 +18,12 @@ package l4 import ( "errors" - "net/http" + "net" + "net/netip" "testing" "time" - - "github.com/trickstercache/trickster/v2/pkg/backends" - "github.com/trickstercache/trickster/v2/pkg/backends/alb/pool" - "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" - bo "github.com/trickstercache/trickster/v2/pkg/backends/options" ) -func originBackend(t *testing.T, name, addr string) backends.Backend { - t.Helper() - o := bo.New() - o.OriginURL = "tcp://" + addr - if err := o.Initialize(name); err != nil { - t.Fatal(err) - } - b, err := backends.New(name, o, nil, http.NotFoundHandler(), nil) - if err != nil { - t.Fatal(err) - } - return b -} - -type pooledBackend struct { - backends.Backend - p pool.Pool -} - -func (b *pooledBackend) Pool() pool.Pool { return b.p } - -func newPool(t *testing.T, members ...*pool.Target) pool.Pool { - t.Helper() - p := pool.New(members, int(healthcheck.StatusUnchecked)) - t.Cleanup(p.Stop) - return p -} - -func member(b backends.Backend, weight int, status int32) *pool.Target { - st := healthcheck.NewStatus(b.Name(), "", "", status, time.Time{}, nil) - return pool.NewWeightedTarget(http.NotFoundHandler(), st, b, weight) -} - -func pooledOf(t *testing.T, name string, members ...*pool.Target) backends.Backend { - t.Helper() - b, err := backends.New(name, bo.New(), nil, http.NotFoundHandler(), nil) - if err != nil { - t.Fatal(err) - } - return &pooledBackend{Backend: b, p: newPool(t, members...)} -} - func TestTableLookup(t *testing.T) { tbl := NewTable() if !tbl.Empty() { @@ -128,126 +82,69 @@ func TestTableLookup(t *testing.T) { } } -func TestStaticAndFromBackend(t *testing.T) { - if got, ok := Static("h:1").Addr(); !ok || got != "h:1" { - t.Errorf("Static = %v, %v", got, ok) +func TestStatic(t *testing.T) { + route, ok := Static("h:1").Pick(Flow{}) + if !ok || route.Addr() != "h:1" { + t.Fatalf("Static = %v, %v", route, ok) + } + // a fixed address has no other route to fall back to, and nothing to report to + if !route.Final() { + t.Error("a static route is not final") + } + route.Dialed(time.Millisecond, nil) + route.FirstByte() + route.Closed(nil) + if again, _ := Static("h:1").Pick(Flow{}); again.Addr() != "h:1" { + t.Errorf("second pick = %v", again) } // an address under the reserved .invalid domain can never resolve, so it is refused without - // a lookup, as a backend is whose origin names one + // a lookup for _, addr := range []string{"unresolved.kgw.invalid:1", "x.INVALID.:9", "x.invalid"} { if !Refusing(addr) { t.Errorf("Refusing(%q) = false", addr) } - if _, ok := Static(addr).Addr(); ok { + if _, ok := Static(addr).Pick(Flow{}); ok { t.Errorf("Static(%q) dials", addr) } } if Refusing("invalid.example.com:1") || Refusing("10.0.0.1:1") { t.Error("a resolvable address is refused") } - if _, ok := FromBackend(originBackend(t, "gone", "unresolved.kgw.invalid:1")).Addr(); ok { - t.Error("a backend under .invalid dials") + if allocs := testing.AllocsPerRun(100, func() { _, _ = Static("h:1").Pick(Flow{}) }); allocs > 1 { + t.Errorf("a static pick allocates %v", allocs) } - if FromBackend(nil) != nil { - t.Error("nil backend yields an upstream") - } - hostless, err := backends.New("hostless", bo.New(), nil, http.NotFoundHandler(), nil) - if err != nil { - t.Fatal(err) - } - if FromBackend(hostless) != nil { - t.Error("a backend without an origin host yields an upstream") - } - if got, ok := FromBackend(originBackend(t, "o", "10.0.0.1:9000")).Addr(); !ok || got != "10.0.0.1:9000" { - t.Errorf("origin backend = %v, %v", got, ok) + fixed := Static("h:1") + if allocs := testing.AllocsPerRun(100, func() { _, _ = fixed.Pick(Flow{}) }); allocs != 0 { + t.Errorf("a pick from a built static upstream allocates %v", allocs) } } -func TestPoolUpstreamRotatesWithWeights(t *testing.T) { - a := originBackend(t, "a", "10.0.0.1:1") - b := originBackend(t, "b", "10.0.0.2:1") - down := originBackend(t, "down", "10.0.0.3:1") - up := FromBackend(pooledOf(t, "alb", - member(a, 2, healthcheck.StatusPassing), - member(b, 1, healthcheck.StatusPassing), - member(down, 1, healthcheck.StatusFailing))) - counts := make(map[string]int) - for range 6 { - addr, ok := up.Addr() - if !ok { - t.Fatal("a pool with healthy members refused") - } - counts[addr]++ - } - if counts["10.0.0.1:1"] != 4 || counts["10.0.0.2:1"] != 2 || counts["10.0.0.3:1"] != 0 { - t.Errorf("weighted rotation = %v", counts) - } - // an even pool is a plain rotation - even := FromBackend(pooledOf(t, "even", member(a, 1, healthcheck.StatusPassing), - member(b, 1, healthcheck.StatusPassing))) - first, _ := even.Addr() - second, _ := even.Addr() - third, _ := even.Addr() - if first == second || first != third { - t.Errorf("rotation: %v, %v, %v", first, second, third) - } - empty := FromBackend(pooledOf(t, "empty")) - if got, ok := empty.Addr(); ok { - t.Errorf("empty pool = %v", got) - } - holderless := &pooledBackend{Backend: hostlessBackend(t)} - if got, ok := FromBackend(holderless).Addr(); ok { - t.Errorf("nil pool = %v", got) - } - // a member with nothing to dial refuses its own share rather than passing it on - mixed := FromBackend(pooledOf(t, "mixed", member(a, 1, healthcheck.StatusPassing), - member(hostlessBackend(t), 1, healthcheck.StatusPassing))) - var refused int - for range 4 { - if _, ok := mixed.Addr(); !ok { - refused++ - } - } - if refused != 2 { - t.Errorf("refused %d of 4, want the hostless member's share", refused) +func TestFlowOf(t *testing.T) { + tcp := flowOf("l", ProtocolTLS, &net.TCPAddr{IP: net.ParseIP("192.0.2.7"), Port: 4431}, "shop.example.com") + if tcp.Protocol != ProtocolTLS || tcp.ServerName != "shop.example.com" || + tcp.Client != netip.MustParseAddrPort("192.0.2.7:4431") { + t.Errorf("tcp flow = %+v", tcp) } -} - -func hostlessBackend(t *testing.T) backends.Backend { - t.Helper() - b, err := backends.New("hostless", bo.New(), nil, http.NotFoundHandler(), nil) - if err != nil { - t.Fatal(err) + udp := flowOf("l", ProtocolUDP, &net.UDPAddr{IP: net.ParseIP("2001:db8::9"), Port: 53}, "") + if udp.Client != netip.MustParseAddrPort("[2001:db8::9]:53") { + t.Errorf("udp flow = %+v", udp) } - return b -} - -func TestPoolUpstreamFollowsNestedPools(t *testing.T) { - // an outer pool of two inner pools, as a weighted rule in endpoint mode compiles to - inner1 := pooledOf(t, "inner1", member(originBackend(t, "a", "10.1.0.1:1"), 1, healthcheck.StatusPassing), - member(originBackend(t, "b", "10.1.0.2:1"), 1, healthcheck.StatusPassing)) - inner2 := pooledOf(t, "inner2", member(originBackend(t, "c", "10.2.0.1:1"), 1, healthcheck.StatusPassing)) - outer := FromBackend(pooledOf(t, "outer", member(inner1, 1, healthcheck.StatusPassing), - member(inner2, 1, healthcheck.StatusPassing), member(hostlessBackend(t), 1, healthcheck.StatusPassing))) - seen := make(map[string]int) - var refused int - for range 6 { - addr, ok := outer.Addr() - if !ok { - refused++ - continue - } - seen[addr]++ + // an IPv4 peer of a dual-stack socket arrives mapped into IPv6; it is one client either way + mapped := flowOf("l", ProtocolTCP, &net.TCPAddr{IP: net.ParseIP("::ffff:192.0.2.7"), Port: 80}, "") + if mapped.Client.Addr() != netip.MustParseAddr("192.0.2.7") { + t.Errorf("mapped client = %v", mapped.Client) } - if refused != 2 || seen["10.2.0.1:1"] != 2 || seen["10.1.0.1:1"] != 1 || seen["10.1.0.2:1"] != 1 { - t.Errorf("outer rotation = %v with %d refused; want each inner pool its share, rotating within", - seen, refused) + // any other address is read from its text; one that is not an address leaves no client + if other := flowOf("l", ProtocolTCP, textAddr("198.51.100.4:9"), ""); other.Client != netip.MustParseAddrPort("198.51.100.4:9") { + t.Errorf("text address = %+v", other) } - // a pool nested beyond the depth bound is not followed - deep := FromBackend(pooledOf(t, "l0", member(pooledOf(t, "l1", member(pooledOf(t, "l2", - member(originBackend(t, "z", "10.9.0.1:1"), 1, healthcheck.StatusPassing)), - 1, healthcheck.StatusPassing)), 1, healthcheck.StatusPassing))) - if got, ok := deep.Addr(); ok { - t.Errorf("a pool three deep was followed: %v", got) + if flowOf("l", ProtocolTCP, textAddr("pipe"), "").Client.IsValid() || flowOf("l", ProtocolTCP, nil, "").Client.IsValid() { + t.Error("an unusable peer address produced a client") } } + +type textAddr string + +func (textAddr) Network() string { return "test" } + +func (a textAddr) String() string { return string(a) } diff --git a/pkg/proxy/l4/udp.go b/pkg/proxy/l4/udp.go index 9ce3c2f2f..293aea28d 100644 --- a/pkg/proxy/l4/udp.go +++ b/pkg/proxy/l4/udp.go @@ -22,9 +22,8 @@ import ( "net" "sync" "sync/atomic" + "syscall" "time" - - "github.com/trickstercache/trickster/v2/pkg/observability/metrics" ) // maxDatagram is the largest UDP payload a socket can carry. @@ -127,12 +126,27 @@ type udpSession struct { closed bool // upstream is set once the flow is open, when its writer starts upstream net.Conn + // set once the upstream is dialed; told when the session ends + route Route + // the further upstreams that receive a copy of each datagram; fixed once the flow is open + mirrors []udpMirror + // set when the upstream refused a datagram, which is how an unreachable udp upstream shows + fault error // queue is a ring of the datagrams waiting for the writer, in arrival order queue [maxQueuedDatagrams][]byte head, num int last atomic.Int64 } +// udpMirror is one further upstream a flow's datagrams are copied to +type udpMirror struct { + conn net.Conn + route Route +} + +// MaxUDPMirrors bounds the upstreams one flow is copied to, each of which holds a socket. +const MaxUDPMirrors = 7 + func newUDPSession(client net.Addr) *udpSession { sess := &udpSession{client: client} sess.wake.L = &sess.mu @@ -229,11 +243,11 @@ func (s *PacketServer) isClosed() bool { } func (s *PacketServer) result(r string) { - metrics.ProxyStreamConnections.WithLabelValues(s.name, ProtocolUDP, r).Inc() + s.cfg.Load().observer().Result(r) } func (s *PacketServer) dropped(reason string) { - metrics.ProxyStreamDroppedDatagrams.WithLabelValues(s.name, reason).Inc() + s.cfg.Load().observer().Dropped(reason) } // forward hands a datagram to its client's flow, opening the flow on a worker of its own, so the @@ -271,7 +285,7 @@ func (s *PacketServer) admit(key string, client net.Addr) *udpSession { s.sessions[key] = sess s.opening++ s.workers.Add(1) - metrics.ProxyStreamActiveConnections.WithLabelValues(s.name, ProtocolUDP).Inc() + s.cfg.Load().observer().Opened() return sess } @@ -321,6 +335,7 @@ func (s *PacketServer) enqueue(sess *udpSession, payload []byte) { func (s *PacketServer) writer(sess *udpSession, up net.Conn) { defer s.workers.Done() sess.mu.Lock() + mirrors := sess.mirrors for { for sess.num == 0 && !sess.closed { sess.wake.Wait() @@ -338,10 +353,15 @@ func (s *PacketServer) writer(sess *udpSession, up net.Conn) { // for as long as the write blocks, however soon its ring slot is reused _ = up.SetWriteDeadline(time.Now().Add(udpWriteTimeout)) n, err := up.Write(payload) + for _, m := range mirrors { + // a copy that cannot be written is lost; the flow answers to its own upstream alone + _ = m.conn.SetWriteDeadline(time.Now().Add(udpWriteTimeout)) + _, _ = m.conn.Write(payload) + } s.queued.Add(-int64(len(payload))) switch { case err == nil: - metrics.ProxyStreamBytes.WithLabelValues(s.name, ProtocolUDP, DirectionIn).Add(float64(n)) + s.cfg.Load().observer().Bytes(DirectionIn, int64(n)) case isTimeout(err): s.dropped(DropWriteTimeout) default: @@ -378,16 +398,24 @@ func (s *PacketServer) open(sess *udpSession) (net.Conn, bool) { s.result(ResultNoRoute) return nil, false } - addr, ok := up.Addr() + route, ok := up.Pick(flowOf(s.name, ProtocolUDP, sess.client, "")) if !ok { s.result(ResultNoUpstream) return nil, false } if !s.acquireDial() { + route.Dialed(0, ErrAbandoned) return nil, false } + flow := flowOf(s.name, ProtocolUDP, sess.client, "") ctx, cancel := context.WithTimeout(s.ctx, cfg.options().Connect()) - conn, err := s.resolveAndDial(ctx, addr) + began := time.Now() + conn, err := s.resolveAndDial(ctx, route.Addr()) + route.Dialed(time.Since(began), err) + var mirrors []udpMirror + if err == nil { + mirrors = s.openMirrors(ctx, up, flow, route) + } cancel() s.releaseDial() if err != nil { @@ -399,8 +427,12 @@ func (s *PacketServer) open(sess *udpSession) (net.Conn, bool) { // the close pass has been through this flow; what was dialed after it is closed here sess.mu.Unlock() _ = conn.Close() + route.Closed(nil) + closeMirrors(mirrors) return nil, false } + sess.route = route + sess.mirrors = mirrors // what the flow kept while opening leaves the opening budget and stays queued for the writer var kept int64 for i := range sess.num { @@ -416,6 +448,37 @@ func (s *PacketServer) open(sess *udpSession) (net.Conn, bool) { return conn, true } +// openMirrors dials the further upstreams a mirroring upstream copies the flow to, within the +// dial the flow already holds. One that cannot be dialed is told so and left out. +func (s *PacketServer) openMirrors(ctx context.Context, up Upstream, flow Flow, primary Route) []udpMirror { + mirrorer, ok := up.(Mirrorer) + if !ok { + return nil + } + routes := mirrorer.Mirror(flow, primary) + for _, extra := range routes[min(len(routes), MaxUDPMirrors):] { + extra.Dialed(0, ErrAbandoned) + } + routes = routes[:min(len(routes), MaxUDPMirrors)] + mirrors := make([]udpMirror, 0, len(routes)) + for _, route := range routes { + began := time.Now() + conn, err := s.resolveAndDial(ctx, route.Addr()) + route.Dialed(time.Since(began), err) + if err == nil { + mirrors = append(mirrors, udpMirror{conn: conn, route: route}) + } + } + return mirrors +} + +func closeMirrors(mirrors []udpMirror) { + for _, m := range mirrors { + _ = m.conn.Close() + m.route.Closed(nil) + } +} + // settle moves a flow out of the opening count once it is open func (s *PacketServer) settle(sess *udpSession) { s.mu.Lock() @@ -516,17 +579,25 @@ func (s *PacketServer) reply(sess *udpSession, up net.Conn) { bp := datagramPool.Get().(*[]byte) defer datagramPool.Put(bp) buf := *bp + sess.mu.Lock() + route := sess.route + sess.mu.Unlock() + replied := false for { idle := s.cfg.Load().options().UDPIdle() _ = up.SetReadDeadline(time.Now().Add(idle)) n, err := up.Read(buf) if n > 0 { + if !replied && route != nil { + replied = true + route.FirstByte() + } sess.touch() s.mu.Lock() pc := s.conn s.mu.Unlock() if _, werr := pc.WriteTo(buf[:n], sess.client); werr == nil { - metrics.ProxyStreamBytes.WithLabelValues(s.name, ProtocolUDP, DirectionOut).Add(float64(n)) + s.cfg.Load().observer().Bytes(DirectionOut, int64(n)) } } if err == nil { @@ -536,6 +607,12 @@ func (s *PacketServer) reply(sess *udpSession, up net.Conn) { // the upstream was silent but the client was not; the session lives on continue } + if errors.Is(err, syscall.ECONNREFUSED) { + // nothing is listening at the upstream: the only sign a udp member is down + sess.mu.Lock() + sess.fault = err + sess.mu.Unlock() + } return } } @@ -546,6 +623,12 @@ func (s *PacketServer) end(key string, sess *udpSession) { if sess.upstream != nil { _ = sess.upstream.Close() } + if sess.route != nil { + sess.route.Closed(sess.fault) + sess.route = nil + } + closeMirrors(sess.mirrors) + sess.mirrors = nil if sess.state == flowOpening { for sess.num > 0 { n := int64(len(sess.pop())) @@ -560,7 +643,7 @@ func (s *PacketServer) end(key string, sess *udpSession) { delete(s.sessions, key) } s.mu.Unlock() - metrics.ProxyStreamActiveConnections.WithLabelValues(s.name, ProtocolUDP).Dec() + s.cfg.Load().observer().Ended() s.workers.Done() } diff --git a/pkg/proxy/l4/udp_bounds_test.go b/pkg/proxy/l4/udp_bounds_test.go index f73cc161a..274de2b46 100644 --- a/pkg/proxy/l4/udp_bounds_test.go +++ b/pkg/proxy/l4/udp_bounds_test.go @@ -26,11 +26,8 @@ import ( "testing" "time" - "github.com/trickstercache/trickster/v2/pkg/observability/metrics" "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" - - "github.com/prometheus/client_golang/prometheus/testutil" ) // countingUpstream refuses its first n selections and then yields addr @@ -40,11 +37,11 @@ type countingUpstream struct { calls atomic.Int32 } -func (c *countingUpstream) Addr() (string, bool) { +func (c *countingUpstream) Pick(f Flow) (Route, bool) { if c.calls.Add(1) <= c.refuse { - return "", false + return nil, false } - return c.addr, true + return Static(c.addr).Pick(f) } func TestFailedFlowsDoNotFillTheSessionBound(t *testing.T) { @@ -448,8 +445,10 @@ func TestBlockedUpstreamWriteStallsOnlyItsFlow(t *testing.T) { release := make(chan struct{}) var first atomic.Pointer[stallingConn] var d net.Dialer + counts := &countingObserver{} srv, addr, _ := startPacketServerWith(t, &Config{ - Table: tableOf(t, map[string]Upstream{"": Static(echo)}), + Observer: counts, + Table: tableOf(t, map[string]Upstream{"": Static(echo)}), }, func(ctx context.Context, a string) (net.Conn, error) { conn, err := d.DialContext(ctx, "udp", a) if err != nil { @@ -462,8 +461,8 @@ func TestBlockedUpstreamWriteStallsOnlyItsFlow(t *testing.T) { } return conn, nil }) - dropsBefore := testutil.ToFloat64(metrics.ProxyStreamDroppedDatagrams.WithLabelValues("udp-test", DropQueueFull)) - timeoutsBefore := testutil.ToFloat64(metrics.ProxyStreamDroppedDatagrams.WithLabelValues("udp-test", DropWriteTimeout)) + dropsBefore := counts.dropped(DropQueueFull) + timeoutsBefore := counts.dropped(DropWriteTimeout) stuck := udpClient(t, addr) if _, err := stuck.Write([]byte("first")); err != nil { t.Fatal(err) @@ -484,7 +483,7 @@ func TestBlockedUpstreamWriteStallsOnlyItsFlow(t *testing.T) { } } waitFor(t, func() bool { - return testutil.ToFloat64(metrics.ProxyStreamDroppedDatagrams.WithLabelValues("udp-test", DropQueueFull))-dropsBefore >= 8 + return counts.dropped(DropQueueFull)-dropsBefore >= 8 }) srv.mu.Lock() sess := srv.sessions[stuck.LocalAddr().String()] @@ -500,7 +499,7 @@ func TestBlockedUpstreamWriteStallsOnlyItsFlow(t *testing.T) { } // the first write times out at the write bound and its datagram is dropped, not the flow waitFor(t, func() bool { - return testutil.ToFloat64(metrics.ProxyStreamDroppedDatagrams.WithLabelValues("udp-test", DropWriteTimeout)) > timeoutsBefore + return counts.dropped(DropWriteTimeout) > timeoutsBefore }) close(release) // once released the queue drains in order and the flow relays again @@ -529,8 +528,10 @@ func TestInFlightWritesStayWithinTheQueuedBudget(t *testing.T) { release := make(chan struct{}) var inFlight atomic.Int64 var d net.Dialer + counts := &countingObserver{} srv, addr, _ := startPacketServerWith(t, &Config{ - Table: tableOf(t, map[string]Upstream{"": Static(echo)}), + Observer: counts, + Table: tableOf(t, map[string]Upstream{"": Static(echo)}), }, func(ctx context.Context, a string) (net.Conn, error) { conn, err := d.DialContext(ctx, "udp", a) if err != nil { @@ -538,7 +539,7 @@ func TestInFlightWritesStayWithinTheQueuedBudget(t *testing.T) { } return &stallingConn{Conn: conn, release: release, inFlight: &inFlight}, nil }) - dropsBefore := testutil.ToFloat64(metrics.ProxyStreamDroppedDatagrams.WithLabelValues("udp-test", DropQueueFull)) + dropsBefore := counts.dropped(DropQueueFull) const flows = 24 payload := make([]byte, 8192) clients := make([]*net.UDPConn, flows) @@ -569,7 +570,7 @@ func TestInFlightWritesStayWithinTheQueuedBudget(t *testing.T) { if held := inFlight.Load(); held <= 0 { t.Error("no write was stalled") } - drops := testutil.ToFloat64(metrics.ProxyStreamDroppedDatagrams.WithLabelValues("udp-test", DropQueueFull)) - dropsBefore + drops := counts.dropped(DropQueueFull) - dropsBefore if drops <= 0 { t.Error("nothing beyond the budget was dropped") } diff --git a/pkg/proxy/l4/udp_test.go b/pkg/proxy/l4/udp_test.go index cd8ac02ae..c361afd2b 100644 --- a/pkg/proxy/l4/udp_test.go +++ b/pkg/proxy/l4/udp_test.go @@ -24,7 +24,6 @@ import ( "testing" "time" - "github.com/trickstercache/trickster/v2/pkg/backends/healthcheck" "github.com/trickstercache/trickster/v2/pkg/parsing/timeconv" "github.com/trickstercache/trickster/v2/pkg/proxy/l4/options" ) @@ -148,10 +147,7 @@ func TestPacketServerRelaysSessions(t *testing.T) { func TestPacketServerRotatesAcrossPool(t *testing.T) { a, b := udpEcho(t, "a:"), udpEcho(t, "b:") - pl := pooledOf(t, "alb", - member(originBackend(t, "a", a), 1, healthcheck.StatusPassing), - member(originBackend(t, "b", b), 1, healthcheck.StatusPassing)) - _, addr, _ := startPacketServer(t, &Config{Table: tableOf(t, map[string]Upstream{"": FromBackend(pl)})}) + _, addr, _ := startPacketServer(t, &Config{Table: tableOf(t, map[string]Upstream{"": rotate(a, b)})}) seen := make(map[string]int) for range 4 { seen[datagram(t, udpClient(t, addr), "x")]++ @@ -164,7 +160,7 @@ func TestPacketServerRotatesAcrossPool(t *testing.T) { func TestPacketServerDropsWhatItCannotRoute(t *testing.T) { for name, tbl := range map[string]*Table{ "empty": NewTable(), - "no_member": tableOf(t, map[string]Upstream{"": FromBackend(pooledOf(t, "none"))}), + "no_member": tableOf(t, map[string]Upstream{"": rotate()}), "bad_addr": tableOf(t, map[string]Upstream{"": Static("not an address")}), } { t.Run(name, func(t *testing.T) { @@ -353,11 +349,9 @@ func TestPacketServerMixedPoolRefusesTheInvalidShare(t *testing.T) { // expires rather than looking its upstream up again per datagram echo := udpEcho(t, "echo:") shortFailedLifetime(t, time.Second) - pl := pooledOf(t, "alb", - member(originBackend(t, "live", echo), 1, healthcheck.StatusPassing), - member(originBackend(t, "gone", "unresolved.kgw.invalid:1"), 1, healthcheck.StatusPassing)) + up := rotate(echo, "unresolved.kgw.invalid:1") srv, addr, _ := startPacketServer(t, &Config{ - Table: tableOf(t, map[string]Upstream{"": FromBackend(pl)}), + Table: tableOf(t, map[string]Upstream{"": up}), Options: &options.Options{IdleTimeout: timeconv.Duration(time.Second)}, }) var served, refused int diff --git a/pkg/proxy/l4/upstream.go b/pkg/proxy/l4/upstream.go new file mode 100644 index 000000000..940fc8883 --- /dev/null +++ b/pkg/proxy/l4/upstream.go @@ -0,0 +1,153 @@ +/* + * Copyright 2026 The Trickster Authors + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package l4 + +import ( + "errors" + "net" + "net/netip" + "time" + + "github.com/trickstercache/trickster/v2/pkg/proxy/hostnames" +) + +// Flow is what the relay knows about a connection or session before it chooses an upstream. +type Flow struct { + // Listener is the name of the listener that accepted the flow + Listener string + // Protocol is the listener's: tcp, tls or udp + Protocol string + // Client is the peer, or the source a trusted PROXY protocol header named + Client netip.AddrPort + // ServerName is the TLS server name the client offered; tls only + ServerName string + // Proxy reads the PROXY protocol header the connection arrived behind; nil without one + Proxy ProxyHeader +} + +// ProxyHeader is implemented by a client connection that was accepted behind a PROXY protocol +// header, which a listener's accepted connections may be. +type ProxyHeader interface { + // ProxyTLV returns the value of the first version 2 TLV of the given type, if it has one. + ProxyTLV(typ byte) ([]byte, bool) +} + +// Upstream chooses where a flow is relayed to. +type Upstream interface { + // Pick commits the flow to one route, or returns false to refuse it. + Pick(Flow) (Route, bool) +} + +// Retrier is optionally implemented by an Upstream that may offer another route when a dial +// fails. The relay asks only for a route that is not Final, never for a udp flow, and keeps +// every attempt inside the one connect timeout. +type Retrier interface { + // Retry returns a route to try in place of failed, or false when there is none to offer. + Retry(f Flow, failed Route) (Route, bool) +} + +// Racer is optionally implemented by an Upstream that connects a tcp or tls flow to several +// routes at once. The relay keeps the first to connect and reports ErrAbandoned to the rest. +type Racer interface { + // Race returns the routes to connect to together; none refuses the flow. + Race(Flow) []Route +} + +// Mirrorer is optionally implemented by an Upstream whose udp flows are copied to further +// routes: each receives every datagram the client sends, and its replies are discarded. +type Mirrorer interface { + // Mirror returns the routes to copy the flow to, beside the primary route it was given. + Mirror(f Flow, primary Route) []Route +} + +// Route is one committed choice of upstream. The relay reports what became of it: Dialed +// once, then, only if the dial succeeded, Closed once when the relay ends. +type Route interface { + // Addr is the host:port to dial. + Addr() string + // Final reports that a failed dial must not be retried on another route. + Final() bool + // Dialed reports how long the dial took and whether it failed. A failed dial ends the route. + Dialed(time.Duration, error) + // FirstByte reports the first byte or datagram received from the upstream. + FirstByte() + // Closed reports that the relayed connection or session has ended. err is nil unless the + // upstream turned out to be unreachable after all, as a udp upstream that answers a + // datagram with a port-unreachable does. + Closed(err error) +} + +// ErrAbandoned is reported to Route.Dialed when the relay gave the route up before dialing it, +// which says nothing about the upstream. +var ErrAbandoned = errors.New("route abandoned before it was dialed") + +// Refusing reports whether an address can never be dialed: one whose host is under the +// reserved .invalid domain, which is how an upstream that must refuse its share is expressed. +func Refusing(addr string) bool { + return hostnames.Reserved(addr) +} + +// Static returns an upstream with one fixed address, or one that refuses every flow when the +// address can never be dialed. +func Static(addr string) Upstream { + if Refusing(addr) { + return refusingUpstream{} + } + return &staticUpstream{addr: addr} +} + +// staticUpstream is its own route: there is nothing to report to and nowhere else to go +type staticUpstream struct { + addr string +} + +func (s *staticUpstream) Pick(Flow) (Route, bool) { return s, true } + +func (s *staticUpstream) Addr() string { return s.addr } + +func (s *staticUpstream) Final() bool { return true } + +func (s *staticUpstream) Dialed(time.Duration, error) {} + +func (s *staticUpstream) FirstByte() {} + +func (s *staticUpstream) Closed(error) {} + +// refusingUpstream refuses every flow, without a lookup or a dial +type refusingUpstream struct{} + +func (refusingUpstream) Pick(Flow) (Route, bool) { return nil, false } + +// flowOf describes a flow from its peer address, which a PROXY protocol listener has already +// replaced with the real client's +func flowOf(listener, protocol string, peer net.Addr, serverName string) Flow { + f := Flow{Listener: listener, Protocol: protocol, ServerName: serverName} + switch a := peer.(type) { + case *net.TCPAddr: + f.Client = a.AddrPort() + case *net.UDPAddr: + f.Client = a.AddrPort() + case nil: + default: + if ap, err := netip.ParseAddrPort(a.String()); err == nil { + f.Client = ap + } + } + if f.Client.IsValid() { + f.Client = netip.AddrPortFrom(f.Client.Addr().Unmap(), f.Client.Port()) + } + return f +} diff --git a/pkg/proxy/listener/native/native.go b/pkg/proxy/listener/native/native.go index b06e471a6..dd75dfc61 100644 --- a/pkg/proxy/listener/native/native.go +++ b/pkg/proxy/listener/native/native.go @@ -60,6 +60,9 @@ type Adapter interface { ValidateListener(*listenerconfig.Options) error ValidateBackend(*bo.Options) error ValidateUserRouter(*config.Config, string, *bo.Options) error + // ValidateBalancer validates an ALB whose selection strategy commits each of the + // listener's sessions to one member of its pool. + ValidateBalancer(*config.Config, string, *bo.Options) error Describe(*config.Config, string) (Descriptor, error) Build(BuildRequest) (listener.ProtocolServer, error) RouteResolver(BuildRequest) backends.RouteResolver diff --git a/pkg/proxy/listener/proxyprotocol.go b/pkg/proxy/listener/proxyprotocol.go index 172bafa8d..2bbe70c36 100644 --- a/pkg/proxy/listener/proxyprotocol.go +++ b/pkg/proxy/listener/proxyprotocol.go @@ -64,3 +64,27 @@ func (o *ProxyProtocolOptions) policy(c proxyproto.ConnPolicyOptions) (proxyprot } return proxyproto.SKIP, nil } + +// ProxyTLV returns the value of the first PROXY protocol v2 TLV of the given type that the +// connection's header carried; a connection without such a header has none. +func (o *observedConnection) ProxyTLV(typ byte) ([]byte, bool) { + pc, ok := o.Conn.(*proxyproto.Conn) + if !ok { + return nil, false + } + // the header is read here if nothing has read it yet, within the listener's header timeout + h := pc.ProxyHeader() + if h == nil { + return nil, false + } + tlvs, err := h.TLVs() + if err != nil { + return nil, false + } + for _, tlv := range tlvs { + if byte(tlv.Type) == typ { + return tlv.Value, true + } + } + return nil, false +} diff --git a/pkg/proxy/listener/proxyprotocol_test.go b/pkg/proxy/listener/proxyprotocol_test.go index afb0e65bd..4bfd9e4dc 100644 --- a/pkg/proxy/listener/proxyprotocol_test.go +++ b/pkg/proxy/listener/proxyprotocol_test.go @@ -125,3 +125,76 @@ func TestListenerProxyProtocol(t *testing.T) { require.Equal(t, "400 Bad Request", status) }) } + +func TestObservedConnectionProxyTLV(t *testing.T) { + tcp, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + ln := (&ProxyProtocolOptions{Enabled: true}).wrap(tcp) + defer ln.Close() + accept := func(t *testing.T, send func(net.Conn)) *observedConnection { + t.Helper() + client, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + t.Cleanup(func() { _ = client.Close() }) + send(client) + c, err := ln.Accept() + require.NoError(t, err) + t.Cleanup(func() { _ = c.Close() }) + return &observedConnection{Conn: c} + } + const vpce = 0xEA + withTLVs := accept(t, func(c net.Conn) { + h := proxyproto.HeaderProxyFromAddrs(2, + &net.TCPAddr{IP: net.ParseIP("203.0.113.9"), Port: 4242}, + &net.TCPAddr{IP: net.ParseIP("10.0.0.1"), Port: 80}) + require.NoError(t, h.SetTLVs([]proxyproto.TLV{ + {Type: proxyproto.PP2_TYPE_AUTHORITY, Value: []byte("shop.example.com")}, + {Type: vpce, Value: []byte("first")}, + {Type: vpce, Value: []byte("second")}, + })) + _, err := h.WriteTo(c) + require.NoError(t, err) + }) + v, ok := withTLVs.ProxyTLV(vpce) + require.True(t, ok) + require.Equal(t, "first", string(v)) + v, ok = withTLVs.ProxyTLV(byte(proxyproto.PP2_TYPE_AUTHORITY)) + require.True(t, ok) + require.Equal(t, "shop.example.com", string(v)) + _, ok = withTLVs.ProxyTLV(0xEB) + require.False(t, ok) + + // a version 1 header has no TLVs, and a connection with no header has nothing at all + v1 := accept(t, func(c net.Conn) { + _, err := io.WriteString(c, "PROXY TCP4 203.0.113.9 10.0.0.1 4242 80\r\n") + require.NoError(t, err) + }) + _, ok = v1.ProxyTLV(vpce) + require.False(t, ok) + bare := accept(t, func(c net.Conn) { + _, err := io.WriteString(c, "hello") + require.NoError(t, err) + }) + _, ok = bare.ProxyTLV(vpce) + require.False(t, ok) + client, server := net.Pipe() + defer client.Close() + defer server.Close() + _, ok = (&observedConnection{Conn: server}).ProxyTLV(vpce) + require.False(t, ok) + + // TLVs that do not parse are no TLVs + broken := accept(t, func(c net.Conn) { + h := proxyproto.HeaderProxyFromAddrs(2, + &net.TCPAddr{IP: net.ParseIP("203.0.113.9"), Port: 4242}, + &net.TCPAddr{IP: net.ParseIP("10.0.0.1"), Port: 80}) + raw, err := h.Format() + require.NoError(t, err) + // one trailing byte cannot be a TLV, which needs three; the length covers it + raw[15]++ + _, err = c.Write(append(raw, 0xEA)) + require.NoError(t, err) + }) + _, ok = broken.ProxyTLV(vpce) + require.False(t, ok) +} diff --git a/pkg/proxy/pgwire/adapter.go b/pkg/proxy/pgwire/adapter.go index f60428a97..d34b07343 100644 --- a/pkg/proxy/pgwire/adapter.go +++ b/pkg/proxy/pgwire/adapter.go @@ -128,6 +128,10 @@ func (a nativeListenerAdapter) ValidateUserRouter(c *config.Config, name string, return nil } +func (nativeListenerAdapter) ValidateBalancer(*config.Config, string, *bo.Options) error { + return errors.New("postgres session balancing is not supported") +} + func (a nativeListenerAdapter) Describe(c *config.Config, listenerName string) (native.Descriptor, error) { protocolConfig, _, err := a.listenerConfig(c, listenerName) if err != nil { diff --git a/pkg/proxy/pgwire/config_test.go b/pkg/proxy/pgwire/config_test.go index bfffd0a02..d25f1b689 100644 --- a/pkg/proxy/pgwire/config_test.go +++ b/pkg/proxy/pgwire/config_test.go @@ -320,6 +320,9 @@ func TestNativeListenerAdapterContract(t *testing.T) { if err := adapter.ValidateListener(nil); err == nil { t.Fatal("ValidateListener(nil) succeeded") } + if err := adapter.ValidateBalancer(nil, configTestName, nil); err == nil { + t.Fatal("ValidateBalancer() accepted a session balancer") + } defaults := listenerconfig.New(configTestName) if err := adapter.ValidateListener(defaults); err != nil || defaults.Postgres == nil { t.Fatalf("ValidateListener(defaults) = %v, options = %+v", err, defaults.Postgres) diff --git a/pkg/testutil/albpool/albpool.go b/pkg/testutil/albpool/albpool.go index 55a511bf5..563c65897 100644 --- a/pkg/testutil/albpool/albpool.go +++ b/pkg/testutil/albpool/albpool.go @@ -65,14 +65,8 @@ func New(healthyFloor int, hs []http.Handler) (pool.Pool, } // NewHealthy builds a pool with healthyFloor=-1 and pre-sets every target's -// status to StatusPassing. It replaces the -// -// p, _, st := albpool.New(-1, hs) -// for _, s := range st { s.Set(0) } -// time.Sleep(250 * time.Millisecond) -// -// boilerplate. Callers should still WaitHealthy if dispatch needs the live -// list to converge, or invoke p.SetHealthy to bypass refresh entirely. +// status to StatusPassing. The pool is dispatch-ready when it returns: status +// changes reach a pool synchronously, so no wait is needed. func NewHealthy(handlers []http.Handler) (pool.Pool, []*pool.Target, []*healthcheck.Status, ) { diff --git a/testdata/ignore_words.txt b/testdata/ignore_words.txt index 3f8c22fae..7b7880490 100644 --- a/testdata/ignore_words.txt +++ b/testdata/ignore_words.txt @@ -4,4 +4,5 @@ flate ro alo crossreference -aks \ No newline at end of file +aks +worstCase \ No newline at end of file