diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..d6bb74a --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,28 @@ +name: ci +on: + push: + branches: [main] + pull_request: + +jobs: + build-test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.26" + - run: go build ./... + - run: go vet ./... + - run: go test -race ./... + + lint: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-go@v5 + with: + go-version: "1.26" + - uses: golangci/golangci-lint-action@v6 + with: + version: latest diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..f011453 --- /dev/null +++ b/.gitignore @@ -0,0 +1,9 @@ +/bin/ +/dist/ +*.out +cover.out +coverage.* +robin + +# Internal design notes — kept locally, not published. +docs/design.md diff --git a/.golangci.yml b/.golangci.yml new file mode 100644 index 0000000..1f431fb --- /dev/null +++ b/.golangci.yml @@ -0,0 +1,9 @@ +version: "2" +linters: + enable: + - errcheck + - govet + - ineffassign + - staticcheck + - unused + - revive diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..9e1a4d6 --- /dev/null +++ b/Makefile @@ -0,0 +1,41 @@ +BINARY := robin +PKG := github.com/snangue/robin +VERSION ?= $(shell git describe --tags --always --dirty 2>/dev/null || echo dev) +COMMIT ?= $(shell git rev-parse --short HEAD 2>/dev/null || echo none) +DATE ?= $(shell date -u +%Y-%m-%dT%H:%M:%SZ) +LDFLAGS := -s -w \ + -X $(PKG)/internal/version.Version=$(VERSION) \ + -X $(PKG)/internal/version.Commit=$(COMMIT) \ + -X $(PKG)/internal/version.Date=$(DATE) + +.PHONY: build test race cover vet lint tidy docker clean run-file + +build: + CGO_ENABLED=0 go build -trimpath -ldflags "$(LDFLAGS)" -o bin/$(BINARY) ./cmd/robin + +test: + go test ./... + +race: + go test -race ./... + +cover: + go test -coverprofile=cover.out ./... && go tool cover -func=cover.out + +vet: + go vet ./... + +lint: + @command -v golangci-lint >/dev/null 2>&1 && golangci-lint run || echo "golangci-lint not installed; skipping" + +tidy: + go mod tidy + +docker: + docker build -t $(BINARY):$(VERSION) -f deploy/Dockerfile . + +clean: + rm -rf bin cover.out + +run-file: + go run ./cmd/robin diff --git a/README.md b/README.md new file mode 100644 index 0000000..ef50562 --- /dev/null +++ b/README.md @@ -0,0 +1,128 @@ +# Robin + +**A workload-identity injector.** Robin runs beside a workload, sources its short-lived, rotating identity, and injects it as a fresh bearer token on every outbound request — so the application never has to know the credential exists, let alone that it rotates. + +## The problem + +Kubernetes increasingly hands workloads **short-lived, rotating credentials**. Projected ServiceAccount tokens roll over roughly hourly; SPIFFE JWT-SVIDs every few minutes. That is exactly what you want for security — a leaked token expires on its own, fast. + +The catch: **almost no application knows how to consume one.** An app reads its API key or bearer token *once* — from an environment variable, a config field, an `Authorization` header set when its HTTP client is constructed — and then holds that value for its entire lifetime. There is no hook to reload it, no callback when it rotates. Point such an app at a rotating token and it keeps presenting the stale one; minutes (or an hour) later, every request starts failing with `401`. + +So teams fall back to the very thing rotation was meant to eliminate: a **long-lived static secret**, mounted into the app and held there indefinitely. Now the application *is* part of the credential plane — it holds a real secret, and rotating that secret means redeploying. + +Robin breaks the coupling. It sits next to the workload, sources the rotating identity from the platform, and injects a **fresh** bearer on every request. The application points its egress at Robin and carries **no credential at all** — not even a rotating one. It presents *identity*; Robin keeps it current; a downstream broker turns that identity into the real upstream credential. + +## How it works + +![Architecture: robin flow](./docs/robin-flow.png) + +- The **App** makes ordinary HTTP requests to Robin on loopback. No API key, no token-reload logic — at most an inert placeholder header, which Robin overwrites. +- **Robin** is a streaming reverse proxy with a pluggable identity provider. It resolves the workload's native token (always current), sets `Authorization: Bearer`, and forwards. It is body-agnostic and never reads or buffers the request body, so streaming responses pass straight through. +- The **Broker** — any identity-aware egress gateway — validates the presented identity and applies the real upstream credential. Robin itself never mints, exchanges, signs, or federates, so the machine holds **no real provider credential** — universally, for every provider, with no exceptions. + +**Threat model in one line:** anything on the pod's loopback can ask Robin to present the workload's identity — so the proxy plane defaults to loopback-only, with a Unix-domain-socket + `SO_PEERCRED` peer-credential mode for hardened deployments. + +## Quickstart + +**Build:** + +```sh +make build # -> bin/robin (static, CGO-free) +make test # go test ./... +``` + +**Run locally** (file provider, pointing at any broker/echo endpoint): + +```sh +ROBIN_UPSTREAM_URL=https://broker.example:8443 \ +ROBIN_TOKEN_SOURCE=file \ +ROBIN_TOKEN_FILE=/var/run/secrets/tokens/broker-token \ + bin/robin +# app egress -> http://127.0.0.1:4000 ; probe http://:4001/healthz +``` + +**Container image:** + +```sh +docker build -t robin:0.1.0 -f deploy/Dockerfile . # ~14MB distroless, nonroot +``` + +**Kubernetes native sidecar:** see [`deploy/k8s/sidecar-example.yaml`](deploy/k8s/sidecar-example.yaml) — Robin runs as an `initContainer` with `restartPolicy: Always` (K8s 1.29+), the app points its egress at `127.0.0.1:4000`, and probes hit the admin plane on `:4001`. + +## Identity providers + +| Source | `ROBIN_TOKEN_SOURCE` | Rotation handling | +|--------|----------------------|-------------------| +| Kubernetes projected ServiceAccount token | `file` | the kubelet rotates the file in place (~80% TTL); Robin **re-reads per request**, so it never serves a stale token. | +| SPIFFE JWT-SVID (Workload API) | `jwtsvid` | not pushed (unary fetch), so Robin caches and **refreshes ahead of `exp`**, and serves a still-valid cached token if the agent briefly fails. | + +Any other source that yields the workload's own short-lived OIDC/JWT identity fits the same shape: fetch it, forward it, let the broker validate. + +## Configuration + +Config is a flat set of `ROBIN_`-prefixed scalars — **no config language**. The primary plane is **environment variables** (idiomatic for sidecars and systemd units); an optional flat `.env`-style `KEY=value` file is supported for standalone hosts. Precedence: **flags > environment > file**. + +| Var | Default | Notes | +|-----|---------|-------| +| `ROBIN_UPSTREAM_URL` | (required) | broker base URL | +| `ROBIN_TOKEN_SOURCE` | `file` | `file` \| `jwtsvid` | +| `ROBIN_LISTEN_ADDR` | `127.0.0.1:4000` | proxy plane (loopback by default) | +| `ROBIN_LISTEN_UDS` | — | UDS path; enables peer-cred mode | +| `ROBIN_ADMIN_ADDR` | `:4001` | health/readiness/metrics plane | +| `ROBIN_TOKEN_FILE` | `/var/run/secrets/tokens/token` | `file` provider | +| `ROBIN_AUDIENCE` | — | required for `jwtsvid`; **must match the broker** | +| `ROBIN_SPIFFE_SOCKET` | — | `jwtsvid` socket addr (optional; falls back to the go-spiffe default) | +| `ROBIN_SVID_REFRESH_BEFORE` | `60s` | refresh ahead of `exp` (clamped ≤ ½ the observed lifetime) | +| `ROBIN_UPSTREAM_CA_FILE` | — | verify broker TLS | +| `ROBIN_PEERCRED_ALLOW_UIDS` | — | comma-separated UIDs; empty = allow any local peer | + +> **Bind address:** the proxy plane defaults to `127.0.0.1:4000` (loopback); the admin plane defaults to `:4001` (all interfaces) so kubelet probes can reach it — restrict `:4001` ingress with a NetworkPolicy where the platform allows it. + +> **Audience must match end to end.** A mismatch between the token's audience and the broker's expected audience is a hard reject — for projected tokens and SVIDs alike. + +## Deployment topologies + +Same binary; the topology determines how identity is *sourced*, not what Robin does with it. + +- **Native sidecar (default).** An init container with `restartPolicy: Always` (Kubernetes 1.29+) so Robin starts before the app container (no first-call race) and stops after it (no in-flight-egress loss). The workload reaches Robin on `localhost`. +- **Standalone systemd unit.** Robin runs as a host/VM service; identity comes from a node-level SPIRE agent (`jwtsvid`). +- **Per-node DaemonSet** — *advanced/optional.* Fewer instances, but loses per-pod identity fidelity unless SPIRE does per-pod attestation. +- **Standalone egress service** — *generally an anti-pattern.* Loses transparent localhost injection and per-pod identity. + +## Admin endpoints + +Served on a **separate admin listener** (`ROBIN_ADMIN_ADDR`, default `:4001`) — *not* the proxy port, because the proxy forwards every path to the broker (a probe there would be proxied upstream and could leak identity). + +- `GET /healthz` — liveness (always 200). +- `GET /readyz` — readiness; 200 only when an identity can actually be resolved. +- `GET /metrics` — Prometheus exposition (token fetch latency, cache hit/refresh, served-stale, upstream status codes). *(planned)* + +## Failure semantics + +| Condition | Response | +|-----------|----------| +| Identity unavailable (token file missing/empty, Workload API down) | **503** — never forwards a missing/placeholder credential (jwtsvid serves a still-valid cached token first) | +| Broker unreachable | **502** — single attempt, no blind retry | +| Audience mismatch | fail closed with an explicit log (config error, not transient) | + +## Security notes + +- The token value is **never logged** — redaction is structural; it never enters a log record. +- The container runs **nonroot** on a static distroless base. +- Bearer-only by design (forwards a JWT-SVID / projected token). mTLS with an X.509-SVID (proof-of-possession) was considered and deliberately deferred. +- The proxy plane binds loopback by default; a Unix-domain-socket + `SO_PEERCRED` UID allowlist hardens the local trust boundary further. + +## When *not* to use Robin + +- **You control the client's auth path.** An in-process `RoundTripper` / auth hook that fetches the identity is lighter than a proxy — no extra process, no localhost trust boundary. Robin exists for apps you *can't* teach to rotate. +- **You already run a service mesh** (Istio ambient / ztunnel / Cilium). Use its egress identity origination instead of adding a per-pod sidecar. + +## Roadmap + +- **v0.1 (core):** `file` + `jwtsvid` providers, native-sidecar deployment, loopback proxy with broker-forward, `/healthz` + `/readyz`, structured logging, container image. +- **v0.2 (hardening):** `SO_PEERCRED` enforcement on the UDS path, Prometheus `/metrics`, standalone systemd deployment. +- **Future:** more native-OIDC identity sources (Azure AD Workload Identity, GCP Workload Identity Federation, …) behind the same provider interface — Robin stays a generic forwarder; the broker still owns validation and credential minting. + +## License + +[MPL-2.0](LICENSE). diff --git a/deploy/Dockerfile b/deploy/Dockerfile new file mode 100644 index 0000000..c1c0cfb --- /dev/null +++ b/deploy/Dockerfile @@ -0,0 +1,27 @@ +# syntax=docker/dockerfile:1 + +# --- build stage: static, CGO-free binary --- +FROM golang:1.26 AS build +WORKDIR /src +COPY go.mod go.sum ./ +RUN go mod download +COPY . . +ARG VERSION=dev +ARG COMMIT=none +ARG DATE=unknown +RUN CGO_ENABLED=0 GOOS=linux go build \ + -trimpath \ + -ldflags="-s -w \ + -X github.com/snangue/robin/internal/version.Version=${VERSION} \ + -X github.com/snangue/robin/internal/version.Commit=${COMMIT} \ + -X github.com/snangue/robin/internal/version.Date=${DATE}" \ + -o /out/robin ./cmd/robin + +# --- runtime stage: distroless static, non-root --- +# distroless/static ships CA roots and /etc/passwd; the binary is fully static. +FROM gcr.io/distroless/static:nonroot +COPY --from=build /out/robin /robin +USER nonroot:nonroot +# 4000 = proxy plane (bind loopback in a sidecar); 4001 = admin plane (probes/metrics). +EXPOSE 4000 4001 +ENTRYPOINT ["/robin"] diff --git a/deploy/k8s/sidecar-example.yaml b/deploy/k8s/sidecar-example.yaml new file mode 100644 index 0000000..d6dae96 --- /dev/null +++ b/deploy/k8s/sidecar-example.yaml @@ -0,0 +1,93 @@ +# Robin as a native sidecar (Kubernetes 1.29+). +# +# Robin runs as an initContainer with restartPolicy: Always, so it starts before +# the app container (no first-call race) and terminates after it (no in-flight +# egress loss). The app sends egress to Robin on loopback; Robin injects the +# pod's projected ServiceAccount token as a bearer and forwards to the broker. +apiVersion: apps/v1 +kind: Deployment +metadata: + name: myapp-with-robin + labels: + app: myapp +spec: + replicas: 1 + selector: + matchLabels: + app: myapp + template: + metadata: + labels: + app: myapp + spec: + serviceAccountName: myapp + terminationGracePeriodSeconds: 30 # > Robin's 25s in-flight drain budget + containers: + - name: myapp + image: ghcr.io/example/myapp:latest + env: + # Point the app's egress at Robin on loopback instead of the upstream. + - name: UPSTREAM_BASE_URL + value: "http://127.0.0.1:4000" + # Many apps require *some* credential; this one is inert — Robin + # overwrites the Authorization header with the workload's identity. + - name: UPSTREAM_API_KEY + value: "unused-placeholder" + initContainers: + - name: robin + image: ghcr.io/snangue/robin:0.1.0 + restartPolicy: Always # native sidecar: starts first, stops last + args: [] + env: + - name: ROBIN_UPSTREAM_URL + value: "https://broker.example.svc:8443" + - name: ROBIN_TOKEN_SOURCE + value: "file" + # Proxy plane: loopback only — only this pod's app may ask for injection. + - name: ROBIN_LISTEN_ADDR + value: "127.0.0.1:4000" + # Admin plane: all interfaces, so the kubelet can reach the probes. + # It serves only health/readiness/metrics (never identity); restrict + # :4001 ingress with a NetworkPolicy where the platform allows it. + - name: ROBIN_ADMIN_ADDR + value: ":4001" + - name: ROBIN_TOKEN_FILE + value: "/var/run/secrets/tokens/broker-token" + volumeMounts: + - name: broker-token + mountPath: /var/run/secrets/tokens + readOnly: true + livenessProbe: + httpGet: + path: /healthz + port: 4001 + readinessProbe: + httpGet: + path: /readyz + port: 4001 + resources: + requests: + cpu: 10m + memory: 16Mi + limits: + memory: 64Mi + securityContext: + allowPrivilegeEscalation: false + readOnlyRootFilesystem: true + runAsNonRoot: true + capabilities: + drop: ["ALL"] + volumes: + - name: broker-token + projected: + sources: + - serviceAccountToken: + path: broker-token + # Must match the audience the broker validates. + audience: "broker.example" + expirationSeconds: 3600 +--- +apiVersion: v1 +kind: ServiceAccount +metadata: + name: myapp diff --git a/docs/robin-flow.png b/docs/robin-flow.png new file mode 100644 index 0000000..ae632d2 Binary files /dev/null and b/docs/robin-flow.png differ diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..9f709a7 --- /dev/null +++ b/go.mod @@ -0,0 +1,16 @@ +module github.com/snangue/robin + +go 1.26 + +require github.com/spiffe/go-spiffe/v2 v2.8.0 + +require ( + github.com/Microsoft/go-winio v0.6.2 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 // indirect + golang.org/x/net v0.48.0 // indirect + golang.org/x/sys v0.39.0 // indirect + golang.org/x/text v0.32.0 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20251202230838-ff82c1b0f217 // indirect + google.golang.org/grpc v1.79.3 // indirect + google.golang.org/protobuf v1.36.11 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..1d66d32 --- /dev/null +++ b/go.sum @@ -0,0 +1,52 @@ +github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY= +github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA= +github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= +github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= +github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= +github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= +github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= +github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek= +github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/spiffe/go-spiffe/v2 v2.8.0 h1:vHCTEZYhpXZ9y6JkIouIdHLJobWGUFn2467/WsXHHjA= +github.com/spiffe/go-spiffe/v2 v2.8.0/go.mod h1:47Q0Q9/AqGha8QLHp+kxpH4Wca7X7EnOtlIJy3mxZ3U= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64= +go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y= +go.opentelemetry.io/otel v1.39.0 h1:8yPrr/S0ND9QEfTfdP9V+SiwT4E0G7Y5MO7p85nis48= +go.opentelemetry.io/otel v1.39.0/go.mod h1:kLlFTywNWrFyEdH0oj2xK0bFYZtHRYUdv1NklR/tgc8= +go.opentelemetry.io/otel/metric v1.39.0 h1:d1UzonvEZriVfpNKEVmHXbdf909uGTOQjA0HF0Ls5Q0= +go.opentelemetry.io/otel/metric v1.39.0/go.mod h1:jrZSWL33sD7bBxg1xjrqyDjnuzTUB0x1nBERXd7Ftcs= +go.opentelemetry.io/otel/sdk v1.39.0 h1:nMLYcjVsvdui1B/4FRkwjzoRVsMK8uL/cj0OyhKzt18= +go.opentelemetry.io/otel/sdk v1.39.0/go.mod h1:vDojkC4/jsTJsE+kh+LXYQlbL8CgrEcwmt1ENZszdJE= +go.opentelemetry.io/otel/sdk/metric v1.39.0 h1:cXMVVFVgsIf2YL6QkRF4Urbr/aMInf+2WKg+sEJTtB8= +go.opentelemetry.io/otel/sdk/metric v1.39.0/go.mod h1:xq9HEVH7qeX69/JnwEfp6fVq5wosJsY1mt4lLfYdVew= +go.opentelemetry.io/otel/trace v1.39.0 h1:2d2vfpEDmCJ5zVYz7ijaJdOF59xLomrvj7bjt6/qCJI= +go.opentelemetry.io/otel/trace v1.39.0/go.mod h1:88w4/PnZSazkGzz/w84VHpQafiU4EtqqlVdxWy+rNOA= +golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU= +golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY= +golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk= +golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= +golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU= +golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY= +gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk= +gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E= +google.golang.org/genproto/googleapis/rpc v0.0.0-20251202230838-ff82c1b0f217 h1:gRkg/vSppuSQoDjxyiGfN4Upv/h/DQmIR10ZU8dh4Ww= +google.golang.org/genproto/googleapis/rpc v0.0.0-20251202230838-ff82c1b0f217/go.mod h1:7i2o+ce6H/6BluujYR+kqX3GKH+dChPTQU19wjRPiGk= +google.golang.org/grpc v1.79.3 h1:sybAEdRIEtvcD68Gx7dmnwjZKlyfuc61Dyo9pGXXkKE= +google.golang.org/grpc v1.79.3/go.mod h1:KmT0Kjez+0dde/v2j9vzwoAScgEPx/Bw1CYChhHLrHQ= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..9078292 --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,174 @@ +// Package config loads Robin's flat ROBIN_* configuration from flags, +// environment, and an optional .env-style file (precedence: flags > env > file). +package config + +import ( + "flag" + "fmt" + "io" + "log/slog" + "net/url" + "os" + "strconv" + "strings" + "time" + + "github.com/snangue/robin/internal/obs" +) + +// Config is Robin's fully-resolved configuration. +type Config struct { + UpstreamURL string // ROBIN_UPSTREAM_URL (required) — broker base URL + TokenSource string // ROBIN_TOKEN_SOURCE — file | jwtsvid + ListenAddr string // ROBIN_LISTEN_ADDR — proxy listener + ListenUDS string // ROBIN_LISTEN_UDS — UDS path (enables peer-cred mode) + AdminAddr string // ROBIN_ADMIN_ADDR — admin/health/metrics listener + TokenFile string // ROBIN_TOKEN_FILE — file provider token path + Audience string // ROBIN_AUDIENCE — required for jwtsvid + SPIFFESocket string // ROBIN_SPIFFE_SOCKET — optional Workload API socket addr + SVIDRefreshBefore time.Duration // ROBIN_SVID_REFRESH_BEFORE — refresh lead before exp + UpstreamCAFile string // ROBIN_UPSTREAM_CA_FILE — CA to verify broker TLS + PeerCredAllowUIDs []int // ROBIN_PEERCRED_ALLOW_UIDS — allowed peer UIDs (UDS mode) + LogLevel slog.Level // ROBIN_LOG_LEVEL +} + +// Load resolves configuration with precedence flags > env > file. getenv and +// args are injected so precedence is testable without touching process state. +func Load(args []string, getenv func(string) string) (Config, error) { + if getenv == nil { + getenv = os.Getenv + } + + path := configPath(args, getenv) + fileMap := map[string]string{} + if path != "" { + f, err := os.Open(path) + if err != nil { + return Config{}, fmt.Errorf("config: open %s: %w", path, err) + } + defer f.Close() + fileMap, err = parseDotenv(f) + if err != nil { + return Config{}, fmt.Errorf("config: parse %s: %w", path, err) + } + } + + // val resolves a key by env first, then file, then the provided default. + val := func(key, def string) string { + if v := getenv(key); v != "" { + return v + } + if v, ok := fileMap[key]; ok && v != "" { + return v + } + return def + } + + var ( + cfg Config + refreshStr, uidStr, levelStr string + cfgPath string + ) + + fs := flag.NewFlagSet("robin", flag.ContinueOnError) + fs.SetOutput(io.Discard) + // --config is resolved before flag parsing (see configPath) so the file is + // already loaded; it is registered here only so Parse accepts the flag. + fs.StringVar(&cfgPath, "config", path, "path to a flat KEY=value config file") + fs.StringVar(&cfg.UpstreamURL, "upstream-url", val("ROBIN_UPSTREAM_URL", ""), "broker base URL (required)") + fs.StringVar(&cfg.TokenSource, "token-source", val("ROBIN_TOKEN_SOURCE", "file"), "identity source: file|jwtsvid") + fs.StringVar(&cfg.ListenAddr, "listen-addr", val("ROBIN_LISTEN_ADDR", "127.0.0.1:4000"), "proxy listen address") + fs.StringVar(&cfg.ListenUDS, "listen-uds", val("ROBIN_LISTEN_UDS", ""), "proxy UDS path (enables peer-cred mode)") + fs.StringVar(&cfg.AdminAddr, "admin-addr", val("ROBIN_ADMIN_ADDR", ":4001"), "admin listen address") + fs.StringVar(&cfg.TokenFile, "token-file", val("ROBIN_TOKEN_FILE", "/var/run/secrets/tokens/token"), "file provider token path") + fs.StringVar(&cfg.Audience, "audience", val("ROBIN_AUDIENCE", ""), "token audience (required for jwtsvid)") + fs.StringVar(&cfg.SPIFFESocket, "spiffe-socket", val("ROBIN_SPIFFE_SOCKET", ""), "SPIFFE Workload API socket address") + fs.StringVar(&refreshStr, "svid-refresh-before", val("ROBIN_SVID_REFRESH_BEFORE", "60s"), "refresh SVID this long before expiry") + fs.StringVar(&cfg.UpstreamCAFile, "upstream-ca-file", val("ROBIN_UPSTREAM_CA_FILE", ""), "CA file to verify broker TLS") + fs.StringVar(&uidStr, "peercred-allow-uids", val("ROBIN_PEERCRED_ALLOW_UIDS", ""), "comma-separated allowed peer UIDs") + fs.StringVar(&levelStr, "log-level", val("ROBIN_LOG_LEVEL", "info"), "log level: debug|info|warn|error") + + if err := fs.Parse(args); err != nil { + return Config{}, fmt.Errorf("config: %w", err) + } + + d, err := time.ParseDuration(refreshStr) + if err != nil { + return Config{}, fmt.Errorf("config: invalid svid-refresh-before %q: %w", refreshStr, err) + } + cfg.SVIDRefreshBefore = d + + uids, err := parseUIDs(uidStr) + if err != nil { + return Config{}, fmt.Errorf("config: invalid peercred-allow-uids %q: %w", uidStr, err) + } + cfg.PeerCredAllowUIDs = uids + + cfg.LogLevel = obs.ParseLevel(levelStr) + + return cfg, cfg.Validate() +} + +// Validate enforces fail-closed configuration rules. +func (c Config) Validate() error { + if c.UpstreamURL == "" { + return fmt.Errorf("config: ROBIN_UPSTREAM_URL is required") + } + if u, err := url.Parse(c.UpstreamURL); err != nil || !u.IsAbs() || u.Host == "" { + return fmt.Errorf("config: ROBIN_UPSTREAM_URL %q must be an absolute URL", c.UpstreamURL) + } + switch c.TokenSource { + case "file": + // Token file is read at request time; nothing else required here. + case "jwtsvid": + if c.Audience == "" { + return fmt.Errorf("config: ROBIN_AUDIENCE is required for token-source=jwtsvid") + } + default: + return fmt.Errorf("config: ROBIN_TOKEN_SOURCE %q must be file or jwtsvid", c.TokenSource) + } + if c.SVIDRefreshBefore <= 0 { + return fmt.Errorf("config: ROBIN_SVID_REFRESH_BEFORE must be positive") + } + if len(c.PeerCredAllowUIDs) > 0 && c.ListenUDS == "" { + return fmt.Errorf("config: ROBIN_PEERCRED_ALLOW_UIDS set but ROBIN_LISTEN_UDS is empty") + } + return nil +} + +// configPath resolves the config-file path from a --config flag or ROBIN_CONFIG. +func configPath(args []string, getenv func(string) string) string { + for i, a := range args { + switch { + case a == "--config" || a == "-config": + if i+1 < len(args) { + return args[i+1] + } + case strings.HasPrefix(a, "--config="): + return strings.TrimPrefix(a, "--config=") + case strings.HasPrefix(a, "-config="): + return strings.TrimPrefix(a, "-config=") + } + } + return getenv("ROBIN_CONFIG") +} + +func parseUIDs(s string) ([]int, error) { + s = strings.TrimSpace(s) + if s == "" { + return nil, nil + } + var uids []int + for _, p := range strings.Split(s, ",") { + p = strings.TrimSpace(p) + if p == "" { + continue + } + n, err := strconv.Atoi(p) + if err != nil { + return nil, fmt.Errorf("%q is not a valid uid", p) + } + uids = append(uids, n) + } + return uids, nil +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..4588598 --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,103 @@ +package config + +import ( + "os" + "path/filepath" + "testing" +) + +func mapEnv(m map[string]string) func(string) string { + return func(k string) string { return m[k] } +} + +func TestLoadPrecedence(t *testing.T) { + dir := t.TempDir() + envFile := filepath.Join(dir, "robin.env") + if err := os.WriteFile(envFile, []byte("ROBIN_UPSTREAM_URL=https://file.example\nROBIN_TOKEN_SOURCE=file\n"), 0o600); err != nil { + t.Fatal(err) + } + + tests := []struct { + name string + args []string + env map[string]string + want string + }{ + {"file only", nil, map[string]string{"ROBIN_CONFIG": envFile}, "https://file.example"}, + {"env over file", nil, map[string]string{"ROBIN_CONFIG": envFile, "ROBIN_UPSTREAM_URL": "https://env.example"}, "https://env.example"}, + {"flag over env", []string{"--upstream-url=https://flag.example"}, map[string]string{"ROBIN_CONFIG": envFile, "ROBIN_UPSTREAM_URL": "https://env.example"}, "https://flag.example"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg, err := Load(tt.args, mapEnv(tt.env)) + if err != nil { + t.Fatalf("Load: %v", err) + } + if cfg.UpstreamURL != tt.want { + t.Errorf("UpstreamURL = %q, want %q", cfg.UpstreamURL, tt.want) + } + }) + } +} + +func TestLoadDefaults(t *testing.T) { + cfg, err := Load(nil, mapEnv(map[string]string{"ROBIN_UPSTREAM_URL": "https://b.example"})) + if err != nil { + t.Fatal(err) + } + if cfg.TokenSource != "file" || cfg.ListenAddr != "127.0.0.1:4000" || cfg.AdminAddr != ":4001" { + t.Errorf("unexpected defaults: %+v", cfg) + } + if cfg.SVIDRefreshBefore.Seconds() != 60 { + t.Errorf("SVIDRefreshBefore = %v, want 60s", cfg.SVIDRefreshBefore) + } +} + +func TestLoadValidation(t *testing.T) { + tests := []struct { + name string + env map[string]string + wantErr bool + }{ + {"missing upstream", nil, true}, + {"bad upstream url", map[string]string{"ROBIN_UPSTREAM_URL": "not-a-url"}, true}, + {"jwtsvid without audience", map[string]string{"ROBIN_UPSTREAM_URL": "https://b", "ROBIN_TOKEN_SOURCE": "jwtsvid"}, true}, + {"unknown token source", map[string]string{"ROBIN_UPSTREAM_URL": "https://b", "ROBIN_TOKEN_SOURCE": "bogus"}, true}, + {"valid file", map[string]string{"ROBIN_UPSTREAM_URL": "https://b.example"}, false}, + {"valid jwtsvid", map[string]string{"ROBIN_UPSTREAM_URL": "https://b.example", "ROBIN_TOKEN_SOURCE": "jwtsvid", "ROBIN_AUDIENCE": "aud"}, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := Load(nil, mapEnv(tt.env)) + if (err != nil) != tt.wantErr { + t.Errorf("Load err = %v, wantErr %v", err, tt.wantErr) + } + }) + } +} + +func TestParseUIDs(t *testing.T) { + got, err := parseUIDs(" 1000, 1001 ,1002") + if err != nil { + t.Fatal(err) + } + if len(got) != 3 || got[0] != 1000 || got[2] != 1002 { + t.Errorf("parseUIDs = %v", got) + } + if _, err := parseUIDs("abc"); err == nil { + t.Error("expected error for non-numeric uid") + } + if got, _ := parseUIDs(""); got != nil { + t.Errorf("empty should be nil, got %v", got) + } +} + +func TestPeerCredRequiresUDS(t *testing.T) { + _, err := Load(nil, mapEnv(map[string]string{ + "ROBIN_UPSTREAM_URL": "https://b.example", + "ROBIN_PEERCRED_ALLOW_UIDS": "1000", + })) + if err == nil { + t.Error("expected error: peercred uids without UDS") + } +} diff --git a/internal/config/dotenv.go b/internal/config/dotenv.go new file mode 100644 index 0000000..30eb231 --- /dev/null +++ b/internal/config/dotenv.go @@ -0,0 +1,45 @@ +package config + +import ( + "bufio" + "fmt" + "io" + "strings" +) + +// parseDotenv parses a flat KEY=value file: blank lines and #-comments are +// skipped, an optional leading "export " is stripped, and a single layer of +// surrounding single/double quotes is removed. No interpolation, no nesting. +func parseDotenv(r io.Reader) (map[string]string, error) { + out := map[string]string{} + sc := bufio.NewScanner(r) + for line := 1; sc.Scan(); line++ { + raw := strings.TrimSpace(sc.Text()) + if raw == "" || strings.HasPrefix(raw, "#") { + continue + } + raw = strings.TrimPrefix(raw, "export ") + eq := strings.IndexByte(raw, '=') + if eq < 0 { + return nil, fmt.Errorf("line %d: missing '='", line) + } + key := strings.TrimSpace(raw[:eq]) + if key == "" { + return nil, fmt.Errorf("line %d: empty key", line) + } + out[key] = unquote(strings.TrimSpace(raw[eq+1:])) + } + if err := sc.Err(); err != nil { + return nil, err + } + return out, nil +} + +func unquote(s string) string { + if len(s) >= 2 { + if c := s[0]; (c == '"' || c == '\'') && s[len(s)-1] == c { + return s[1 : len(s)-1] + } + } + return s +} diff --git a/internal/config/dotenv_test.go b/internal/config/dotenv_test.go new file mode 100644 index 0000000..a6c1b30 --- /dev/null +++ b/internal/config/dotenv_test.go @@ -0,0 +1,41 @@ +package config + +import ( + "strings" + "testing" +) + +func TestParseDotenv(t *testing.T) { + in := strings.Join([]string{ + "# a comment", + "", + "ROBIN_UPSTREAM_URL=https://b.example", + "export ROBIN_TOKEN_SOURCE=file", + `ROBIN_AUDIENCE="quoted-aud"`, + "ROBIN_LISTEN_ADDR='127.0.0.1:4000'", + "ROBIN_TOKEN_FILE=/var/run/secrets/tokens/token=weird", + }, "\n") + + m, err := parseDotenv(strings.NewReader(in)) + if err != nil { + t.Fatal(err) + } + want := map[string]string{ + "ROBIN_UPSTREAM_URL": "https://b.example", + "ROBIN_TOKEN_SOURCE": "file", + "ROBIN_AUDIENCE": "quoted-aud", + "ROBIN_LISTEN_ADDR": "127.0.0.1:4000", + "ROBIN_TOKEN_FILE": "/var/run/secrets/tokens/token=weird", + } + for k, v := range want { + if m[k] != v { + t.Errorf("%s = %q, want %q", k, m[k], v) + } + } +} + +func TestParseDotenvError(t *testing.T) { + if _, err := parseDotenv(strings.NewReader("NOEQUALS")); err == nil { + t.Error("expected error for line missing '='") + } +} diff --git a/internal/identity/factory.go b/internal/identity/factory.go new file mode 100644 index 0000000..10d1524 --- /dev/null +++ b/internal/identity/factory.go @@ -0,0 +1,24 @@ +package identity + +import ( + "context" + "fmt" + + "github.com/snangue/robin/internal/config" +) + +// New builds the identity provider selected by cfg.TokenSource. +func New(ctx context.Context, cfg config.Config) (Provider, error) { + switch cfg.TokenSource { + case "file": + return NewFileProvider(cfg.TokenFile), nil + case "jwtsvid": + fetcher, err := newJWTSVIDSource(ctx, cfg) + if err != nil { + return nil, err + } + return newJWTSVIDProviderWith(fetcher, cfg.Audience, cfg.SVIDRefreshBefore), nil + default: + return nil, fmt.Errorf("identity: unknown token source %q", cfg.TokenSource) + } +} diff --git a/internal/identity/file.go b/internal/identity/file.go new file mode 100644 index 0000000..e177068 --- /dev/null +++ b/internal/identity/file.go @@ -0,0 +1,36 @@ +package identity + +import ( + "context" + "fmt" + "os" + "strings" +) + +// FileProvider serves a Kubernetes projected ServiceAccount token from a file. +// The kubelet atomically rotates the file in place (~80% TTL), so the provider +// re-reads it on every request and never caches. +type FileProvider struct { + path string +} + +// NewFileProvider returns a FileProvider reading the token at path. +func NewFileProvider(path string) *FileProvider { + return &FileProvider{path: path} +} + +// Token reads, trims, and returns the token file contents. +func (p *FileProvider) Token(_ context.Context) (string, error) { + b, err := os.ReadFile(p.path) + if err != nil { + return "", fmt.Errorf("identity/file: read %s: %w", p.path, err) + } + tok := strings.TrimSpace(string(b)) + if tok == "" { + return "", ErrNoToken + } + return tok, nil +} + +// Close is a no-op; the file provider holds no resources. +func (p *FileProvider) Close() error { return nil } diff --git a/internal/identity/file_test.go b/internal/identity/file_test.go new file mode 100644 index 0000000..d039296 --- /dev/null +++ b/internal/identity/file_test.go @@ -0,0 +1,56 @@ +package identity + +import ( + "context" + "errors" + "os" + "path/filepath" + "testing" +) + +func TestFileProviderReadAndTrim(t *testing.T) { + path := filepath.Join(t.TempDir(), "token") + if err := os.WriteFile(path, []byte(" abc.def.ghi\n"), 0o600); err != nil { + t.Fatal(err) + } + tok, err := NewFileProvider(path).Token(context.Background()) + if err != nil { + t.Fatal(err) + } + if tok != "abc.def.ghi" { + t.Errorf("token = %q, want trimmed abc.def.ghi", tok) + } +} + +func TestFileProviderRotation(t *testing.T) { + path := filepath.Join(t.TempDir(), "token") + if err := os.WriteFile(path, []byte("first"), 0o600); err != nil { + t.Fatal(err) + } + p := NewFileProvider(path) + first, _ := p.Token(context.Background()) + if err := os.WriteFile(path, []byte("second"), 0o600); err != nil { // simulate kubelet swap + t.Fatal(err) + } + second, _ := p.Token(context.Background()) + if first != "first" || second != "second" { + t.Errorf("expected per-request re-read, got %q then %q", first, second) + } +} + +func TestFileProviderEmpty(t *testing.T) { + path := filepath.Join(t.TempDir(), "token") + if err := os.WriteFile(path, []byte(" \n"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := NewFileProvider(path).Token(context.Background()); !errors.Is(err, ErrNoToken) { + t.Errorf("want ErrNoToken, got %v", err) + } +} + +func TestFileProviderMissing(t *testing.T) { + p := NewFileProvider(filepath.Join(t.TempDir(), "nope")) + if _, err := p.Token(context.Background()); err == nil { + t.Error("want error for missing file") + } +} diff --git a/internal/identity/jwtsvid.go b/internal/identity/jwtsvid.go new file mode 100644 index 0000000..49e6d6d --- /dev/null +++ b/internal/identity/jwtsvid.go @@ -0,0 +1,148 @@ +package identity + +import ( + "context" + "fmt" + "sync" + "time" +) + +// jwtFetcher is the minimal seam over the SPIFFE Workload API. The production +// adapter (spiffeFetcher) wraps a *workloadapi.JWTSource; tests inject a fake, +// so the cache / refresh / serve-stale logic needs no running SPIRE agent. +type jwtFetcher interface { + fetch(ctx context.Context, audience string) (token string, expiry time.Time, err error) + close() error +} + +type cachedSVID struct { + token string + expiry time.Time + refreshAt time.Time +} + +type refreshCall struct { + done chan struct{} + token string + err error +} + +// JWTSVIDProvider serves a SPIFFE JWT-SVID, refreshing it ahead of expiry and +// serving a still-valid cached token when the Workload API is briefly down. +// FetchJWTSVID is unary (not pushed), so the provider owns the refresh policy. +type JWTSVIDProvider struct { + fetcher jwtFetcher + audience string + refreshBefore time.Duration + now func() time.Time // injectable clock for tests + + mu sync.Mutex + cache map[string]cachedSVID + inflight map[string]*refreshCall +} + +func newJWTSVIDProviderWith(f jwtFetcher, audience string, refreshBefore time.Duration) *JWTSVIDProvider { + return &JWTSVIDProvider{ + fetcher: f, + audience: audience, + refreshBefore: refreshBefore, + now: time.Now, + cache: map[string]cachedSVID{}, + inflight: map[string]*refreshCall{}, + } +} + +// Token returns a valid JWT-SVID for the configured audience. +func (p *JWTSVIDProvider) Token(ctx context.Context) (string, error) { + aud := p.audience + + p.mu.Lock() + c, ok := p.cache[aud] + p.mu.Unlock() + + now := p.now() + if ok && now.Before(c.refreshAt) { + return c.token, nil // fresh — fast path, no fetch + } + + tok, err := p.refresh(ctx, aud) + if err != nil { + if ok && now.Before(c.expiry) { + return c.token, nil // serve stale: refresh failed but token still valid + } + return "", err + } + return tok, nil +} + +// refresh fetches a new SVID, collapsing concurrent refreshes so a burst of +// requests triggers a single Workload API call. A fetched token is rejected if +// it is empty or already expired, so the provider fails closed (like the file +// provider) rather than handing the proxy a degenerate "Bearer " credential. +func (p *JWTSVIDProvider) refresh(ctx context.Context, aud string) (tok string, err error) { + p.mu.Lock() + // Re-check under the lock: a concurrent refresh may have just populated a + // fresh entry between our cache read in Token and acquiring this lock. + if c, ok := p.cache[aud]; ok && p.now().Before(c.refreshAt) { + p.mu.Unlock() + return c.token, nil + } + if call, ok := p.inflight[aud]; ok { + p.mu.Unlock() + select { + case <-call.done: + return call.token, call.err + case <-ctx.Done(): + return "", ctx.Err() + } + } + call := &refreshCall{done: make(chan struct{})} + p.inflight[aud] = call + p.mu.Unlock() + + var exp time.Time + // Always release leadership and wake waiters — even if fetch panics — so a + // misbehaving Workload API client can never permanently wedge the provider. + defer func() { + if r := recover(); r != nil { + tok, exp, err = "", time.Time{}, fmt.Errorf("identity/jwtsvid: fetch panicked: %v", r) + } + p.mu.Lock() + delete(p.inflight, aud) + if err == nil { + p.cache[aud] = cachedSVID{token: tok, expiry: exp, refreshAt: p.refreshAt(exp)} + } + p.mu.Unlock() + // Share one result (and one error shape) with any collapsed waiters. + call.token, call.err = tok, err + close(call.done) + }() + + tok, exp, err = p.fetcher.fetch(ctx, aud) + switch { + case err != nil: + err = fmt.Errorf("identity/jwtsvid: fetch: %w", err) + case tok == "": + err = fmt.Errorf("identity/jwtsvid: %w", ErrNoToken) + case !exp.After(p.now()): + err = fmt.Errorf("identity/jwtsvid: fetched SVID already expired at %s", exp.UTC().Format(time.RFC3339)) + } + if err != nil { + tok = "" + } + return tok, err +} + +// refreshAt returns when a token expiring at exp should be proactively +// refreshed: refreshBefore ahead of exp, clamped to at most half the observed +// lifetime so very short TTLs are not refreshed too eagerly. +func (p *JWTSVIDProvider) refreshAt(exp time.Time) time.Time { + lead := p.refreshBefore + if ttl := exp.Sub(p.now()); ttl > 0 && ttl/2 < lead { + lead = ttl / 2 + } + return exp.Add(-lead) +} + +// Close releases the underlying Workload API source. +func (p *JWTSVIDProvider) Close() error { return p.fetcher.close() } diff --git a/internal/identity/jwtsvid_source.go b/internal/identity/jwtsvid_source.go new file mode 100644 index 0000000..018e576 --- /dev/null +++ b/internal/identity/jwtsvid_source.go @@ -0,0 +1,42 @@ +package identity + +import ( + "context" + "fmt" + "time" + + "github.com/spiffe/go-spiffe/v2/svid/jwtsvid" + "github.com/spiffe/go-spiffe/v2/workloadapi" + + "github.com/snangue/robin/internal/config" +) + +// spiffeFetcher is the production jwtFetcher, backed by the SPIFFE Workload API. +type spiffeFetcher struct { + src *workloadapi.JWTSource +} + +func (f *spiffeFetcher) fetch(ctx context.Context, aud string) (string, time.Time, error) { + svid, err := f.src.FetchJWTSVID(ctx, jwtsvid.Params{Audience: aud}) + if err != nil { + return "", time.Time{}, err + } + return svid.Marshal(), svid.Expiry, nil +} + +func (f *spiffeFetcher) close() error { return f.src.Close() } + +// newJWTSVIDSource creates a Workload API JWT source. The socket address is +// passed explicitly when set; otherwise go-spiffe's default applies (which +// honors SPIFFE_ENDPOINT_SOCKET). +func newJWTSVIDSource(ctx context.Context, cfg config.Config) (jwtFetcher, error) { + var opts []workloadapi.JWTSourceOption + if cfg.SPIFFESocket != "" { + opts = append(opts, workloadapi.WithClientOptions(workloadapi.WithAddr(cfg.SPIFFESocket))) + } + src, err := workloadapi.NewJWTSource(ctx, opts...) + if err != nil { + return nil, fmt.Errorf("identity/jwtsvid: create source: %w", err) + } + return &spiffeFetcher{src: src}, nil +} diff --git a/internal/identity/jwtsvid_test.go b/internal/identity/jwtsvid_test.go new file mode 100644 index 0000000..c974a23 --- /dev/null +++ b/internal/identity/jwtsvid_test.go @@ -0,0 +1,252 @@ +package identity + +import ( + "context" + "errors" + "sync" + "testing" + "time" +) + +// fakeFetcher is a programmable jwtFetcher. fetch signals `entered` (if set) +// when it begins, then blocks on `gate` (if set) — letting tests pin a fetch +// in flight deterministically. +type fakeFetcher struct { + mu sync.Mutex + calls int + tok string + exp time.Time + err error + gate chan struct{} + entered chan struct{} +} + +func (f *fakeFetcher) fetch(_ context.Context, _ string) (string, time.Time, error) { + if f.entered != nil { + f.entered <- struct{}{} + } + if f.gate != nil { + <-f.gate + } + f.mu.Lock() + defer f.mu.Unlock() + f.calls++ + return f.tok, f.exp, f.err +} + +func (f *fakeFetcher) close() error { return nil } + +func (f *fakeFetcher) callCount() int { + f.mu.Lock() + defer f.mu.Unlock() + return f.calls +} + +func (f *fakeFetcher) set(tok string, exp time.Time, err error) { + f.mu.Lock() + defer f.mu.Unlock() + f.tok, f.exp, f.err = tok, exp, err +} + +// panicFetcher panics until stop is set — to prove a panicking Workload API +// client cannot permanently wedge the provider. +type panicFetcher struct { + stop bool + tok string + exp time.Time +} + +func (f *panicFetcher) fetch(context.Context, string) (string, time.Time, error) { + if !f.stop { + panic("workload api client blew up") + } + return f.tok, f.exp, nil +} + +func (f *panicFetcher) close() error { return nil } + +// clock is a thread-safe injectable clock. +type clock struct { + mu sync.Mutex + t time.Time +} + +func (c *clock) now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return c.t +} + +func (c *clock) advance(d time.Duration) { + c.mu.Lock() + defer c.mu.Unlock() + c.t = c.t.Add(d) +} + +func newTestProvider(f jwtFetcher, now func() time.Time) *JWTSVIDProvider { + p := newJWTSVIDProviderWith(f, "broker.example", 60*time.Second) + p.now = now + return p +} + +var base = time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC) + +func TestJWTSVIDCacheHit(t *testing.T) { + clk := &clock{t: base} + f := &fakeFetcher{tok: "tok1", exp: base.Add(5 * time.Minute)} + p := newTestProvider(f, clk.now) + + if tok, err := p.Token(context.Background()); err != nil || tok != "tok1" { + t.Fatalf("first Token = %q, %v", tok, err) + } + clk.advance(30 * time.Second) // still before refreshAt (exp-60s) + if tok, err := p.Token(context.Background()); err != nil || tok != "tok1" { + t.Fatalf("second Token = %q, %v", tok, err) + } + if f.callCount() != 1 { + t.Errorf("fetch calls = %d, want 1 (cache hit)", f.callCount()) + } +} + +func TestJWTSVIDProactiveRefresh(t *testing.T) { + clk := &clock{t: base} + f := &fakeFetcher{tok: "tok1", exp: base.Add(5 * time.Minute)} + p := newTestProvider(f, clk.now) + + if _, err := p.Token(context.Background()); err != nil { + t.Fatal(err) + } + f.set("tok2", clk.now().Add(5*time.Minute), nil) + clk.advance(4*time.Minute + time.Second) // past refreshAt (exp-60s), before exp + + tok, err := p.Token(context.Background()) + if err != nil || tok != "tok2" { + t.Fatalf("Token = %q, %v, want tok2", tok, err) + } + if f.callCount() != 2 { + t.Errorf("fetch calls = %d, want 2", f.callCount()) + } +} + +func TestJWTSVIDServeStaleOnError(t *testing.T) { + clk := &clock{t: base} + f := &fakeFetcher{tok: "tok1", exp: base.Add(5 * time.Minute)} + p := newTestProvider(f, clk.now) + + if _, err := p.Token(context.Background()); err != nil { + t.Fatal(err) + } + f.set("tok1", base.Add(5*time.Minute), errors.New("agent down")) + clk.advance(4*time.Minute + 30*time.Second) // refresh window, still before exp + + tok, err := p.Token(context.Background()) + if err != nil { + t.Fatalf("expected stale token served, got error %v", err) + } + if tok != "tok1" { + t.Errorf("Token = %q, want stale tok1", tok) + } +} + +func TestJWTSVIDHardFailWhenColdAndErroring(t *testing.T) { + clk := &clock{t: base} + f := &fakeFetcher{err: errors.New("agent down")} + p := newTestProvider(f, clk.now) + + if _, err := p.Token(context.Background()); err == nil { + t.Error("want error when cache is cold and fetch fails") + } +} + +func TestJWTSVIDHardFailWhenExpiredAndErroring(t *testing.T) { + clk := &clock{t: base} + f := &fakeFetcher{tok: "tok1", exp: base.Add(5 * time.Minute)} + p := newTestProvider(f, clk.now) + if _, err := p.Token(context.Background()); err != nil { + t.Fatal(err) + } + f.set("tok1", base.Add(5*time.Minute), errors.New("agent down")) + clk.advance(6 * time.Minute) // past exp: stale token no longer valid + + if _, err := p.Token(context.Background()); err == nil { + t.Error("want error when cached token has expired and fetch fails") + } +} + +func TestJWTSVIDRejectsEmptyToken(t *testing.T) { + clk := &clock{t: base} + f := &fakeFetcher{tok: "", exp: base.Add(5 * time.Minute)} + p := newTestProvider(f, clk.now) + if _, err := p.Token(context.Background()); err == nil { + t.Error("want error: an empty token must not be returned as a valid bearer") + } +} + +func TestJWTSVIDRejectsExpiredToken(t *testing.T) { + clk := &clock{t: base} + f := &fakeFetcher{tok: "stale", exp: base.Add(-time.Minute)} // already expired + p := newTestProvider(f, clk.now) + if _, err := p.Token(context.Background()); err == nil { + t.Error("want error: an already-expired token must not be returned as valid") + } +} + +func TestJWTSVIDSurvivesFetchPanic(t *testing.T) { + clk := &clock{t: base} + f := &panicFetcher{} + p := newTestProvider(f, clk.now) + + if _, err := p.Token(context.Background()); err == nil { + t.Fatal("want error when fetch panics") + } + // The provider must not be wedged: a later successful fetch works. + f.stop = true + f.tok = "tok1" + f.exp = base.Add(5 * time.Minute) + if tok, err := p.Token(context.Background()); err != nil || tok != "tok1" { + t.Fatalf("provider wedged after panic: tok=%q err=%v", tok, err) + } +} + +func TestJWTSVIDConcurrencyCollapse(t *testing.T) { + clk := &clock{t: base} + gate := make(chan struct{}) + entered := make(chan struct{}, 1) + f := &fakeFetcher{tok: "tok1", exp: base.Add(5 * time.Minute), gate: gate, entered: entered} + p := newTestProvider(f, clk.now) + + // Leader enters fetch and blocks there, holding inflight while the cache is + // still empty — so any concurrent caller is forced onto the collapse path. + leader := make(chan string, 1) + go func() { + tok, _ := p.Token(context.Background()) + leader <- tok + }() + <-entered // leader is now inside fetch: inflight set, cache empty + + const n = 20 + var wg sync.WaitGroup + toks := make([]string, n) + errs := make([]error, n) + for i := 0; i < n; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + toks[i], errs[i] = p.Token(context.Background()) + }(i) + } + close(gate) // release the single in-flight fetch + wg.Wait() + + if got := <-leader; got != "tok1" { + t.Errorf("leader token = %q, want tok1", got) + } + if f.callCount() != 1 { + t.Errorf("fetch calls = %d, want exactly 1 (collapsed)", f.callCount()) + } + for i := 0; i < n; i++ { + if errs[i] != nil || toks[i] != "tok1" { + t.Errorf("waiter %d: %q, %v", i, toks[i], errs[i]) + } + } +} diff --git a/internal/identity/provider.go b/internal/identity/provider.go new file mode 100644 index 0000000..439654c --- /dev/null +++ b/internal/identity/provider.go @@ -0,0 +1,25 @@ +// Package identity resolves the workload's native identity as a bearer token. +// Each provider yields a token for the configured audience; the proxy applies +// it uniformly without knowing which provider produced it. +package identity + +import ( + "context" + "errors" +) + +// Provider resolves the workload's native identity as a bearer token. +type Provider interface { + // Token returns a bearer token valid for the configured audience. + // Implementations must be safe for concurrent use. + Token(ctx context.Context) (string, error) + // Close releases any background resources held by the provider. + Close() error +} + +// Sentinel errors. ErrAudienceMismatch is a fail-closed configuration error, +// not a transient condition. +var ( + ErrNoToken = errors.New("identity: empty token") + ErrAudienceMismatch = errors.New("identity: audience mismatch") +) diff --git a/internal/obs/log.go b/internal/obs/log.go new file mode 100644 index 0000000..2b26aba --- /dev/null +++ b/internal/obs/log.go @@ -0,0 +1,46 @@ +// Package obs provides observability primitives: structured logging and, in +// v0.2, metrics. The token value is deliberately never a logged field. +package obs + +import ( + "log/slog" + "os" + "strings" +) + +// Structured log field keys, centralized so they stay consistent — and so the +// token value is conspicuously absent from the set. +const ( + FieldProvider = "provider" + FieldAudience = "audience" + FieldDecision = "decision" + FieldUpstreamStatus = "upstream_status" + FieldError = "err" +) + +// Values for the FieldDecision log field. +const ( + DecisionInjected = "injected" + DecisionFailed = "failed" + DecisionUpstreamError = "upstream_error" +) + +// NewLogger returns a JSON slog.Logger at the given level, writing to stdout. +func NewLogger(level slog.Level) *slog.Logger { + h := slog.NewJSONHandler(os.Stdout, &slog.HandlerOptions{Level: level}) + return slog.New(h) +} + +// ParseLevel maps a level string to an slog.Level, defaulting to Info. +func ParseLevel(s string) slog.Level { + switch strings.ToLower(strings.TrimSpace(s)) { + case "debug": + return slog.LevelDebug + case "warn", "warning": + return slog.LevelWarn + case "error": + return slog.LevelError + default: + return slog.LevelInfo + } +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go new file mode 100644 index 0000000..2a99786 --- /dev/null +++ b/internal/proxy/proxy.go @@ -0,0 +1,98 @@ +// Package proxy is Robin's reverse-proxy core: it resolves the workload's +// identity and injects it as a bearer token on every request forwarded to the +// broker. It is body-agnostic and streaming — the request body is never read. +package proxy + +import ( + "context" + "errors" + "fmt" + "log/slog" + "net/http" + "net/http/httputil" + "net/url" + + "github.com/snangue/robin/internal/config" + "github.com/snangue/robin/internal/identity" + "github.com/snangue/robin/internal/obs" +) + +type ctxKey int + +const tokenCtxKey ctxKey = 0 + +// Handler resolves identity and proxies requests to the broker, injecting the +// token as a bearer header. +type Handler struct { + provider identity.Provider + rp *httputil.ReverseProxy + log *slog.Logger + source string + audience string +} + +// New builds a proxy Handler targeting cfg.UpstreamURL. +func New(cfg config.Config, p identity.Provider, log *slog.Logger) (*Handler, error) { + target, err := url.Parse(cfg.UpstreamURL) + if err != nil { + return nil, fmt.Errorf("proxy: parse upstream URL: %w", err) + } + transport, err := buildTransport(cfg) + if err != nil { + return nil, err + } + + h := &Handler{ + provider: p, + log: log, + source: cfg.TokenSource, + audience: cfg.Audience, + } + h.rp = &httputil.ReverseProxy{ + Transport: transport, + Rewrite: func(pr *httputil.ProxyRequest) { + pr.SetURL(target) + // Inject identity, overwriting any inbound placeholder credential. + tok, _ := pr.In.Context().Value(tokenCtxKey).(string) + pr.Out.Header.Set("Authorization", "Bearer "+tok) + }, + ModifyResponse: func(resp *http.Response) error { + // One access-log line per request, with the broker's status. + h.log.Info("request injected", + slog.String(obs.FieldProvider, h.source), + slog.String(obs.FieldAudience, h.audience), + slog.String(obs.FieldDecision, obs.DecisionInjected), + slog.Int(obs.FieldUpstreamStatus, resp.StatusCode)) + return nil + }, + ErrorHandler: func(w http.ResponseWriter, _ *http.Request, err error) { + // Single attempt — surface as 502, never retry into a thundering herd. + h.log.Error("broker unreachable", + slog.String(obs.FieldProvider, h.source), + slog.String(obs.FieldDecision, obs.DecisionUpstreamError), + slog.String(obs.FieldError, err.Error())) + w.WriteHeader(http.StatusBadGateway) + }, + } + return h, nil +} + +// ServeHTTP resolves the workload's identity and forwards the request. If the +// identity cannot be resolved it returns 503 and never contacts the broker. +func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + tok, err := h.provider.Token(r.Context()) + if err == nil && tok == "" { + // Fail closed: never forward an empty/placeholder credential. + err = errors.New("provider returned an empty token") + } + if err != nil { + h.log.Warn("identity unavailable", + slog.String(obs.FieldProvider, h.source), + slog.String(obs.FieldDecision, obs.DecisionFailed), + slog.String(obs.FieldError, err.Error())) + http.Error(w, "identity unavailable", http.StatusServiceUnavailable) + return + } + ctx := context.WithValue(r.Context(), tokenCtxKey, tok) + h.rp.ServeHTTP(w, r.WithContext(ctx)) +} diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go new file mode 100644 index 0000000..ab92ee3 --- /dev/null +++ b/internal/proxy/proxy_test.go @@ -0,0 +1,131 @@ +package proxy + +import ( + "context" + "errors" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/snangue/robin/internal/config" +) + +type stubProvider struct { + tok string + err error +} + +func (s stubProvider) Token(context.Context) (string, error) { return s.tok, s.err } +func (s stubProvider) Close() error { return nil } + +func discardLogger() *slog.Logger { return slog.New(slog.NewTextHandler(io.Discard, nil)) } + +func newHandler(t *testing.T, upstream string, p stubProvider) *Handler { + t.Helper() + h, err := New(config.Config{UpstreamURL: upstream, TokenSource: "file"}, p, discardLogger()) + if err != nil { + t.Fatalf("New: %v", err) + } + return h +} + +func TestInjectOverwritesPlaceholderAndForwardsBody(t *testing.T) { + var gotAuth, gotBody, gotPath string + broker := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + gotAuth = r.Header.Get("Authorization") + gotPath = r.URL.Path + b, _ := io.ReadAll(r.Body) + gotBody = string(b) + w.WriteHeader(http.StatusNoContent) + })) + defer broker.Close() + + srv := httptest.NewServer(newHandler(t, broker.URL, stubProvider{tok: "realtok"})) + defer srv.Close() + + req, _ := http.NewRequest(http.MethodPost, srv.URL+"/v1/chat", strings.NewReader("hello-body")) + req.Header.Set("Authorization", "Bearer placeholder") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + + if gotAuth != "Bearer realtok" { + t.Errorf("Authorization = %q, want Bearer realtok (placeholder overwritten)", gotAuth) + } + if gotBody != "hello-body" { + t.Errorf("forwarded body = %q, want hello-body", gotBody) + } + if gotPath != "/v1/chat" { + t.Errorf("forwarded path = %q, want /v1/chat", gotPath) + } +} + +func TestProviderErrorReturns503(t *testing.T) { + called := false + broker := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + called = true + })) + defer broker.Close() + + srv := httptest.NewServer(newHandler(t, broker.URL, stubProvider{err: errors.New("boom")})) + defer srv.Close() + + resp, err := http.Get(srv.URL) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusServiceUnavailable { + t.Errorf("status = %d, want 503", resp.StatusCode) + } + if called { + t.Error("broker must not be contacted when identity is unavailable") + } +} + +func TestEmptyTokenReturns503(t *testing.T) { + called := false + broker := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + called = true + })) + defer broker.Close() + + // Provider returns ("", nil) — must fail closed, never forward "Bearer ". + srv := httptest.NewServer(newHandler(t, broker.URL, stubProvider{tok: ""})) + defer srv.Close() + + resp, err := http.Get(srv.URL) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusServiceUnavailable { + t.Errorf("status = %d, want 503 for empty token", resp.StatusCode) + } + if called { + t.Error("broker must not be contacted when the token is empty") + } +} + +func TestBrokerDownReturns502(t *testing.T) { + dead := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + deadURL := dead.URL + dead.Close() // nothing is listening now + + srv := httptest.NewServer(newHandler(t, deadURL, stubProvider{tok: "tok"})) + defer srv.Close() + + resp, err := http.Get(srv.URL) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusBadGateway { + t.Errorf("status = %d, want 502", resp.StatusCode) + } +} diff --git a/internal/proxy/transport.go b/internal/proxy/transport.go new file mode 100644 index 0000000..dbfc083 --- /dev/null +++ b/internal/proxy/transport.go @@ -0,0 +1,37 @@ +package proxy + +import ( + "crypto/tls" + "crypto/x509" + "fmt" + "net/http" + "os" + + "github.com/snangue/robin/internal/config" +) + +// buildTransport clones the default transport (preserving connection-pool and +// timeout defaults) and, when a CA file is configured, pins the broker's trust +// roots for TLS verification. +func buildTransport(cfg config.Config) (*http.Transport, error) { + base, ok := http.DefaultTransport.(*http.Transport) + if !ok { + return nil, fmt.Errorf("proxy: unexpected default transport type %T", http.DefaultTransport) + } + t := base.Clone() + // Pin a TLS 1.2 floor on every path, not only when a CA file is configured. + tlsCfg := &tls.Config{MinVersion: tls.VersionTLS12} + if cfg.UpstreamCAFile != "" { + pem, err := os.ReadFile(cfg.UpstreamCAFile) + if err != nil { + return nil, fmt.Errorf("proxy: read upstream CA %s: %w", cfg.UpstreamCAFile, err) + } + pool := x509.NewCertPool() + if !pool.AppendCertsFromPEM(pem) { + return nil, fmt.Errorf("proxy: no certificates found in %s", cfg.UpstreamCAFile) + } + tlsCfg.RootCAs = pool + } + t.TLSClientConfig = tlsCfg + return t, nil +} diff --git a/internal/proxy/transport_test.go b/internal/proxy/transport_test.go new file mode 100644 index 0000000..f1f23ee --- /dev/null +++ b/internal/proxy/transport_test.go @@ -0,0 +1,60 @@ +package proxy + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "os" + "path/filepath" + "testing" + + "github.com/snangue/robin/internal/config" +) + +func TestBuildTransportNoCA(t *testing.T) { + tr, err := buildTransport(config.Config{}) + if err != nil { + t.Fatal(err) + } + if tr.TLSClientConfig != nil && tr.TLSClientConfig.RootCAs != nil { + t.Error("expected no custom RootCAs when CA file is unset") + } +} + +func TestBuildTransportBadCA(t *testing.T) { + path := filepath.Join(t.TempDir(), "ca.pem") + if err := os.WriteFile(path, []byte("not a pem"), 0o600); err != nil { + t.Fatal(err) + } + if _, err := buildTransport(config.Config{UpstreamCAFile: path}); err == nil { + t.Error("expected error for a file with no valid certificates") + } +} + +func TestBuildTransportGoodCA(t *testing.T) { + priv, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatal(err) + } + tmpl := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "test-ca"}} + der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &priv.PublicKey, priv) + if err != nil { + t.Fatal(err) + } + path := filepath.Join(t.TempDir(), "ca.pem") + if err := os.WriteFile(path, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o600); err != nil { + t.Fatal(err) + } + + tr, err := buildTransport(config.Config{UpstreamCAFile: path}) + if err != nil { + t.Fatal(err) + } + if tr.TLSClientConfig == nil || tr.TLSClientConfig.RootCAs == nil { + t.Error("expected RootCAs to be populated from the CA file") + } +} diff --git a/internal/server/admin.go b/internal/server/admin.go new file mode 100644 index 0000000..7accf9d --- /dev/null +++ b/internal/server/admin.go @@ -0,0 +1,33 @@ +package server + +import ( + "io" + "net/http" + + "github.com/snangue/robin/internal/identity" +) + +// AdminMux builds the admin handler served on a listener separate from the +// proxy: liveness, readiness, and (in v0.2) metrics. These must not live on the +// proxy listener, which forwards every path to the broker. +func AdminMux(p identity.Provider) http.Handler { + mux := http.NewServeMux() + + // Liveness: the process is up. + mux.HandleFunc("/healthz", func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = io.WriteString(w, "ok") + }) + + // Readiness: an identity can actually be resolved. + mux.HandleFunc("/readyz", func(w http.ResponseWriter, r *http.Request) { + if _, err := p.Token(r.Context()); err != nil { + http.Error(w, "not ready", http.StatusServiceUnavailable) + return + } + w.WriteHeader(http.StatusOK) + _, _ = io.WriteString(w, "ready") + }) + + return mux +} diff --git a/internal/server/listen.go b/internal/server/listen.go new file mode 100644 index 0000000..4d8f826 --- /dev/null +++ b/internal/server/listen.go @@ -0,0 +1,35 @@ +package server + +import ( + "fmt" + "net" + "os" + + "github.com/snangue/robin/internal/config" +) + +// proxyListener builds the proxy listener: a Unix domain socket when configured +// (peer-cred hardening lands in v0.2), otherwise a TCP listener. +func proxyListener(cfg config.Config) (net.Listener, error) { + if cfg.ListenUDS != "" { + // Remove a stale socket left behind by a previous run. + if err := os.Remove(cfg.ListenUDS); err != nil && !os.IsNotExist(err) { + return nil, fmt.Errorf("server: remove stale socket %s: %w", cfg.ListenUDS, err) + } + ln, err := net.Listen("unix", cfg.ListenUDS) + if err != nil { + return nil, fmt.Errorf("server: listen unix %s: %w", cfg.ListenUDS, err) + } + // 0o600 until v0.2 SO_PEERCRED enforcement lands: owner-only, not group-wide. + if err := os.Chmod(cfg.ListenUDS, 0o600); err != nil { + _ = ln.Close() + return nil, fmt.Errorf("server: chmod %s: %w", cfg.ListenUDS, err) + } + return ln, nil + } + ln, err := net.Listen("tcp", cfg.ListenAddr) + if err != nil { + return nil, fmt.Errorf("server: listen tcp %s: %w", cfg.ListenAddr, err) + } + return ln, nil +} diff --git a/internal/server/server.go b/internal/server/server.go new file mode 100644 index 0000000..830d1a5 --- /dev/null +++ b/internal/server/server.go @@ -0,0 +1,89 @@ +// Package server runs Robin's two listeners — the proxy plane and a separate +// admin plane (health/readiness/metrics) — with graceful, drain-on-SIGTERM +// shutdown so in-flight egress is not severed during pod termination. +package server + +import ( + "context" + "errors" + "fmt" + "log/slog" + "net" + "net/http" + "time" + + "github.com/snangue/robin/internal/config" +) + +// shutdownTimeout bounds the in-flight drain. Deployments must set their +// termination grace period strictly larger (see deploy/ in v0.2). +const shutdownTimeout = 25 * time.Second + +// Server owns the proxy and admin listeners. +type Server struct { + proxy *http.Server + admin *http.Server + proxyLn net.Listener + adminLn net.Listener + log *slog.Logger +} + +// New binds both listeners (proxy: TCP or UDS; admin: TCP) up front, so their +// resolved addresses are available before Run and addressable in tests. +func New(cfg config.Config, proxyHandler, adminHandler http.Handler, log *slog.Logger) (*Server, error) { + proxyLn, err := proxyListener(cfg) + if err != nil { + return nil, err + } + adminLn, err := net.Listen("tcp", cfg.AdminAddr) + if err != nil { + _ = proxyLn.Close() + return nil, fmt.Errorf("server: listen admin %s: %w", cfg.AdminAddr, err) + } + return &Server{ + proxy: &http.Server{Handler: proxyHandler}, + admin: &http.Server{Handler: adminHandler}, + proxyLn: proxyLn, + adminLn: adminLn, + log: log, + }, nil +} + +// ProxyAddr returns the resolved proxy listener address. +func (s *Server) ProxyAddr() net.Addr { return s.proxyLn.Addr() } + +// AdminAddr returns the resolved admin listener address. +func (s *Server) AdminAddr() net.Addr { return s.adminLn.Addr() } + +// Run serves both planes until ctx is canceled (SIGTERM), then drains +// in-flight requests within shutdownTimeout. +func (s *Server) Run(ctx context.Context) error { + errc := make(chan error, 2) + serve := func(srv *http.Server, ln net.Listener) { + if err := srv.Serve(ln); err != nil && !errors.Is(err, http.ErrServerClosed) { + errc <- err + } + } + go serve(s.proxy, s.proxyLn) + go serve(s.admin, s.adminLn) + + s.log.Info("robin started", + slog.String("proxy", s.proxyLn.Addr().String()), + slog.String("admin", s.adminLn.Addr().String())) + + shutdown := func() error { + shutCtx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) + defer cancel() + // Drain the proxy first (finish in-flight egress), then the admin plane. + return errors.Join(s.proxy.Shutdown(shutCtx), s.admin.Shutdown(shutCtx)) + } + + select { + case err := <-errc: + // One plane failed to serve: shut the other down too, never leave it running. + return errors.Join(err, shutdown()) + case <-ctx.Done(): + s.log.Info("shutting down, draining in-flight egress") + return shutdown() + } +} diff --git a/internal/server/server_test.go b/internal/server/server_test.go new file mode 100644 index 0000000..bc1f4ec --- /dev/null +++ b/internal/server/server_test.go @@ -0,0 +1,146 @@ +package server + +import ( + "context" + "errors" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "testing" + + "github.com/snangue/robin/internal/config" + "github.com/snangue/robin/internal/proxy" +) + +type stubProvider struct { + tok string + err error +} + +func (s stubProvider) Token(context.Context) (string, error) { return s.tok, s.err } +func (s stubProvider) Close() error { return nil } + +func discardLogger() *slog.Logger { return slog.New(slog.NewTextHandler(io.Discard, nil)) } + +func startServer(t *testing.T, broker string, p stubProvider) *Server { + t.Helper() + cfg := config.Config{UpstreamURL: broker, TokenSource: "file", ListenAddr: "127.0.0.1:0", AdminAddr: "127.0.0.1:0"} + h, err := proxy.New(cfg, p, discardLogger()) + if err != nil { + t.Fatal(err) + } + srv, err := New(cfg, h, AdminMux(p), discardLogger()) + if err != nil { + t.Fatal(err) + } + return srv +} + +func TestHealthReadyAndProxy(t *testing.T) { + broker := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-Auth", r.Header.Get("Authorization")) + w.WriteHeader(http.StatusNoContent) + })) + defer broker.Close() + + srv := startServer(t, broker.URL, stubProvider{tok: "tok"}) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- srv.Run(ctx) }() + + get := func(url string) *http.Response { + t.Helper() + resp, err := http.Get(url) + if err != nil { + t.Fatal(err) + } + return resp + } + + if r := get("http://" + srv.AdminAddr().String() + "/healthz"); r.StatusCode != http.StatusOK { + t.Errorf("/healthz = %d, want 200", r.StatusCode) + r.Body.Close() + } else { + r.Body.Close() + } + + if r := get("http://" + srv.AdminAddr().String() + "/readyz"); r.StatusCode != http.StatusOK { + t.Errorf("/readyz = %d, want 200", r.StatusCode) + r.Body.Close() + } else { + r.Body.Close() + } + + r := get("http://" + srv.ProxyAddr().String() + "/v1/x") + if got := r.Header.Get("X-Auth"); got != "Bearer tok" { + t.Errorf("proxied Authorization = %q, want Bearer tok", got) + } + r.Body.Close() + + cancel() + if err := <-done; err != nil { + t.Errorf("Run returned %v", err) + } +} + +func TestReadyzReflectsProviderFailure(t *testing.T) { + broker := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})) + defer broker.Close() + + srv := startServer(t, broker.URL, stubProvider{err: errors.New("workload api down")}) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- srv.Run(ctx) }() + + resp, err := http.Get("http://" + srv.AdminAddr().String() + "/readyz") + if err != nil { + t.Fatal(err) + } + if resp.StatusCode != http.StatusServiceUnavailable { + t.Errorf("/readyz = %d, want 503 when provider cannot resolve", resp.StatusCode) + } + resp.Body.Close() + + cancel() + <-done +} + +func TestGracefulShutdownDrainsInflight(t *testing.T) { + received := make(chan struct{}) + release := make(chan struct{}) + broker := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + close(received) + <-release + w.WriteHeader(http.StatusOK) + })) + defer broker.Close() + + srv := startServer(t, broker.URL, stubProvider{tok: "tok"}) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan error, 1) + go func() { done <- srv.Run(ctx) }() + + addr := srv.ProxyAddr().String() + codes := make(chan int, 1) + go func() { + resp, err := http.Get("http://" + addr + "/slow") + if err != nil { + codes <- -1 + return + } + codes <- resp.StatusCode + resp.Body.Close() + }() + + <-received // request is in-flight at the broker + cancel() // SIGTERM-equivalent: begin graceful shutdown + close(release) // let the broker finish responding + + if code := <-codes; code != http.StatusOK { + t.Errorf("in-flight request got %d, want 200 (should drain, not be severed)", code) + } + if err := <-done; err != nil { + t.Errorf("Run returned %v", err) + } +} diff --git a/internal/version/version.go b/internal/version/version.go new file mode 100644 index 0000000..9124e36 --- /dev/null +++ b/internal/version/version.go @@ -0,0 +1,15 @@ +// Package version holds build metadata injected at link time via -ldflags. +package version + +// Build metadata, set via -X linker flags (see Makefile). Defaults apply to +// `go run` and `go test` builds. +var ( + Version = "dev" + Commit = "none" + Date = "unknown" +) + +// String returns a human-readable version line. +func String() string { + return Version + " (commit " + Commit + ", built " + Date + ")" +}