From fcd75f64b92927abd1a32c3cb7080d8de5bf5e37 Mon Sep 17 00:00:00 2001
From: Kieran Klukas
Date: Mon, 27 Jul 2026 03:15:09 -0400
Subject: [PATCH] feat: proper indiko device oauth
---
.gitignore | 1 +
README.md | 58 ++--
cmd/lard-client/main.go | 34 ++-
cmd/lard/main.go | 98 +------
internal/client/config.go | 8 +-
internal/client/device.go | 189 +++++++-----
internal/client/oauth.go | 417 +++------------------------
internal/client/oauth_test.go | 151 ----------
internal/collector/collector.go | 283 ++----------------
internal/collector/collector_test.go | 156 +---------
internal/collector/device.go | 217 --------------
internal/collector/device_test.go | 313 --------------------
internal/collector/devicehandler.go | 282 ------------------
internal/setup/setup.go | 95 +++---
14 files changed, 299 insertions(+), 2003 deletions(-)
delete mode 100644 internal/client/oauth_test.go
delete mode 100644 internal/collector/device.go
delete mode 100644 internal/collector/device_test.go
delete mode 100644 internal/collector/devicehandler.go
diff --git a/.gitignore b/.gitignore
index ad8f118..a70c657 100644
--- a/.gitignore
+++ b/.gitignore
@@ -5,4 +5,5 @@ crush.json
*.db-wal
/lard
/lard-client
+/bin/
/memory
diff --git a/README.md b/README.md
index 89619a1..b763868 100644
--- a/README.md
+++ b/README.md
@@ -35,7 +35,8 @@ consolidates itself once uploads go quiet.
## client
```
-lard-client login [--url URL] [--token TOKEN] [--root DIR...]
+lard-client login [--url URL] [--token TOKEN] [--root DIR...] [-f]
+lard-client logout # revoke + forget credentials
lard-client status # server, auth, agent at a glance
lard-client backfill [--root DIR...] # every session ever, idempotent
lard-client sync [--workspace DIR...] # new sessions only
@@ -100,8 +101,7 @@ paths are `profile`, `areas/`, `topics/`, `people/`.
| `LARD_OAUTH_USERS` | | comma list of indiko `me` urls allowed to call lard |
| `LARD_OAUTH_SCOPES` | | comma list of scopes every token must carry |
| `LARD_COLLECTOR_CLIENT_ID` | | oauth client id collectors should use |
-| `LARD_COLLECTOR_CLIENT_SECRET` | | its secret; enables server-side code exchange |
-| `LARD_COLLECTOR_PORTS` | `40714-40718` | permitted localhost callback ports |
+| `LARD_COLLECTOR_SCOPES` | `profile` | scopes the collector should request |
| `LARD_CONSOLIDATE_AFTER` | `5m` | quiet period before a pass; `off` to disable |
| `LARD_CONSOLIDATE_MAX_WAIT` | `30m` | cap on that wait during constant uploads |
| `HYPER_API_KEY` | | hyper API key for consolidation |
@@ -137,49 +137,37 @@ memory. lard warns at boot if you skip it.
a collector cannot invent its own client id: the auth server decides which
clients exist, and lard decides which it trusts, so a guessed id gets rejected.
-set `LARD_COLLECTOR_CLIENT_ID` and lard publishes it at `/auth/collector`, and
-`lard-client login` adopts it. that id is then trusted automatically, so it does
-not also need to be in `LARD_OAUTH_CLIENT_IDS`.
+set `LARD_COLLECTOR_CLIENT_ID` to a client id registered with your auth server
+and lard publishes it at `/auth/collector`; `lard-client login` adopts it. that
+id is then trusted automatically, so it does not also need to be in
+`LARD_OAUTH_CLIENT_IDS`.
-if you also set `LARD_COLLECTOR_CLIENT_SECRET`, lard performs the code exchange
-itself. the collector keeps its pkce verifier, the server keeps the secret, and
-neither can complete the exchange alone. that is what lets a pre-registered
-confidential client work from a laptop CLI without shipping the secret around.
+### login (device grant)
-### brokered login (device flow)
-
-with a collector client id and `LARD_PUBLIC_URL` set, lard brokers the whole
-authorization, following the shape of the oauth device grant ([rfc 8628]):
+the only login flow is the oauth device authorization grant ([rfc 8628]), run
+against the authorization server directly — lard is not involved beyond handing
+the collector its client id:
```
-POST /auth/collector/device client starts, gets a code + url
-GET /auth/collector/device/verify user opens this, gets sent to the provider
-GET ...device/callback provider redirects here; lard exchanges
-POST ...device/token client polls until the token appears
+POST {as}/auth/device client gets a device code + user code + url
+GET {as}/device?code=XXXX-XXXX user approves, from any browser anywhere
+POST {as}/auth/token client polls until the token appears
```
-the redirect lands on lard, not on the collector, which is the whole point: a
-server url is reachable from whatever browser the user actually has. the
-collector needs no listener, no port forward, and no browser of its own.
-
-**register the callback with your provider.** lard logs the exact url at boot:
-
-```
-INFO register this redirect URI with your auth provider
- redirect_uri=https://lard.your.domain/auth/collector/device/callback
-```
+no listener, no port forward, and no browser on the collector's machine, so
+ssh, containers, and headless boxes all work the same way. no client secret
+either: the device code itself is the proof of possession, so the collector is
+an ordinary public client and there is nothing sensitive to ship.
-that url is built from `LARD_PUBLIC_URL`, so set it to the server's external
-address before registering.
+requirements on the provider: it must serve rfc 8414 metadata advertising
+`device_authorization_endpoint` and support the device grant (indiko does).
+login fails with a clear message otherwise.
-pending sessions live in memory for 10 minutes, device codes are single use, and
-polling faster than once a second gets `slow_down`.
+with no collector registration configured (`LARD_COLLECTOR_CLIENT_ID` unset),
+`lard-client login` fails and tells you to set it.
[rfc 8628]: https://datatracker.ietf.org/doc/html/rfc8628
-with no collector registration configured, `lard-client` falls back to a
-`http://localhost:/` client id and tells you to allowlist it.
-
### indiko notes
indiko has no dynamic client registration, so mcp clients that insist on
diff --git a/cmd/lard-client/main.go b/cmd/lard-client/main.go
index 6cbe056..7f723e8 100644
--- a/cmd/lard-client/main.go
+++ b/cmd/lard-client/main.go
@@ -46,6 +46,7 @@ Start with 'lard-client login', then 'lard-client backfill'.`,
}
root.AddCommand(
loginCmd(),
+ logoutCmd(),
backfillCmd(),
syncCmd(),
daemonCmd(),
@@ -73,7 +74,7 @@ authorize. On a headless machine pass --url and --token instead.`,
if err != nil {
return err
}
- printConnected(cmd.Context(), cfg, opts.CallbackPort)
+ printConnected(cfg)
return nil
},
}
@@ -81,26 +82,27 @@ authorize. On a headless machine pass --url and --token instead.`,
f.StringVar(&opts.URL, "url", "", "server base URL (asked for if omitted)")
f.StringVar(&opts.Token, "token", "", "shared secret instead of the browser login")
f.StringSliceVar(&opts.Roots, "root", nil, "directory to scan for Crush sessions (repeatable)")
- f.IntVar(&opts.CallbackPort, "port", client.CallbackPort, "localhost port for the OAuth callback")
f.BoolVar(&opts.NoBrowser, "no-browser", false, "print the authorization URL instead of opening it")
+ f.BoolVarP(&opts.Force, "force", "f", false, "re-authenticate, revoking the old grant first")
return cmd
}
-func printConnected(ctx context.Context, cfg *client.Config, port int) {
- fmt.Printf("Connected to %s via %s.\n", cfg.URL, cfg.AuthMode())
- // Only mention the client id when we had to invent one, since that is the
- // case where the operator must add it to the server's allowlist. A
- // server-published id is trusted by definition.
- if cfg.AuthMode() == "oauth" {
- if _, err := client.FetchRegistration(ctx, cfg.URL); err != nil {
- if port <= 0 {
- port = client.CallbackPort
- }
- fmt.Printf("OAuth client id: %s\n", client.ClientID(port))
- fmt.Println(" This server publishes no collector registration, so add that id to")
- fmt.Println(" its LARD_OAUTH_CLIENT_IDS, or set LARD_COLLECTOR_CLIENT_ID there instead.")
- }
+func logoutCmd() *cobra.Command {
+ return &cobra.Command{
+ Use: "logout",
+ Short: "Forget this machine's credentials",
+ Long: `Revoke the refresh token at the authorization server and remove the
+saved credentials. The server URL and roots are kept, so 'lard-client login'
+only has to re-authenticate.`,
+ Args: cobra.NoArgs,
+ RunE: func(cmd *cobra.Command, _ []string) error {
+ return setup.Logout(cmd.Context())
+ },
}
+}
+
+func printConnected(cfg *client.Config) {
+ fmt.Printf("Connected to %s via %s.\n", cfg.URL, cfg.AuthMode())
fmt.Printf("Saved %s\n\nNext: lard-client backfill\n", client.DefaultConfigPath())
}
diff --git a/cmd/lard/main.go b/cmd/lard/main.go
index 4a1800c..11dc8c3 100644
--- a/cmd/lard/main.go
+++ b/cmd/lard/main.go
@@ -3,7 +3,6 @@ package main
import (
"context"
- "encoding/json"
"errors"
"fmt"
"log/slog"
@@ -11,7 +10,6 @@ import (
"os"
"os/signal"
"path/filepath"
- "strconv"
"strings"
"syscall"
"time"
@@ -80,40 +78,16 @@ func run() error {
CollectorClientID: os.Getenv("LARD_COLLECTOR_CLIENT_ID"),
}
- // The collector registration: what identity edge collectors adopt, and
- // whether this server exchanges their codes for them.
+ // The collector registration: which OAuth client edge collectors adopt.
+ // Login itself is the device grant against the authorization server, so
+ // this server only publishes the identity.
collectorCfg := collector.Config{
- ClientID: os.Getenv("LARD_COLLECTOR_CLIENT_ID"),
- ClientSecret: os.Getenv("LARD_COLLECTOR_CLIENT_SECRET"),
- Ports: envPorts("LARD_COLLECTOR_PORTS"),
- Scopes: envList("LARD_COLLECTOR_SCOPES"),
+ ClientID: os.Getenv("LARD_COLLECTOR_CLIENT_ID"),
+ Scopes: envList("LARD_COLLECTOR_SCOPES"),
}
- // The brokered device login needs both endpoints up front: this server, not
- // the collector, is the one that talks to the authorization server.
- var deviceCfg collector.DeviceConfig
+ collectorH := collector.New(collectorCfg)
if collectorCfg.Configured() {
- if meta, err := discoverAuthMetadata(ctx, cfg.IndikoURL); err != nil {
- slog.Warn("collector: cannot discover auth endpoints; brokered login disabled", "error", err)
- collectorCfg.ClientSecret = ""
- } else {
- collectorCfg.TokenEndpoint = meta.TokenEndpoint
- deviceCfg = collector.DeviceConfig{
- PublicURL: cfg.PublicURL,
- AuthorizationEndpoint: meta.AuthorizationEndpoint,
- }
- }
- }
- collectorH := collector.New(collectorCfg, deviceCfg, auth.PathCollector)
- if collectorCfg.Configured() {
- slog.Info("collector registration published",
- "client_id", collectorCfg.ClientID,
- "confidential", collectorCfg.Confidential(),
- "device_flow", collectorH.DeviceAvailable())
- if uri := collectorH.CallbackURI(); uri != "" {
- slog.Info("register this redirect URI with your auth provider", "redirect_uri", uri)
- } else if cfg.PublicURL == "" {
- slog.Warn("brokered device login needs LARD_PUBLIC_URL set to this server's external URL")
- }
+ slog.Info("collector registration published", "client_id", collectorCfg.ClientID)
}
for _, warn := range cfg.Validate() {
slog.Warn("auth: " + warn)
@@ -128,13 +102,6 @@ func run() error {
mux.Handle(auth.PathProtectedResource+"/", auth.ProtectedResourceMetadata(cfg))
mux.Handle(auth.PathAuthServer, auth.AuthServerMetadata(cfg))
mux.HandleFunc("GET "+auth.PathCollector, collectorH.Register)
- mux.HandleFunc("POST "+auth.PathCollector+"/exchange", collectorH.Exchange)
- mux.HandleFunc("POST "+auth.PathCollector+"/refresh", collectorH.Refresh)
- // Brokered device login: the collector polls, the user's browser visits.
- mux.HandleFunc("POST "+auth.PathCollector+collector.PathDevice, collectorH.StartDevice)
- mux.HandleFunc("POST "+auth.PathCollector+collector.PathDeviceToken, collectorH.PollDevice)
- mux.HandleFunc("GET "+auth.PathCollector+collector.PathVerify, collectorH.Verify)
- mux.HandleFunc("GET "+auth.PathCollector+collector.PathCallback, collectorH.Callback)
mux.Handle("/", api.Handler())
addr := envOr("LARD_ADDR", ":7477")
@@ -211,54 +178,3 @@ func defaultMemDir() string {
}
return "memory"
}
-
-// envPorts reads a comma-separated list of port numbers.
-func envPorts(key string) []int {
- var out []int
- for _, s := range envList(key) {
- if n, err := strconv.Atoi(s); err == nil && n > 0 && n < 65536 {
- out = append(out, n)
- } else {
- slog.Warn("ignoring invalid port", "key", key, "value", s)
- }
- }
- return out
-}
-
-// authMetadata is the subset of RFC 8414 metadata the collector flows need.
-type authMetadata struct {
- AuthorizationEndpoint string `json:"authorization_endpoint"`
- TokenEndpoint string `json:"token_endpoint"`
-}
-
-// discoverAuthMetadata reads the authorization server's metadata, so no
-// endpoint path is hardcoded and swapping providers needs no code change.
-func discoverAuthMetadata(ctx context.Context, authServerURL string) (*authMetadata, error) {
- if authServerURL == "" {
- return nil, errors.New("no authorization server configured")
- }
- ctx, cancel := context.WithTimeout(ctx, 10*time.Second)
- defer cancel()
- url := strings.TrimRight(authServerURL, "/") + auth.PathAuthServer
- req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
- if err != nil {
- return nil, err
- }
- req.Header.Set("accept", "application/json")
- resp, err := http.DefaultClient.Do(req)
- if err != nil {
- return nil, err
- }
- defer resp.Body.Close()
- if resp.StatusCode != http.StatusOK {
- return nil, fmt.Errorf("%s returned %d", url, resp.StatusCode)
- }
- var meta authMetadata
- if err := json.NewDecoder(resp.Body).Decode(&meta); err != nil {
- return nil, err
- }
- if meta.TokenEndpoint == "" || meta.AuthorizationEndpoint == "" {
- return nil, errors.New("metadata is missing endpoints")
- }
- return &meta, nil
-}
diff --git a/internal/client/config.go b/internal/client/config.go
index b9339bc..f1a1931 100644
--- a/internal/client/config.go
+++ b/internal/client/config.go
@@ -36,9 +36,9 @@ type OAuthToken struct {
AccessToken string `json:"accessToken"`
RefreshToken string `json:"refreshToken,omitempty"`
Expiry time.Time `json:"expiry,omitempty"`
- // CallbackPort pins the port the login used, since the OAuth client id is
- // derived from it and a refresh must present the same id.
- CallbackPort int `json:"callbackPort,omitempty"`
+ // ClientID is the public OAuth client this token was minted for. A refresh
+ // must present the same id, so it is pinned here rather than re-derived.
+ ClientID string `json:"clientId,omitempty"`
}
// expired reports whether the access token is gone or about to lapse. The
@@ -154,7 +154,7 @@ func (c *Config) Bearer(ctx context.Context, path string) (string, error) {
if !c.OAuth.expired() {
return c.OAuth.AccessToken, nil
}
- tok, err := RefreshToken(ctx, c.URL, c.OAuth.RefreshToken, c.OAuth.CallbackPort)
+ tok, err := RefreshToken(ctx, c.URL, c.OAuth.RefreshToken, c.OAuth.ClientID)
if err != nil {
return "", err
}
diff --git a/internal/client/device.go b/internal/client/device.go
index 997d6ba..c43fd6e 100644
--- a/internal/client/device.go
+++ b/internal/client/device.go
@@ -1,12 +1,13 @@
package client
import (
- "bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
+ "net/url"
+ "os"
"strings"
"time"
@@ -14,7 +15,10 @@ import (
"golang.org/x/oauth2"
)
-// Poll outcomes, matching RFC 8628's error codes.
+// deviceGrantType is the RFC 8628 section 3.4 grant used when polling.
+const deviceGrantType = "urn:ietf:params:oauth:grant-type:device_code"
+
+// Poll outcomes, matching RFC 8628 section 3.5's error codes.
const (
errAuthorizationPending = "authorization_pending"
errSlowDown = "slow_down"
@@ -22,25 +26,32 @@ const (
errAccessDenied = "access_denied"
)
-// DeviceAuth is what the server hands back when a brokered login starts.
-type DeviceAuth struct {
- DeviceCode string `json:"deviceCode"`
- UserCode string `json:"userCode"`
- VerificationURI string `json:"verificationUri"`
- VerificationURIComplete string `json:"verificationUriComplete"`
- ExpiresIn int `json:"expiresIn"`
- Interval int `json:"interval"`
-}
-
-// LoginDevice authorizes this machine through the server rather than through a
-// local callback listener.
+// LoginDevice authorizes this machine with the OAuth device authorization
+// grant (RFC 8628), talking to the authorization server directly.
//
-// The collector opens no ports and needs no browser of its own: the user visits
-// a URL on the lard server from any browser anywhere, and this function polls
-// until the server reports a token. That is what makes login work unchanged over
-// SSH, inside a container, and on a headless box.
-func LoginDevice(ctx context.Context, serverURL string, openBrowser bool) (*oauth2.Token, error) {
- auth, err := startDevice(ctx, serverURL)
+// No listener, no browser, and no client secret are needed on this machine:
+// the device code itself is the proof of possession. The user approves the
+// user code at the AS's verification URI from any browser anywhere, which is
+// what makes login work unchanged over SSH, in a container, and on a headless
+// box.
+func LoginDevice(ctx context.Context, serverURL, clientID string, scopes []string, openBrowser bool) (*oauth2.Token, error) {
+ eps, err := Discover(ctx, serverURL)
+ if err != nil {
+ return nil, err
+ }
+ if eps.DeviceAuthorization == "" {
+ return nil, errors.New("the authorization server does not offer the device grant")
+ }
+ if clientID == "" {
+ if reg, err := FetchRegistration(ctx, serverURL); err == nil {
+ clientID = reg.ClientID
+ }
+ }
+ if clientID == "" {
+ return nil, errors.New("no client id: the server publishes no collector registration")
+ }
+
+ auth, err := startDevice(ctx, eps.DeviceAuthorization, clientID, scopes)
if err != nil {
return nil, err
}
@@ -48,7 +59,7 @@ func LoginDevice(ctx context.Context, serverURL string, openBrowser bool) (*oaut
interval := time.Duration(auth.Interval) * time.Second
if interval <= 0 {
- interval = 2 * time.Second
+ interval = 5 * time.Second // RFC 8628 §3.5 default
}
deadline := time.Now().Add(time.Duration(auth.ExpiresIn) * time.Second)
if auth.ExpiresIn <= 0 {
@@ -64,7 +75,7 @@ func LoginDevice(ctx context.Context, serverURL string, openBrowser bool) (*oaut
if time.Now().After(deadline) {
return nil, errors.New("this login expired; run 'lard-client login' again")
}
- tok, code, err := pollDevice(ctx, serverURL, auth.DeviceCode)
+ tok, code, err := pollDevice(ctx, eps.Token, clientID, auth.DeviceCode)
if err != nil {
return nil, err
}
@@ -75,9 +86,9 @@ func LoginDevice(ctx context.Context, serverURL string, openBrowser bool) (*oaut
case errAuthorizationPending:
// Still waiting on the human.
case errSlowDown:
- // The server says we are polling too fast; back off permanently
- // rather than just for this round.
- interval += time.Second
+ // RFC 8628 §3.5: increase the interval by 5 seconds for this and
+ // all subsequent requests.
+ interval += 5 * time.Second
case errExpiredToken:
return nil, errors.New("this login expired; run 'lard-client login' again")
case errAccessDenied:
@@ -88,52 +99,36 @@ func LoginDevice(ctx context.Context, serverURL string, openBrowser bool) (*oaut
}
}
-// printDevicePrompt shows the URL and code.
-//
-// The URL is always printed, even when a browser opens, because opening one is
-// a guess that is wrong over SSH and impossible on a headless machine. A
-// printed URL costs one line and always works.
-func printDevicePrompt(auth *DeviceAuth, openBrowser bool) {
- target := auth.VerificationURIComplete
- if target == "" {
- target = auth.VerificationURI
- }
- opened := false
- if openBrowser && isLocal() {
- opened = browser.OpenURL(target) == nil
- }
- if opened {
- fmt.Println("Opened your browser to authorize. If it opened on the wrong")
- fmt.Println("machine, use this URL from any browser instead:")
- } else {
- fmt.Println("Open this URL in any browser to authorize:")
- }
- // Bare and unstyled on its own line: terminals linkify a lone URL, and
- // decoration breaks copy and paste.
- fmt.Printf("\n%s\n\n", target)
- if auth.UserCode != "" {
- fmt.Printf("Your code: %s\n", auth.UserCode)
- }
- fmt.Println("Waiting for you to authorize...")
+// DeviceAuth mirrors the RFC 8628 device authorization response.
+type DeviceAuth struct {
+ DeviceCode string `json:"device_code"`
+ UserCode string `json:"user_code"`
+ VerificationURI string `json:"verification_uri"`
+ VerificationURIComplete string `json:"verification_uri_complete"`
+ ExpiresIn int `json:"expires_in"`
+ Interval int `json:"interval"`
}
-func startDevice(ctx context.Context, serverURL string) (*DeviceAuth, error) {
+// startDevice asks the authorization server for a device/user code pair.
+func startDevice(ctx context.Context, endpoint, clientID string, scopes []string) (*DeviceAuth, error) {
ctx, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
- req, err := http.NewRequestWithContext(ctx, http.MethodPost,
- strings.TrimRight(serverURL, "/")+"/auth/collector/device", strings.NewReader("{}"))
+ form := url.Values{"client_id": {clientID}}
+ if len(scopes) > 0 {
+ form.Set("scope", strings.Join(scopes, " "))
+ }
+ req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint,
+ strings.NewReader(form.Encode()))
if err != nil {
return nil, err
}
- req.Header.Set("content-type", "application/json")
+ req.Header.Set("content-type", "application/x-www-form-urlencoded")
+ req.Header.Set("accept", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
return nil, fmt.Errorf("starting authorization: %w", err)
}
defer resp.Body.Close()
- if resp.StatusCode == http.StatusNotFound {
- return nil, errors.New("this server does not offer brokered login")
- }
var auth DeviceAuth
if err := json.NewDecoder(resp.Body).Decode(&auth); err != nil {
return nil, fmt.Errorf("starting authorization: unreadable response (status %d)", resp.StatusCode)
@@ -144,21 +139,24 @@ func startDevice(ctx context.Context, serverURL string) (*DeviceAuth, error) {
return &auth, nil
}
-// pollDevice asks once whether authorization finished. A non-empty code string
-// is a status (pending, slow_down, ...) rather than a transport failure.
-func pollDevice(ctx context.Context, serverURL, deviceCode string) (*oauth2.Token, string, error) {
- body, err := json.Marshal(map[string]string{"deviceCode": deviceCode})
- if err != nil {
- return nil, "", err
- }
+// pollDevice asks the token endpoint once whether authorization finished. A
+// non-empty code string is an RFC 8628 status (pending, slow_down, ...)
+// rather than a transport failure.
+func pollDevice(ctx context.Context, tokenEndpoint, clientID, deviceCode string) (*oauth2.Token, string, error) {
ctx, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
- req, err := http.NewRequestWithContext(ctx, http.MethodPost,
- strings.TrimRight(serverURL, "/")+"/auth/collector/device/token", bytes.NewReader(body))
+ form := url.Values{
+ "grant_type": {deviceGrantType},
+ "device_code": {deviceCode},
+ "client_id": {clientID},
+ }
+ req, err := http.NewRequestWithContext(ctx, http.MethodPost, tokenEndpoint,
+ strings.NewReader(form.Encode()))
if err != nil {
return nil, "", err
}
- req.Header.Set("content-type", "application/json")
+ req.Header.Set("content-type", "application/x-www-form-urlencoded")
+ req.Header.Set("accept", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
// A transient network blip should not abandon a login the user is
@@ -167,10 +165,11 @@ func pollDevice(ctx context.Context, serverURL, deviceCode string) (*oauth2.Toke
}
defer resp.Body.Close()
var out struct {
- AccessToken string `json:"accessToken"`
- RefreshToken string `json:"refreshToken"`
- Expiry time.Time `json:"expiry"`
- Error string `json:"error"`
+ AccessToken string `json:"access_token"`
+ RefreshToken string `json:"refresh_token"`
+ TokenType string `json:"token_type"`
+ ExpiresIn int64 `json:"expires_in"`
+ Error string `json:"error"`
}
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
return nil, "", fmt.Errorf("polling: unreadable response (status %d)", resp.StatusCode)
@@ -181,10 +180,48 @@ func pollDevice(ctx context.Context, serverURL, deviceCode string) (*oauth2.Toke
if out.AccessToken == "" {
return nil, "", fmt.Errorf("polling: server returned status %d", resp.StatusCode)
}
- return &oauth2.Token{
+ tok := &oauth2.Token{
AccessToken: out.AccessToken,
RefreshToken: out.RefreshToken,
- Expiry: out.Expiry,
TokenType: "Bearer",
- }, "", nil
+ }
+ if out.ExpiresIn > 0 {
+ tok.Expiry = time.Now().Add(time.Duration(out.ExpiresIn) * time.Second)
+ }
+ return tok, "", nil
+}
+
+// isLocal guesses whether a browser on this machine is reachable by the user.
+// SSH sets these, and their absence is the common case on a laptop.
+func isLocal() bool {
+ return os.Getenv("SSH_CONNECTION") == "" && os.Getenv("SSH_CLIENT") == "" && os.Getenv("SSH_TTY") == ""
+}
+
+// printDevicePrompt shows the URL and code.
+//
+// The URL is always printed, even when a browser opens, because opening one is
+// a guess that is wrong over SSH and impossible on a headless machine. A
+// printed URL costs one line and always works.
+func printDevicePrompt(auth *DeviceAuth, openBrowser bool) {
+ target := auth.VerificationURIComplete
+ if target == "" {
+ target = auth.VerificationURI
+ }
+ opened := false
+ if openBrowser && isLocal() {
+ opened = browser.OpenURL(target) == nil
+ }
+ if opened {
+ fmt.Println("Opened your browser to authorize. If it opened on the wrong")
+ fmt.Println("machine, use this URL from any browser instead:")
+ } else {
+ fmt.Println("Open this URL in any browser to authorize:")
+ }
+ // Bare and unstyled on its own line: terminals linkify a lone URL, and
+ // decoration breaks copy and paste.
+ fmt.Printf("\n%s\n\n", target)
+ if auth.UserCode != "" {
+ fmt.Printf("Your code: %s\n", auth.UserCode)
+ }
+ fmt.Println("Waiting for you to authorize...")
}
diff --git a/internal/client/oauth.go b/internal/client/oauth.go
index b28b2b2..e37a3d9 100644
--- a/internal/client/oauth.go
+++ b/internal/client/oauth.go
@@ -1,75 +1,27 @@
package client
import (
- "bytes"
"context"
- "crypto/rand"
- "encoding/base64"
"encoding/json"
"errors"
"fmt"
- "html"
- "log/slog"
- "net"
"net/http"
"net/url"
- "os"
- "strconv"
"strings"
"time"
- "github.com/pkg/browser"
"golang.org/x/oauth2"
)
-// CallbackPort is the localhost port the login flow listens on by default.
-// Crush uses 40704-40713 for its own MCP flow, so this sits just past that
-// range. The server may publish a different set, which wins.
-const CallbackPort = 40714
-
-// ClientID is the fallback OAuth client id, used only when the server publishes
-// no collector registration. Deriving it from the callback port satisfies
-// authorization servers that require the client id host to match the redirect
-// host, but a server-published id is always preferred: the server is the one
-// that decides which clients it trusts.
-func ClientID(port int) string {
- return fmt.Sprintf("http://localhost:%d/", port)
-}
-
// Registration is the OAuth client identity a lard server tells collectors to
// use, fetched from its collector endpoint.
type Registration struct {
- ClientID string `json:"clientId"`
- RedirectURIs []string `json:"redirectUris"`
- Scopes []string `json:"scopes"`
- // ServerExchange means the server holds a client secret and will trade the
- // authorization code for a token on our behalf, so the secret stays off
- // this machine.
- ServerExchange bool `json:"serverExchange"`
- // DeviceFlow means the server brokers the whole authorization, so this
- // machine needs no callback listener and no browser of its own.
- DeviceFlow bool `json:"deviceFlow"`
-}
-
-// Ports extracts the callback ports from the published redirect URIs.
-func (r *Registration) Ports() []int {
- var out []int
- for _, raw := range r.RedirectURIs {
- u, err := url.Parse(raw)
- if err != nil {
- continue
- }
- if _, portStr, err := net.SplitHostPort(u.Host); err == nil {
- if n, err := strconv.Atoi(portStr); err == nil {
- out = append(out, n)
- }
- }
- }
- return out
+ ClientID string `json:"clientId"`
+ Scopes []string `json:"scopes"`
}
// FetchRegistration asks the server which OAuth client to be. A 404 means the
-// server publishes none, and the caller falls back to a localhost client id.
+// server publishes none.
func FetchRegistration(ctx context.Context, serverURL string) (*Registration, error) {
ctx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
@@ -85,15 +37,17 @@ func FetchRegistration(ctx context.Context, serverURL string) (*Registration, er
// endpoints are the OAuth endpoints discovered from a lard server.
type endpoints struct {
- Issuer string
- Authorization string
- Token string
+ Issuer string
+ Authorization string
+ Token string
+ DeviceAuthorization string
+ Revocation string
}
// Discover walks the RFC 9728 -> RFC 8414 chain: ask the lard server which
// authorization server protects it, then ask that server where its endpoints
-// are. Nothing about indiko is hardcoded, so pointing lard at a different
-// provider needs no client change.
+// are. Nothing about the provider is hardcoded, so pointing lard at a
+// different authorization server needs no client change.
func Discover(ctx context.Context, serverURL string) (*endpoints, error) {
ctx, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
@@ -111,9 +65,11 @@ func Discover(ctx context.Context, serverURL string) (*endpoints, error) {
as := strings.TrimRight(prm.AuthorizationServers[0], "/")
var meta struct {
- Issuer string `json:"issuer"`
- Authorization string `json:"authorization_endpoint"`
- Token string `json:"token_endpoint"`
+ Issuer string `json:"issuer"`
+ Authorization string `json:"authorization_endpoint"`
+ Token string `json:"token_endpoint"`
+ DeviceAuthorization string `json:"device_authorization_endpoint"`
+ Revocation string `json:"revocation_endpoint"`
}
if err := getJSON(ctx, as+"/.well-known/oauth-authorization-server", &meta); err != nil {
return nil, fmt.Errorf("discover authorization server: %w", err)
@@ -121,7 +77,13 @@ func Discover(ctx context.Context, serverURL string) (*endpoints, error) {
if meta.Authorization == "" || meta.Token == "" {
return nil, errors.New("authorization server metadata is missing endpoints")
}
- return &endpoints{Issuer: meta.Issuer, Authorization: meta.Authorization, Token: meta.Token}, nil
+ return &endpoints{
+ Issuer: meta.Issuer,
+ Authorization: meta.Authorization,
+ Token: meta.Token,
+ DeviceAuthorization: meta.DeviceAuthorization,
+ Revocation: meta.Revocation,
+ }, nil
}
func getJSON(ctx context.Context, url string, out any) error {
@@ -141,244 +103,20 @@ func getJSON(ctx context.Context, url string, out any) error {
return json.NewDecoder(resp.Body).Decode(out)
}
-// Login runs the browser authorization-code flow with PKCE and returns the
-// token. The caller persists it.
-// Login runs the browser authorization-code flow with PKCE and returns the
-// token.
-//
-// The OAuth client identity comes from the server when it publishes one, since
-// the server is what decides which clients it accepts. A client id guessed here
-// would be rejected by the very server we are trying to reach.
-func Login(ctx context.Context, serverURL string, port int, openBrowser bool) (*oauth2.Token, error) {
- eps, err := Discover(ctx, serverURL)
- if err != nil {
- return nil, err
- }
-
- clientID := ""
- scopes := []string{"profile"}
- serverExchange := false
- ports := []int{}
- if reg, err := FetchRegistration(ctx, serverURL); err == nil {
- clientID = reg.ClientID
- serverExchange = reg.ServerExchange
- if len(reg.Scopes) > 0 {
- scopes = reg.Scopes
- }
- ports = reg.Ports()
- } else {
- slog.Debug("no collector registration published; using a localhost client id", "error", err)
- }
- // An explicitly requested port wins, then the server's list, then the
- // default. The chosen port must appear in the server's list or it will
- // refuse the exchange.
- if port > 0 {
- ports = append([]int{port}, ports...)
- }
- if len(ports) == 0 {
- ports = []int{CallbackPort}
- }
-
- // Bind before sending the user to the browser, so a busy port fails now
- // rather than after they have authenticated.
- ln, boundPort, err := listenAny(ports)
- if err != nil {
- return nil, err
- }
- defer ln.Close()
-
- if clientID == "" {
- clientID = ClientID(boundPort)
- }
- redirectURI := fmt.Sprintf("http://localhost:%d/callback", boundPort)
- cfg := &oauth2.Config{
- ClientID: clientID,
- RedirectURL: redirectURI,
- Endpoint: oauth2.Endpoint{
- AuthURL: eps.Authorization,
- TokenURL: eps.Token,
- AuthStyle: oauth2.AuthStyleInParams,
- },
- Scopes: scopes,
- }
-
- verifier := oauth2.GenerateVerifier()
- state, err := randomState()
- if err != nil {
- return nil, err
- }
- authURL := cfg.AuthCodeURL(state,
- oauth2.S256ChallengeOption(verifier),
- oauth2.AccessTypeOffline, // ask for a refresh token
- )
-
- // Over SSH the local browser is the wrong browser, so don't try to launch
- // one: the printed URL is what the user actually needs.
- code, err := awaitCode(ctx, ln, state, authURL, openBrowser && isLocal())
- if err != nil {
- return nil, err
- }
-
- // A confidential client's secret lives on the server, so the server
- // finishes the exchange. We still hold the PKCE verifier, so neither side
- // can complete it alone.
- if serverExchange {
- return exchangeViaServer(ctx, serverURL, code, verifier, redirectURI)
- }
- tok, err := cfg.Exchange(ctx, code, oauth2.VerifierOption(verifier))
- if err != nil {
- return nil, fmt.Errorf("exchange code for token: %w", err)
- }
- return tok, nil
-}
-
-// listenAny binds the first available port from the list.
-func listenAny(ports []int) (net.Listener, int, error) {
- var lastErr error
- for _, p := range ports {
- ln, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", p))
- if err == nil {
- return ln, p, nil
- }
- lastErr = err
- }
- return nil, 0, fmt.Errorf("no callback port available (tried %v): %w", ports, lastErr)
-}
-
-// awaitCode serves the redirect callback and returns the authorization code.
-func awaitCode(ctx context.Context, ln net.Listener, state, authURL string, openBrowser bool) (string, error) {
- type result struct {
- code string
- err error
- }
- results := make(chan result, 1)
- srv := &http.Server{
- Handler: http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- q := r.URL.Query()
- switch {
- case q.Get("error") != "":
- finish(w, "Authorization failed: "+q.Get("error"))
- results <- result{err: fmt.Errorf("authorization denied: %s", q.Get("error"))}
- case q.Get("state") != state:
- finish(w, "Authorization failed: state mismatch.")
- results <- result{err: errors.New("state mismatch; possible CSRF, aborting")}
- case q.Get("code") == "":
- finish(w, "Authorization failed: no code returned.")
- results <- result{err: errors.New("no authorization code in callback")}
- default:
- finish(w, "Connected. You can close this tab and return to the terminal.")
- results <- result{code: q.Get("code")}
- }
- }),
- ReadHeaderTimeout: 10 * time.Second,
- }
- go func() { _ = srv.Serve(ln) }()
- defer func() {
- shutdownCtx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
- defer cancel()
- _ = srv.Shutdown(shutdownCtx)
- }()
-
- printAuthURL(authURL, listenPort(ln), openBrowser)
-
- // Give the human a few minutes to find their passkey.
- waitCtx, cancel := context.WithTimeout(ctx, 5*time.Minute)
- defer cancel()
- select {
- case <-waitCtx.Done():
- return "", errors.New("timed out waiting for authorization")
- case res := <-results:
- return res.code, res.err
- }
-}
-
-// exchangeViaServer asks lard to trade the code for a token using its own
-// confidential client credentials.
-func exchangeViaServer(ctx context.Context, serverURL, code, verifier, redirectURI string) (*oauth2.Token, error) {
- body, err := json.Marshal(map[string]string{
- "code": code,
- "codeVerifier": verifier,
- "redirectUri": redirectURI,
- })
- if err != nil {
- return nil, err
- }
- ctx, cancel := context.WithTimeout(ctx, 30*time.Second)
- defer cancel()
- req, err := http.NewRequestWithContext(ctx, http.MethodPost,
- strings.TrimRight(serverURL, "/")+"/auth/collector/exchange", bytes.NewReader(body))
- if err != nil {
- return nil, err
- }
- req.Header.Set("content-type", "application/json")
- resp, err := http.DefaultClient.Do(req)
- if err != nil {
- return nil, fmt.Errorf("server-side exchange: %w", err)
- }
- defer resp.Body.Close()
- var out struct {
- AccessToken string `json:"accessToken"`
- RefreshToken string `json:"refreshToken"`
- Expiry time.Time `json:"expiry"`
- Error string `json:"error"`
- }
- if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
- return nil, fmt.Errorf("server-side exchange returned unreadable JSON (status %d)", resp.StatusCode)
- }
- if out.Error != "" {
- return nil, fmt.Errorf("server-side exchange: %s", out.Error)
- }
- if out.AccessToken == "" {
- return nil, fmt.Errorf("server-side exchange returned status %d", resp.StatusCode)
- }
- return &oauth2.Token{
- AccessToken: out.AccessToken,
- RefreshToken: out.RefreshToken,
- Expiry: out.Expiry,
- TokenType: "Bearer",
- }, nil
-}
-
-func finish(w http.ResponseWriter, msg string) {
- w.Header().Set("content-type", "text/html; charset=utf-8")
- fmt.Fprintf(w, `lard
-
-
-%s
`, html.EscapeString(msg))
-}
-
-func randomState() (string, error) {
- b := make([]byte, 16)
- if _, err := rand.Read(b); err != nil {
- return "", err
- }
- return base64.RawURLEncoding.EncodeToString(b), nil
-}
-
-// RefreshToken exchanges a refresh token for a fresh access token. The daemon
-// uses this, since it must never try to open a browser.
-//
-// A confidential registration is refreshed through the server for the same
-// reason its code was exchanged there: the client secret is required and lives
-// only on the server.
-func RefreshToken(ctx context.Context, serverURL, refresh string, port int) (*oauth2.Token, error) {
+// RefreshToken exchanges a refresh token for a fresh access token at the
+// authorization server. The daemon uses this, since it must never try to open
+// a browser. clientID is the public client the refresh token was minted for.
+func RefreshToken(ctx context.Context, serverURL, refresh, clientID string) (*oauth2.Token, error) {
if refresh == "" {
return nil, errors.New("no refresh token saved; run 'lard-client login'")
}
- if reg, err := FetchRegistration(ctx, serverURL); err == nil && reg.ServerExchange {
- return refreshViaServer(ctx, serverURL, refresh)
- }
- if port <= 0 {
- port = CallbackPort
+ if clientID == "" {
+ return nil, errors.New("no client id saved; run 'lard-client login' again")
}
eps, err := Discover(ctx, serverURL)
if err != nil {
return nil, err
}
- clientID := ClientID(port)
- if reg, err := FetchRegistration(ctx, serverURL); err == nil {
- clientID = reg.ClientID
- }
cfg := &oauth2.Config{
ClientID: clientID,
Endpoint: oauth2.Endpoint{
@@ -394,95 +132,32 @@ func RefreshToken(ctx context.Context, serverURL, refresh string, port int) (*oa
return tok, nil
}
-// refreshViaServer asks lard to refresh using its confidential credentials.
-func refreshViaServer(ctx context.Context, serverURL, refresh string) (*oauth2.Token, error) {
- body, err := json.Marshal(map[string]string{"refreshToken": refresh})
+// RevokeToken tells the authorization server to kill a token (RFC 7009), so a
+// logged-out machine's refresh token stops working anywhere, not just locally.
+// Best-effort: the server answers 200 even for a token it doesn't know.
+func RevokeToken(ctx context.Context, serverURL, token string) error {
+ eps, err := Discover(ctx, serverURL)
if err != nil {
- return nil, err
+ return err
+ }
+ if eps.Revocation == "" {
+ return errors.New("the authorization server does not offer token revocation")
}
- ctx, cancel := context.WithTimeout(ctx, 30*time.Second)
+ ctx, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
- req, err := http.NewRequestWithContext(ctx, http.MethodPost,
- strings.TrimRight(serverURL, "/")+"/auth/collector/refresh", bytes.NewReader(body))
+ req, err := http.NewRequestWithContext(ctx, http.MethodPost, eps.Revocation,
+ strings.NewReader(url.Values{"token": {token}}.Encode()))
if err != nil {
- return nil, err
+ return err
}
- req.Header.Set("content-type", "application/json")
+ req.Header.Set("content-type", "application/x-www-form-urlencoded")
resp, err := http.DefaultClient.Do(req)
if err != nil {
- return nil, fmt.Errorf("server-side refresh: %w", err)
+ return fmt.Errorf("revoking token: %w", err)
}
defer resp.Body.Close()
- var out struct {
- AccessToken string `json:"accessToken"`
- RefreshToken string `json:"refreshToken"`
- Expiry time.Time `json:"expiry"`
- Error string `json:"error"`
- }
- if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
- return nil, fmt.Errorf("server-side refresh returned unreadable JSON (status %d)", resp.StatusCode)
- }
- if out.Error != "" {
- return nil, fmt.Errorf("server-side refresh: %s (re-run 'lard-client login')", out.Error)
- }
- if out.AccessToken == "" {
- return nil, fmt.Errorf("server-side refresh returned status %d", resp.StatusCode)
- }
- tok := &oauth2.Token{AccessToken: out.AccessToken, Expiry: out.Expiry, TokenType: "Bearer"}
- // Keep the existing refresh token when the server does not rotate it.
- tok.RefreshToken = out.RefreshToken
- if tok.RefreshToken == "" {
- tok.RefreshToken = refresh
- }
- return tok, nil
-}
-
-// printAuthURL shows the authorization URL and, always, the raw URL itself.
-//
-// The URL is printed even when a browser opens successfully. Opening a browser
-// is a guess: over SSH it launches on the wrong machine, in a container it
-// launches nowhere, and a headless box has none. Printing the URL costs one
-// screen of text and is the difference between a working login and a dead end.
-func printAuthURL(authURL string, port int, openBrowser bool) {
- opened := false
- if openBrowser {
- opened = browser.OpenURL(authURL) == nil
- }
- if opened {
- fmt.Println("Opened your browser to authorize.")
- fmt.Println("If it opened on the wrong machine, use this URL instead:")
- } else {
- fmt.Println("Open this URL to authorize:")
- }
- // Bare, on its own line, unwrapped and unstyled: terminals turn a lone URL
- // into a click target, and anything decorative breaks copy and paste.
- fmt.Printf("\n%s\n\n", authURL)
- fmt.Printf("Waiting for the callback on localhost:%d...\n", port)
- if !isLocal() {
- fmt.Printf("\nThis looks like a remote session. The callback goes to localhost:%d\n", port)
- fmt.Printf("on *this* machine, so forward it from your laptop first:\n")
- fmt.Printf(" ssh -L %d:localhost:%d %s\n", port, port, remoteHostHint())
- }
-}
-
-// listenPort reports the port a listener is bound to.
-func listenPort(ln net.Listener) int {
- if addr, ok := ln.Addr().(*net.TCPAddr); ok {
- return addr.Port
- }
- return 0
-}
-
-// isLocal guesses whether a browser on this machine is reachable by the user.
-// SSH sets these, and their absence is the common case on a laptop.
-func isLocal() bool {
- return os.Getenv("SSH_CONNECTION") == "" && os.Getenv("SSH_CLIENT") == "" && os.Getenv("SSH_TTY") == ""
-}
-
-// remoteHostHint names this host for the suggested ssh command.
-func remoteHostHint() string {
- if h, err := os.Hostname(); err == nil && h != "" {
- return h
+ if resp.StatusCode != http.StatusOK {
+ return fmt.Errorf("revoking token: status %d", resp.StatusCode)
}
- return "this-host"
+ return nil
}
diff --git a/internal/client/oauth_test.go b/internal/client/oauth_test.go
deleted file mode 100644
index 100325f..0000000
--- a/internal/client/oauth_test.go
+++ /dev/null
@@ -1,151 +0,0 @@
-package client
-
-import (
- "context"
- "net"
- "net/http"
- "net/http/httptest"
- "os"
- "strings"
- "testing"
-)
-
-// fakeLard serves the discovery chain a real lard publishes.
-func fakeLard(t *testing.T) (lardURL string) {
- t.Helper()
- as := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if r.URL.Path != "/.well-known/oauth-authorization-server" {
- http.NotFound(w, r)
- return
- }
- w.Header().Set("content-type", "application/json")
- w.Write([]byte(`{"issuer":"` + asBase(r) + `","authorization_endpoint":"` + asBase(r) + `/auth/authorize","token_endpoint":"` + asBase(r) + `/auth/token","code_challenge_methods_supported":["S256"]}`))
- }))
- t.Cleanup(as.Close)
-
- lard := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- if !strings.HasPrefix(r.URL.Path, "/.well-known/oauth-protected-resource") {
- http.NotFound(w, r)
- return
- }
- w.Header().Set("content-type", "application/json")
- w.Write([]byte(`{"resource":"x","authorization_servers":["` + as.URL + `"]}`))
- }))
- t.Cleanup(lard.Close)
- return lard.URL
-}
-
-func asBase(r *http.Request) string { return "http://" + r.Host }
-
-func TestDiscoverWalksResourceToAuthServer(t *testing.T) {
- eps, err := Discover(context.Background(), fakeLard(t))
- if err != nil {
- t.Fatal(err)
- }
- if !strings.HasSuffix(eps.Authorization, "/auth/authorize") {
- t.Errorf("authorization endpoint = %q", eps.Authorization)
- }
- if !strings.HasSuffix(eps.Token, "/auth/token") {
- t.Errorf("token endpoint = %q", eps.Token)
- }
- if eps.Issuer == "" {
- t.Error("issuer empty")
- }
-}
-
-// A server with auth off publishes nothing, and the client must say so plainly
-// rather than opening a browser at a nonexistent endpoint.
-func TestDiscoverFailsWhenServerHasNoOAuth(t *testing.T) {
- srv := httptest.NewServer(http.HandlerFunc(http.NotFound))
- defer srv.Close()
- if _, err := Discover(context.Background(), srv.URL); err == nil {
- t.Fatal("want an error when discovery is unavailable")
- }
-}
-
-// The client id must be derived from the callback port, since indiko rejects a
-// client id whose host differs from the redirect host.
-func TestClientIDMatchesCallbackHost(t *testing.T) {
- if got := ClientID(40714); got != "http://localhost:40714/" {
- t.Fatalf("got %q", got)
- }
-}
-
-// Over SSH a local browser is the wrong browser, so the flow must know the
-// difference and lean on the printed URL instead.
-func TestIsLocalDetectsSSH(t *testing.T) {
- t.Setenv("SSH_CONNECTION", "")
- t.Setenv("SSH_CLIENT", "")
- t.Setenv("SSH_TTY", "")
- os.Unsetenv("SSH_CONNECTION")
- os.Unsetenv("SSH_CLIENT")
- os.Unsetenv("SSH_TTY")
- if !isLocal() {
- t.Error("no SSH vars should read as local")
- }
- for _, k := range []string{"SSH_CONNECTION", "SSH_CLIENT", "SSH_TTY"} {
- t.Run(k, func(t *testing.T) {
- t.Setenv(k, "set")
- if isLocal() {
- t.Errorf("%s set should read as remote", k)
- }
- })
- }
-}
-
-func TestListenPortReportsBoundPort(t *testing.T) {
- ln, err := net.Listen("tcp", "127.0.0.1:0")
- if err != nil {
- t.Fatal(err)
- }
- defer ln.Close()
- if got := listenPort(ln); got <= 0 {
- t.Fatalf("listenPort = %d, want a real port", got)
- }
-}
-
-// listenAny must fall through to the next port rather than failing, so a
-// leftover process on one port does not block login.
-func TestListenAnyFallsThroughBusyPorts(t *testing.T) {
- busy, err := net.Listen("tcp", "127.0.0.1:0")
- if err != nil {
- t.Fatal(err)
- }
- defer busy.Close()
- busyPort := busy.Addr().(*net.TCPAddr).Port
-
- ln, port, err := listenAny([]int{busyPort, 0})
- if err != nil {
- t.Fatal(err)
- }
- defer ln.Close()
- if port == busyPort {
- t.Fatalf("bound the busy port %d", busyPort)
- }
-}
-
-func TestListenAnyFailsWhenAllBusy(t *testing.T) {
- busy, err := net.Listen("tcp", "127.0.0.1:0")
- if err != nil {
- t.Fatal(err)
- }
- defer busy.Close()
- p := busy.Addr().(*net.TCPAddr).Port
- if _, _, err := listenAny([]int{p}); err == nil {
- t.Fatal("want an error when every port is taken")
- }
-}
-
-// The registration's redirect URIs are the source of truth for which ports the
-// server will accept, so parsing them must not silently drop any.
-func TestRegistrationPorts(t *testing.T) {
- reg := &Registration{RedirectURIs: []string{
- "http://localhost:40714/callback",
- "http://localhost:40715/callback",
- "not a url at all::",
- }}
- got := reg.Ports()
- if len(got) != 2 || got[0] != 40714 || got[1] != 40715 {
- t.Fatalf("Ports() = %v", got)
- }
-}
diff --git a/internal/collector/collector.go b/internal/collector/collector.go
index d157562..3a14925 100644
--- a/internal/collector/collector.go
+++ b/internal/collector/collector.go
@@ -1,75 +1,35 @@
-// Package collector serves the OAuth registration that edge collectors use to
-// authenticate.
+// Package collector serves the OAuth client identity that edge collectors
+// adopt when logging in.
//
-// The problem this solves: a collector cannot invent its own OAuth client id.
-// The authorization server decides which clients exist, and lard decides which
-// clients it trusts, so a client id guessed by the collector is one the server
-// will reject. Instead the server publishes the identity to use, and the
-// collector adopts it.
-//
-// When the registration is confidential (a client secret), the collector must
-// not hold that secret: it is a public CLI on a laptop. So the server also
-// performs the code exchange on the collector's behalf, keeping the secret
-// server-side while the collector keeps its PKCE verifier. The collector proves
-// possession with PKCE; the server proves the client's identity with the
-// secret.
+// A collector cannot invent its own client id: the authorization server
+// decides which clients exist, and lard decides which clients it trusts, so a
+// client id guessed by the collector is one the server will reject. Instead
+// the server publishes the identity to use, and the collector adopts it, then
+// runs the OAuth device authorization grant (RFC 8628) against the
+// authorization server directly. Login, exchange, and refresh all happen
+// provider-side; this package is just the registration document.
package collector
import (
- "context"
"encoding/json"
- "errors"
- "fmt"
- "log/slog"
- "net"
"net/http"
- "net/url"
- "slices"
- "strconv"
- "strings"
- "time"
)
-// DefaultPorts are the localhost callback ports a collector may bind, in
-// preference order. They are published so the collector picks one the
-// authorization server already knows, and validated so a stolen code cannot be
-// redirected somewhere else.
-var DefaultPorts = []int{40714, 40715, 40716, 40717, 40718}
-
// DefaultScopes is what a collector asks for: enough to identify the user.
var DefaultScopes = []string{"profile"}
// Config describes the collector OAuth registration this server hands out.
type Config struct {
// ClientID is the OAuth client collectors authenticate as. Empty means no
- // registration is published, and collectors fall back to their own
- // localhost client id.
+ // registration is published.
ClientID string
- // ClientSecret, when set, marks the registration confidential. The server
- // then exchanges codes itself so the secret never leaves this process.
- ClientSecret string
- // Ports are the permitted localhost callback ports.
- Ports []int
// Scopes the collector should request.
Scopes []string
- // TokenEndpoint is the authorization server's token endpoint, used for the
- // server-side exchange.
- TokenEndpoint string
}
// Configured reports whether a registration is published.
func (c Config) Configured() bool { return c.ClientID != "" }
-// Confidential reports whether the exchange must happen server-side.
-func (c Config) Confidential() bool { return c.ClientSecret != "" }
-
-func (c Config) ports() []int {
- if len(c.Ports) > 0 {
- return c.Ports
- }
- return DefaultPorts
-}
-
func (c Config) scopes() []string {
if len(c.Scopes) > 0 {
return c.Scopes
@@ -77,43 +37,20 @@ func (c Config) scopes() []string {
return DefaultScopes
}
-// redirectURIs lists the callbacks the collector may use.
-func (c Config) redirectURIs() []string {
- var out []string
- for _, p := range c.ports() {
- out = append(out, fmt.Sprintf("http://localhost:%d/callback", p))
- }
- return out
-}
-
// Registration is the document a collector fetches before logging in.
type Registration struct {
- ClientID string `json:"clientId"`
- RedirectURIs []string `json:"redirectUris"`
- Scopes []string `json:"scopes"`
- // ServerExchange tells the collector to post its code back here instead of
- // calling the token endpoint directly, because the secret lives on this
- // side.
- ServerExchange bool `json:"serverExchange"`
- // DeviceFlow means this server brokers the whole authorization: the
- // collector never runs a callback listener, so login works over SSH, in a
- // container, and on headless machines.
- DeviceFlow bool `json:"deviceFlow"`
+ ClientID string `json:"clientId"`
+ Scopes []string `json:"scopes"`
}
-// Handler serves the collector registration, the brokered device login, and
-// (for confidential clients) the code exchange.
+// Handler serves the collector registration.
type Handler struct {
- cfg Config
- dev DeviceConfig
- prefix string
- devices *deviceStore
+ cfg Config
}
-// New builds the handler. prefix is the URL prefix the handler is mounted at,
-// needed so the browser-facing URLs it builds are absolute.
-func New(cfg Config, dev DeviceConfig, prefix string) *Handler {
- return &Handler{cfg: cfg, dev: dev, prefix: prefix, devices: newDeviceStore()}
+// New builds the handler.
+func New(cfg Config) *Handler {
+ return &Handler{cfg: cfg}
}
// Register serves GET: the registration a collector should adopt.
@@ -124,189 +61,7 @@ func (h *Handler) Register(w http.ResponseWriter, r *http.Request) {
}
w.Header().Set("content-type", "application/json")
_ = json.NewEncoder(w).Encode(Registration{
- ClientID: h.cfg.ClientID,
- RedirectURIs: h.cfg.redirectURIs(),
- Scopes: h.cfg.scopes(),
- ServerExchange: h.cfg.Confidential(),
- DeviceFlow: h.DeviceAvailable(),
- })
-}
-
-// exchangeRequest is what a collector posts after the user authorizes.
-type exchangeRequest struct {
- Code string `json:"code"`
- CodeVerifier string `json:"codeVerifier"`
- RedirectURI string `json:"redirectUri"`
-}
-
-// Exchange serves POST: trade an authorization code for a token using the
-// server's confidential client credentials.
-//
-// The collector supplies the code and its PKCE verifier; this server supplies
-// the client secret. Neither side alone can complete the exchange, which is the
-// point.
-func (h *Handler) Exchange(w http.ResponseWriter, r *http.Request) {
- if !h.cfg.Confidential() {
- http.Error(w, `{"error":"server-side exchange is not enabled"}`, http.StatusNotFound)
- return
- }
- var req exchangeRequest
- if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 1<<16)).Decode(&req); err != nil {
- writeError(w, http.StatusBadRequest, "malformed request")
- return
- }
- if req.Code == "" || req.CodeVerifier == "" {
- writeError(w, http.StatusBadRequest, "code and codeVerifier are required")
- return
- }
- // Only redirect back to a callback we published. Without this check the
- // endpoint would exchange a code for any redirect a caller names, which
- // turns the server's credentials into an oracle for stolen codes.
- if err := h.validateRedirect(req.RedirectURI); err != nil {
- writeError(w, http.StatusBadRequest, err.Error())
- return
- }
- tok, err := h.exchange(r.Context(), req)
- if err != nil {
- slog.Warn("collector: code exchange failed", "error", err)
- writeError(w, http.StatusBadGateway, err.Error())
- return
- }
- w.Header().Set("content-type", "application/json")
- w.Header().Set("cache-control", "no-store")
- _ = json.NewEncoder(w).Encode(tok)
-}
-
-// validateRedirect confirms the redirect is one of the published localhost
-// callbacks.
-func (h *Handler) validateRedirect(raw string) error {
- if raw == "" {
- return errors.New("redirectUri is required")
- }
- if slices.Contains(h.cfg.redirectURIs(), raw) {
- return nil
- }
- u, err := url.Parse(raw)
- if err != nil {
- return errors.New("redirectUri is not a valid URL")
- }
- host, portStr, err := net.SplitHostPort(u.Host)
- if err != nil || (host != "localhost" && host != "127.0.0.1") {
- return errors.New("redirectUri must be a localhost callback")
- }
- port, err := strconv.Atoi(portStr)
- if err != nil || !slices.Contains(h.cfg.ports(), port) {
- return fmt.Errorf("redirectUri port must be one of %v", h.cfg.ports())
- }
- return nil
-}
-
-// TokenResponse is the subset of the token response a collector needs.
-type TokenResponse struct {
- AccessToken string `json:"accessToken"`
- RefreshToken string `json:"refreshToken,omitempty"`
- Expiry time.Time `json:"expiry,omitempty"`
- TokenType string `json:"tokenType,omitempty"`
-}
-
-func (h *Handler) exchange(ctx context.Context, req exchangeRequest) (*TokenResponse, error) {
- if h.cfg.TokenEndpoint == "" {
- return nil, errors.New("server has no token endpoint configured")
- }
- ctx, cancel := context.WithTimeout(ctx, 20*time.Second)
- defer cancel()
- form := url.Values{
- "grant_type": {"authorization_code"},
- "code": {req.Code},
- "client_id": {h.cfg.ClientID},
- "client_secret": {h.cfg.ClientSecret},
- "redirect_uri": {req.RedirectURI},
- "code_verifier": {req.CodeVerifier},
- }
- return postToken(ctx, h.cfg.TokenEndpoint, form)
-}
-
-// Refresh trades a refresh token for a new access token, again using the
-// server's credentials so the collector never needs them.
-func (h *Handler) Refresh(w http.ResponseWriter, r *http.Request) {
- if !h.cfg.Confidential() {
- http.Error(w, `{"error":"server-side refresh is not enabled"}`, http.StatusNotFound)
- return
- }
- var req struct {
- RefreshToken string `json:"refreshToken"`
- }
- if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 1<<16)).Decode(&req); err != nil {
- writeError(w, http.StatusBadRequest, "malformed request")
- return
- }
- if req.RefreshToken == "" {
- writeError(w, http.StatusBadRequest, "refreshToken is required")
- return
- }
- ctx, cancel := context.WithTimeout(r.Context(), 20*time.Second)
- defer cancel()
- tok, err := postToken(ctx, h.cfg.TokenEndpoint, url.Values{
- "grant_type": {"refresh_token"},
- "refresh_token": {req.RefreshToken},
- "client_id": {h.cfg.ClientID},
- "client_secret": {h.cfg.ClientSecret},
+ ClientID: h.cfg.ClientID,
+ Scopes: h.cfg.scopes(),
})
- if err != nil {
- writeError(w, http.StatusBadGateway, err.Error())
- return
- }
- w.Header().Set("content-type", "application/json")
- w.Header().Set("cache-control", "no-store")
- _ = json.NewEncoder(w).Encode(tok)
-}
-
-func postToken(ctx context.Context, endpoint string, form url.Values) (*TokenResponse, error) {
- httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint,
- strings.NewReader(form.Encode()))
- if err != nil {
- return nil, err
- }
- httpReq.Header.Set("content-type", "application/x-www-form-urlencoded")
- httpReq.Header.Set("accept", "application/json")
- resp, err := http.DefaultClient.Do(httpReq)
- if err != nil {
- return nil, fmt.Errorf("reaching the token endpoint: %w", err)
- }
- defer resp.Body.Close()
- var body struct {
- AccessToken string `json:"access_token"`
- RefreshToken string `json:"refresh_token"`
- TokenType string `json:"token_type"`
- ExpiresIn int64 `json:"expires_in"`
- Error string `json:"error"`
- ErrorDescription string `json:"error_description"`
- }
- if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
- return nil, fmt.Errorf("token endpoint returned unreadable JSON (status %d)", resp.StatusCode)
- }
- if body.Error != "" {
- if body.ErrorDescription != "" {
- return nil, fmt.Errorf("%s: %s", body.Error, body.ErrorDescription)
- }
- return nil, errors.New(body.Error)
- }
- if resp.StatusCode != http.StatusOK || body.AccessToken == "" {
- return nil, fmt.Errorf("token endpoint returned status %d", resp.StatusCode)
- }
- out := &TokenResponse{
- AccessToken: body.AccessToken,
- RefreshToken: body.RefreshToken,
- TokenType: body.TokenType,
- }
- if body.ExpiresIn > 0 {
- out.Expiry = time.Now().Add(time.Duration(body.ExpiresIn) * time.Second)
- }
- return out, nil
-}
-
-func writeError(w http.ResponseWriter, status int, msg string) {
- w.Header().Set("content-type", "application/json")
- w.WriteHeader(status)
- _ = json.NewEncoder(w).Encode(map[string]string{"error": msg})
}
diff --git a/internal/collector/collector_test.go b/internal/collector/collector_test.go
index 170f1d1..b0c265f 100644
--- a/internal/collector/collector_test.go
+++ b/internal/collector/collector_test.go
@@ -10,14 +10,14 @@ import (
func TestRegisterNotFoundWhenUnconfigured(t *testing.T) {
w := httptest.NewRecorder()
- New(Config{}, DeviceConfig{}, "/auth/collector").Register(w, httptest.NewRequest(http.MethodGet, "/auth/collector", nil))
+ New(Config{}).Register(w, httptest.NewRequest(http.MethodGet, "/auth/collector", nil))
if w.Code != http.StatusNotFound {
t.Fatalf("want 404, got %d", w.Code)
}
}
func TestRegisterPublishesIdentity(t *testing.T) {
- h := New(Config{ClientID: "ikc_abc", ClientSecret: "iks_xyz", TokenEndpoint: "https://as/token"}, DeviceConfig{}, "/auth/collector")
+ h := New(Config{ClientID: "ikc_abc"})
w := httptest.NewRecorder()
h.Register(w, httptest.NewRequest(http.MethodGet, "/auth/collector", nil))
if w.Code != http.StatusOK {
@@ -30,158 +30,18 @@ func TestRegisterPublishesIdentity(t *testing.T) {
if reg.ClientID != "ikc_abc" {
t.Errorf("clientId = %q", reg.ClientID)
}
- if !reg.ServerExchange {
- t.Error("a confidential client must ask for a server-side exchange")
- }
- if len(reg.RedirectURIs) == 0 {
- t.Error("no redirect URIs published")
- }
- // The secret must never appear in the published document.
- if strings.Contains(w.Body.String(), "iks_xyz") {
- t.Fatal("client secret leaked into the registration document")
+ if len(reg.Scopes) == 0 {
+ t.Error("no scopes published")
}
}
-func TestPublicClientDoesNotRequestServerExchange(t *testing.T) {
- h := New(Config{ClientID: "https://app.example.com/"}, DeviceConfig{}, "/auth/collector")
+func TestRegisterPublishesConfiguredScopes(t *testing.T) {
+ h := New(Config{ClientID: "ikc_abc", Scopes: []string{"profile", "email"}})
w := httptest.NewRecorder()
h.Register(w, httptest.NewRequest(http.MethodGet, "/auth/collector", nil))
var reg Registration
_ = json.Unmarshal(w.Body.Bytes(), ®)
- if reg.ServerExchange {
- t.Fatal("a public client should exchange its own code")
- }
-}
-
-// The exchange endpoint holds the client secret, so it must only ever redirect
-// to a callback it published. Otherwise a stolen code plus an attacker-chosen
-// redirect turns the server's credentials into an exchange oracle.
-func TestValidateRedirectRejectsForeignTargets(t *testing.T) {
- h := New(Config{ClientID: "ikc_abc", ClientSecret: "s", Ports: []int{40714}}, DeviceConfig{}, "/auth/collector")
- bad := []string{
- "",
- "https://evil.example.com/callback",
- "http://localhost:9999/callback",
- "http://evil.example.com:40714/callback",
- "http://localhost/callback",
- }
- for _, r := range bad {
- if err := h.validateRedirect(r); err == nil {
- t.Errorf("validateRedirect(%q) = nil, want an error", r)
- }
- }
- good := []string{
- "http://localhost:40714/callback",
- "http://127.0.0.1:40714/callback",
- }
- for _, r := range good {
- if err := h.validateRedirect(r); err != nil {
- t.Errorf("validateRedirect(%q) = %v, want nil", r, err)
- }
- }
-}
-
-func TestExchangeRejectsMissingFields(t *testing.T) {
- h := New(Config{ClientID: "ikc_abc", ClientSecret: "s", TokenEndpoint: "https://as/token"}, DeviceConfig{}, "/auth/collector")
- w := httptest.NewRecorder()
- body := strings.NewReader(`{"code":"abc"}`) // no verifier
- h.Exchange(w, httptest.NewRequest(http.MethodPost, "/auth/collector/exchange", body))
- if w.Code != http.StatusBadRequest {
- t.Fatalf("want 400, got %d", w.Code)
- }
-}
-
-func TestExchangeDisabledForPublicClient(t *testing.T) {
- h := New(Config{ClientID: "https://app.example.com/"}, DeviceConfig{}, "/auth/collector")
- w := httptest.NewRecorder()
- h.Exchange(w, httptest.NewRequest(http.MethodPost, "/auth/collector/exchange",
- strings.NewReader(`{"code":"a","codeVerifier":"b","redirectUri":"http://localhost:40714/callback"}`)))
- if w.Code != http.StatusNotFound {
- t.Fatalf("want 404, got %d", w.Code)
- }
-}
-
-// The exchange must send the secret to the authorization server, and pass the
-// collector's PKCE verifier through unchanged.
-func TestExchangeSendsSecretAndVerifier(t *testing.T) {
- var got map[string]string
- as := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- _ = r.ParseForm()
- got = map[string]string{}
- for k := range r.PostForm {
- got[k] = r.PostForm.Get(k)
- }
- w.Header().Set("content-type", "application/json")
- w.Write([]byte(`{"access_token":"AT","refresh_token":"RT","expires_in":3600}`))
- }))
- defer as.Close()
-
- h := New(Config{ClientID: "ikc_abc", ClientSecret: "iks_xyz", TokenEndpoint: as.URL, Ports: []int{40714}}, DeviceConfig{}, "/auth/collector")
- w := httptest.NewRecorder()
- h.Exchange(w, httptest.NewRequest(http.MethodPost, "/auth/collector/exchange",
- strings.NewReader(`{"code":"CODE","codeVerifier":"VERIFIER","redirectUri":"http://localhost:40714/callback"}`)))
- if w.Code != http.StatusOK {
- t.Fatalf("want 200, got %d: %s", w.Code, w.Body.String())
- }
- if got["client_secret"] != "iks_xyz" {
- t.Errorf("client_secret = %q", got["client_secret"])
- }
- if got["code_verifier"] != "VERIFIER" {
- t.Errorf("code_verifier = %q", got["code_verifier"])
- }
- if got["client_id"] != "ikc_abc" {
- t.Errorf("client_id = %q", got["client_id"])
- }
- var out TokenResponse
- _ = json.Unmarshal(w.Body.Bytes(), &out)
- if out.AccessToken != "AT" || out.RefreshToken != "RT" {
- t.Errorf("token = %+v", out)
- }
- if out.Expiry.IsZero() {
- t.Error("expiry not derived from expires_in")
- }
-}
-
-// An authorization-server error must reach the user as a message, not a
-// generic failure.
-func TestExchangeSurfacesUpstreamError(t *testing.T) {
- as := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- w.Header().Set("content-type", "application/json")
- w.WriteHeader(http.StatusBadRequest)
- w.Write([]byte(`{"error":"invalid_grant","error_description":"Authorization code not found"}`))
- }))
- defer as.Close()
-
- h := New(Config{ClientID: "ikc_abc", ClientSecret: "s", TokenEndpoint: as.URL, Ports: []int{40714}}, DeviceConfig{}, "/auth/collector")
- w := httptest.NewRecorder()
- h.Exchange(w, httptest.NewRequest(http.MethodPost, "/auth/collector/exchange",
- strings.NewReader(`{"code":"bad","codeVerifier":"v","redirectUri":"http://localhost:40714/callback"}`)))
- if !strings.Contains(w.Body.String(), "Authorization code not found") {
- t.Fatalf("upstream error not surfaced: %s", w.Body.String())
- }
-}
-
-func TestRefreshUsesServerCredentials(t *testing.T) {
- var got map[string]string
- as := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
- _ = r.ParseForm()
- got = map[string]string{}
- for k := range r.PostForm {
- got[k] = r.PostForm.Get(k)
- }
- w.Header().Set("content-type", "application/json")
- w.Write([]byte(`{"access_token":"AT2","expires_in":3600}`))
- }))
- defer as.Close()
-
- h := New(Config{ClientID: "ikc_abc", ClientSecret: "iks_xyz", TokenEndpoint: as.URL}, DeviceConfig{}, "/auth/collector")
- w := httptest.NewRecorder()
- h.Refresh(w, httptest.NewRequest(http.MethodPost, "/auth/collector/refresh",
- strings.NewReader(`{"refreshToken":"RT"}`)))
- if w.Code != http.StatusOK {
- t.Fatalf("want 200, got %d: %s", w.Code, w.Body.String())
- }
- if got["grant_type"] != "refresh_token" || got["client_secret"] != "iks_xyz" {
- t.Errorf("form = %v", got)
+ if strings.Join(reg.Scopes, " ") != "profile email" {
+ t.Errorf("scopes = %v", reg.Scopes)
}
}
diff --git a/internal/collector/device.go b/internal/collector/device.go
deleted file mode 100644
index 622eaaf..0000000
--- a/internal/collector/device.go
+++ /dev/null
@@ -1,217 +0,0 @@
-package collector
-
-import (
- "crypto/rand"
- "encoding/base32"
- "strings"
- "sync"
- "time"
-
- "golang.org/x/oauth2"
-)
-
-// Device-flow timings, following RFC 8628's guidance.
-const (
- // DeviceCodeTTL is how long the user has to finish authorizing.
- DeviceCodeTTL = 10 * time.Minute
- // DevicePollInterval is the minimum gap between polls the client is told
- // to respect.
- DevicePollInterval = 2 * time.Second
- // devicePollFloor is enforced server-side; a client polling faster than
- // this gets slow_down rather than an answer.
- devicePollFloor = 1 * time.Second
-)
-
-// Device-flow poll errors, using RFC 8628's names so a client can branch on
-// them without parsing prose.
-const (
- ErrAuthorizationPending = "authorization_pending"
- ErrSlowDown = "slow_down"
- ErrExpiredToken = "expired_token"
- ErrAccessDenied = "access_denied"
- ErrInvalidGrant = "invalid_grant"
-)
-
-// deviceSession is one in-flight authorization.
-//
-// It is deliberately memory-only: a pending login is worth less than the ten
-// minutes it lives for, and persisting half-finished credentials to disk buys
-// nothing but a place for them to leak.
-type deviceSession struct {
- DeviceCode string
- UserCode string
- // State is the OAuth state parameter, which is how the callback finds its
- // way back to this session.
- State string
- Verifier string
- Expires time.Time
-
- // Filled in once the user finishes.
- Token *oauth2.Token
- Err string
-
- lastPoll time.Time
-}
-
-// deviceStore holds pending device authorizations.
-type deviceStore struct {
- mu sync.Mutex
- byDevice map[string]*deviceSession
- byState map[string]*deviceSession
- byUser map[string]*deviceSession
-}
-
-func newDeviceStore() *deviceStore {
- return &deviceStore{
- byDevice: map[string]*deviceSession{},
- byState: map[string]*deviceSession{},
- byUser: map[string]*deviceSession{},
- }
-}
-
-// create mints a new session with fresh codes.
-func (d *deviceStore) create(verifier string) (*deviceSession, error) {
- deviceCode, err := randomCode(32)
- if err != nil {
- return nil, err
- }
- state, err := randomCode(16)
- if err != nil {
- return nil, err
- }
- userCode, err := randomUserCode()
- if err != nil {
- return nil, err
- }
- s := &deviceSession{
- DeviceCode: deviceCode,
- UserCode: userCode,
- State: state,
- Verifier: verifier,
- Expires: time.Now().Add(DeviceCodeTTL),
- }
- d.mu.Lock()
- defer d.mu.Unlock()
- d.sweepLocked()
- d.byDevice[deviceCode] = s
- d.byState[state] = s
- d.byUser[userCode] = s
- return s, nil
-}
-
-func (d *deviceStore) byStateCode(state string) (*deviceSession, bool) {
- d.mu.Lock()
- defer d.mu.Unlock()
- s, ok := d.byState[state]
- if !ok || time.Now().After(s.Expires) {
- return nil, false
- }
- return s, true
-}
-
-func (d *deviceStore) byUserCode(code string) (*deviceSession, bool) {
- d.mu.Lock()
- defer d.mu.Unlock()
- s, ok := d.byUser[NormalizeUserCode(code)]
- if !ok || time.Now().After(s.Expires) {
- return nil, false
- }
- return s, true
-}
-
-// complete records the outcome of an authorization.
-func (d *deviceStore) complete(s *deviceSession, tok *oauth2.Token, errMsg string) {
- d.mu.Lock()
- defer d.mu.Unlock()
- s.Token = tok
- s.Err = errMsg
-}
-
-// poll returns the session's outcome, enforcing the poll interval and
-// consuming the session once a token is handed over. Single use: a device code
-// that keeps returning a token is a device code worth stealing.
-func (d *deviceStore) poll(deviceCode string) (*oauth2.Token, string) {
- d.mu.Lock()
- defer d.mu.Unlock()
- s, ok := d.byDevice[deviceCode]
- if !ok {
- return nil, ErrInvalidGrant
- }
- if time.Now().After(s.Expires) {
- d.forgetLocked(s)
- return nil, ErrExpiredToken
- }
- if !s.lastPoll.IsZero() && time.Since(s.lastPoll) < devicePollFloor {
- return nil, ErrSlowDown
- }
- s.lastPoll = time.Now()
- switch {
- case s.Token != nil:
- tok := s.Token
- d.forgetLocked(s)
- return tok, ""
- case s.Err != "":
- errMsg := s.Err
- d.forgetLocked(s)
- return nil, errMsg
- default:
- return nil, ErrAuthorizationPending
- }
-}
-
-func (d *deviceStore) forgetLocked(s *deviceSession) {
- delete(d.byDevice, s.DeviceCode)
- delete(d.byState, s.State)
- delete(d.byUser, s.UserCode)
-}
-
-// sweepLocked drops expired sessions so an abandoned login cannot accumulate.
-func (d *deviceStore) sweepLocked() {
- now := time.Now()
- for _, s := range d.byDevice {
- if now.After(s.Expires) {
- d.forgetLocked(s)
- }
- }
-}
-
-// randomCode returns a URL-safe random string of n bytes of entropy.
-func randomCode(n int) (string, error) {
- b := make([]byte, n)
- if _, err := rand.Read(b); err != nil {
- return "", err
- }
- return strings.ToLower(base32.StdEncoding.WithPadding(base32.NoPadding).EncodeToString(b)), nil
-}
-
-// userCodeAlphabet omits characters that are easy to misread aloud or by eye
-// (0/O, 1/I/L, 2/Z, 5/S, 8/B), since this code gets typed by a human.
-const userCodeAlphabet = "ACDEFGHJKMNPQRTUVWXY34679"
-
-// randomUserCode returns a short human-typable code, formatted XXXX-XXXX.
-func randomUserCode() (string, error) {
- b := make([]byte, 8)
- if _, err := rand.Read(b); err != nil {
- return "", err
- }
- out := make([]byte, 0, 9)
- for i, v := range b {
- if i == 4 {
- out = append(out, '-')
- }
- out = append(out, userCodeAlphabet[int(v)%len(userCodeAlphabet)])
- }
- return string(out), nil
-}
-
-// NormalizeUserCode makes a typed code comparable: upper case, no spaces, and
-// dashes optional, because nobody types the dash reliably.
-func NormalizeUserCode(s string) string {
- s = strings.ToUpper(strings.TrimSpace(s))
- s = strings.ReplaceAll(s, " ", "")
- s = strings.ReplaceAll(s, "-", "")
- if len(s) == 8 {
- return s[:4] + "-" + s[4:]
- }
- return s
-}
diff --git a/internal/collector/device_test.go b/internal/collector/device_test.go
deleted file mode 100644
index febd0b7..0000000
--- a/internal/collector/device_test.go
+++ /dev/null
@@ -1,313 +0,0 @@
-package collector
-
-import (
- "encoding/json"
- "net/http"
- "net/http/httptest"
- "net/url"
- "strings"
- "testing"
- "time"
-
- "golang.org/x/oauth2"
-)
-
-func deviceHandler(t *testing.T, tokenEndpoint string) *Handler {
- t.Helper()
- return New(
- Config{ClientID: "ikc_abc", ClientSecret: "iks_xyz", TokenEndpoint: tokenEndpoint},
- DeviceConfig{PublicURL: "https://lard.example.com", AuthorizationEndpoint: "https://as.example.com/auth/authorize"},
- "/auth/collector",
- )
-}
-
-func startSession(t *testing.T, h *Handler) deviceStartResponse {
- t.Helper()
- w := httptest.NewRecorder()
- h.StartDevice(w, httptest.NewRequest(http.MethodPost, "/auth/collector/device", strings.NewReader("{}")))
- if w.Code != http.StatusOK {
- t.Fatalf("StartDevice: want 200, got %d: %s", w.Code, w.Body.String())
- }
- var out deviceStartResponse
- if err := json.Unmarshal(w.Body.Bytes(), &out); err != nil {
- t.Fatal(err)
- }
- return out
-}
-
-func TestDeviceUnavailableWithoutPublicURL(t *testing.T) {
- h := New(Config{ClientID: "x"}, DeviceConfig{AuthorizationEndpoint: "https://as/a"}, "/auth/collector")
- if h.DeviceAvailable() {
- t.Fatal("brokered login needs a public URL to receive the redirect")
- }
- w := httptest.NewRecorder()
- h.StartDevice(w, httptest.NewRequest(http.MethodPost, "/auth/collector/device", strings.NewReader("{}")))
- if w.Code != http.StatusNotFound {
- t.Fatalf("want 404, got %d", w.Code)
- }
-}
-
-func TestStartDeviceReturnsCodesAndURLs(t *testing.T) {
- out := startSession(t, deviceHandler(t, "https://as/token"))
- if out.DeviceCode == "" || out.UserCode == "" {
- t.Fatalf("missing codes: %+v", out)
- }
- // The verification URL must be on the server, not on the collector: that
- // is what makes it reachable from another machine.
- if !strings.HasPrefix(out.VerificationURI, "https://lard.example.com/auth/collector/device/verify") {
- t.Errorf("verificationUri = %q", out.VerificationURI)
- }
- if !strings.Contains(out.VerificationURIComplete, url.QueryEscape(out.UserCode)) {
- t.Errorf("complete URI should carry the code: %q", out.VerificationURIComplete)
- }
- if out.Interval <= 0 || out.ExpiresIn <= 0 {
- t.Errorf("interval/expiry not set: %+v", out)
- }
-}
-
-// The user code is typed by a human, so it must avoid glyphs that misread.
-func TestUserCodeAvoidsAmbiguousCharacters(t *testing.T) {
- for range 200 {
- code, err := randomUserCode()
- if err != nil {
- t.Fatal(err)
- }
- if len(code) != 9 || code[4] != '-' {
- t.Fatalf("unexpected shape: %q", code)
- }
- if strings.ContainsAny(code, "OIL01258BSZ") {
- t.Fatalf("ambiguous character in %q", code)
- }
- }
-}
-
-func TestNormalizeUserCodeIsForgiving(t *testing.T) {
- want := "ACDE-FGHJ"
- for _, in := range []string{"ACDE-FGHJ", "acde-fghj", "ACDEFGHJ", " acdefghj ", "acde fghj"} {
- if got := NormalizeUserCode(in); got != want {
- t.Errorf("NormalizeUserCode(%q) = %q, want %q", in, got, want)
- }
- }
-}
-
-// Visiting the verification URL must send the user to the authorization server
-// with PKCE and a redirect back to this server.
-func TestVerifyRedirectsToAuthorizationServer(t *testing.T) {
- h := deviceHandler(t, "https://as/token")
- out := startSession(t, h)
-
- w := httptest.NewRecorder()
- h.Verify(w, httptest.NewRequest(http.MethodGet,
- "/auth/collector/device/verify?user_code="+url.QueryEscape(out.UserCode), nil))
- if w.Code != http.StatusFound {
- t.Fatalf("want 302, got %d: %s", w.Code, w.Body.String())
- }
- loc, err := url.Parse(w.Header().Get("Location"))
- if err != nil {
- t.Fatal(err)
- }
- q := loc.Query()
- if q.Get("client_id") != "ikc_abc" {
- t.Errorf("client_id = %q", q.Get("client_id"))
- }
- if q.Get("code_challenge") == "" || q.Get("code_challenge_method") != "S256" {
- t.Error("PKCE challenge missing")
- }
- if q.Get("redirect_uri") != "https://lard.example.com/auth/collector/device/callback" {
- t.Errorf("redirect_uri = %q", q.Get("redirect_uri"))
- }
- if q.Get("state") == "" {
- t.Error("state missing; the callback could not find its session")
- }
-}
-
-func TestVerifyRejectsUnknownCode(t *testing.T) {
- h := deviceHandler(t, "https://as/token")
- w := httptest.NewRecorder()
- h.Verify(w, httptest.NewRequest(http.MethodGet, "/auth/collector/device/verify?user_code=ZZZZ-ZZZZ", nil))
- if w.Code != http.StatusNotFound {
- t.Fatalf("want 404, got %d", w.Code)
- }
-}
-
-func TestVerifyShowsFormWithoutCode(t *testing.T) {
- h := deviceHandler(t, "https://as/token")
- w := httptest.NewRecorder()
- h.Verify(w, httptest.NewRequest(http.MethodGet, "/auth/collector/device/verify", nil))
- if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), "
"+userCodeForm(""))
- return
- }
- http.Redirect(w, r, h.authorizeURL(sess), http.StatusFound)
-}
-
-// authorizeURL builds the authorization request, with the redirect pointing at
-// this server.
-func (h *Handler) authorizeURL(sess *deviceSession) string {
- q := url.Values{
- "response_type": {"code"},
- "client_id": {h.cfg.ClientID},
- "redirect_uri": {h.callbackURI()},
- "state": {sess.State},
- "code_challenge": {oauth2.S256ChallengeFromVerifier(sess.Verifier)},
- "code_challenge_method": {"S256"},
- "scope": {strings.Join(h.cfg.scopes(), " ")},
- "access_type": {"offline"},
- }
- sep := "?"
- if strings.Contains(h.dev.AuthorizationEndpoint, "?") {
- sep = "&"
- }
- return h.dev.AuthorizationEndpoint + sep + q.Encode()
-}
-
-// callbackURI is the redirect this server registers with the authorization
-// server. Register exactly this with your provider.
-func (h *Handler) callbackURI() string {
- return strings.TrimRight(h.dev.PublicURL, "/") + h.prefix + PathCallback
-}
-
-// CallbackURI exposes the redirect URI so it can be logged at boot: it is the
-// one value an operator must register with the authorization server.
-func (h *Handler) CallbackURI() string {
- if !h.DeviceAvailable() {
- return ""
- }
- return h.callbackURI()
-}
-
-// Callback receives the authorization server's redirect, exchanges the code
-// using this server's credentials, and parks the token for the polling client.
-func (h *Handler) Callback(w http.ResponseWriter, r *http.Request) {
- if !h.DeviceAvailable() {
- http.Error(w, "brokered device login is not configured", http.StatusNotFound)
- return
- }
- q := r.URL.Query()
- sess, ok := h.devices.byStateCode(q.Get("state"))
- if !ok {
- // No session for this state: either it expired or the state was forged.
- // Either way there is nothing to complete.
- h.devicePage(w, http.StatusBadRequest, "This login has expired",
- "Start again with lard-client login.
")
- return
- }
- if e := q.Get("error"); e != "" {
- h.devices.complete(sess, nil, ErrAccessDenied)
- h.devicePage(w, http.StatusOK, "Authorization declined",
- "Nothing was connected. You can close this tab.
")
- return
- }
- code := q.Get("code")
- if code == "" {
- h.devices.complete(sess, nil, ErrInvalidGrant)
- h.devicePage(w, http.StatusBadRequest, "No authorization code",
- "The provider did not return a code. Try again.
")
- return
- }
- tok, err := h.exchangeDevice(r.Context(), code, sess.Verifier)
- if err != nil {
- slog.Warn("device login: exchange failed", "error", err)
- h.devices.complete(sess, nil, err.Error())
- h.devicePage(w, http.StatusBadGateway, "Could not finish authorizing",
- ""+html.EscapeString(err.Error())+"
")
- return
- }
- h.devices.complete(sess, tok, "")
- slog.Info("device login completed", "user_code", sess.UserCode)
- h.devicePage(w, http.StatusOK, "Connected",
- "This machine is now connected to lard. You can close this tab and return to your terminal.
")
-}
-
-// exchangeDevice trades the code for a token. The client secret is used when
-// there is one, which is why this happens here rather than on the collector.
-func (h *Handler) exchangeDevice(ctx context.Context, code, verifier string) (*oauth2.Token, error) {
- ctx, cancel := context.WithTimeout(ctx, 20*time.Second)
- defer cancel()
- form := url.Values{
- "grant_type": {"authorization_code"},
- "code": {code},
- "client_id": {h.cfg.ClientID},
- "redirect_uri": {h.callbackURI()},
- "code_verifier": {verifier},
- }
- if h.cfg.ClientSecret != "" {
- form.Set("client_secret", h.cfg.ClientSecret)
- }
- res, err := postToken(ctx, h.cfg.TokenEndpoint, form)
- if err != nil {
- return nil, err
- }
- return &oauth2.Token{
- AccessToken: res.AccessToken,
- RefreshToken: res.RefreshToken,
- Expiry: res.Expiry,
- TokenType: "Bearer",
- }, nil
-}
-
-// PollDevice is the client's polling endpoint. It answers with the token once
-// the user finishes, and with RFC 8628 error codes until then.
-func (h *Handler) PollDevice(w http.ResponseWriter, r *http.Request) {
- if !h.DeviceAvailable() {
- writeError(w, http.StatusNotFound, "brokered device login is not configured")
- return
- }
- var req struct {
- DeviceCode string `json:"deviceCode"`
- }
- if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 1<<16)).Decode(&req); err != nil {
- writeError(w, http.StatusBadRequest, "malformed request")
- return
- }
- if req.DeviceCode == "" {
- writeError(w, http.StatusBadRequest, "deviceCode is required")
- return
- }
- tok, errCode := h.devices.poll(req.DeviceCode)
- w.Header().Set("content-type", "application/json")
- w.Header().Set("cache-control", "no-store")
- if errCode != "" {
- // Pending and slow_down are normal progress, not failures, so they get
- // 200 with an error code the client understands. Anything else is a
- // real 400.
- status := http.StatusBadRequest
- if errCode == ErrAuthorizationPending || errCode == ErrSlowDown {
- status = http.StatusOK
- }
- w.WriteHeader(status)
- _ = json.NewEncoder(w).Encode(map[string]string{"error": errCode})
- return
- }
- _ = json.NewEncoder(w).Encode(TokenResponse{
- AccessToken: tok.AccessToken,
- RefreshToken: tok.RefreshToken,
- Expiry: tok.Expiry,
- TokenType: "Bearer",
- })
-}
-
-// devicePage renders a minimal styled page for the browser side of the flow.
-func (h *Handler) devicePage(w http.ResponseWriter, status int, title, body string) {
- w.Header().Set("content-type", "text/html; charset=utf-8")
- w.WriteHeader(status)
- fmt.Fprintf(w, `%s · lard
-
-%s
%s`,
- html.EscapeString(title), html.EscapeString(title), body)
-}
-
-// userCodeForm is the fallback for a user who read their code off another
-// screen rather than following a link.
-func userCodeForm(prefill string) string {
- return ``
-}
diff --git a/internal/setup/setup.go b/internal/setup/setup.go
index 1241309..94bca96 100644
--- a/internal/setup/setup.go
+++ b/internal/setup/setup.go
@@ -33,11 +33,12 @@ func Interactive() bool {
// Options carry anything already supplied on the command line, so the form
// only asks for what is genuinely missing.
type Options struct {
- URL string
- Token string
- Roots []string
- CallbackPort int
- NoBrowser bool
+ URL string
+ Token string
+ Roots []string
+ NoBrowser bool
+ // Force re-authenticates even if the saved credentials still verify.
+ Force bool
}
// Run resolves a working configuration, prompting where needed, and saves it.
@@ -81,6 +82,39 @@ func Run(ctx context.Context, opts Options) (*client.Config, error) {
return cfg, nil
}
+// Logout forgets this machine's credentials. It first tells the authorization
+// server to revoke the refresh token (RFC 7009), so the grant stops working
+// anywhere rather than just being deleted locally; a failure there is
+// reported but does not stop the local wipe, since an unreachable server
+// should not strand credentials on disk. The config file keeps the server URL
+// and roots so a later `login` only has to re-authenticate.
+func Logout(ctx context.Context) error {
+ path := client.DefaultConfigPath()
+ cfg, err := client.LoadConfig(path)
+ if err != nil {
+ return err
+ }
+ revoked := false
+ if cfg.OAuth != nil && cfg.OAuth.RefreshToken != "" {
+ if err := client.RevokeToken(ctx, cfg.URL, cfg.OAuth.RefreshToken); err != nil {
+ fmt.Fprintf(os.Stderr, "warning: could not revoke the refresh token at the server: %v\n", err)
+ } else {
+ revoked = true
+ }
+ }
+ cfg.OAuth = nil
+ cfg.Token = ""
+ if err := cfg.Save(path); err != nil {
+ return err
+ }
+ if revoked {
+ fmt.Println("Revoked the refresh token and removed local credentials.")
+ } else {
+ fmt.Println("Removed local credentials.")
+ }
+ return nil
+}
+
// needsURL reports whether the URL is still unset in any meaningful sense.
func needsURL(path string) bool {
if _, err := os.Stat(path); err == nil {
@@ -166,42 +200,36 @@ func authenticate(ctx context.Context, cfg *client.Config, opts Options) error {
cfg.OAuth = nil
return nil
}
- // Already have something that works? Don't make the user re-authorize.
- if cfg.AuthMode() != "none" {
+ // Already have something that works? Don't make the user re-authorize,
+ // unless they asked for fresh credentials.
+ if !opts.Force && cfg.AuthMode() != "none" {
if _, err := cfg.Verify(ctx); err == nil {
return nil
}
}
+ // Forcing a re-login rotates the grant: kill the old refresh token at the
+ // server so only the new one works. Best-effort — a dead server shouldn't
+ // block a fresh login.
+ if opts.Force && cfg.OAuth != nil && cfg.OAuth.RefreshToken != "" {
+ if err := client.RevokeToken(ctx, cfg.URL, cfg.OAuth.RefreshToken); err != nil {
+ fmt.Fprintf(os.Stderr, "warning: could not revoke the old refresh token: %v\n", err)
+ }
+ }
// Does the server need credentials at all?
if _, err := cfg.Verify(ctx); err == nil {
return nil
}
- port := opts.CallbackPort
- if port <= 0 {
- port = client.CallbackPort
- }
- // Prefer the brokered flow when the server offers it: it needs no local
- // callback port, so it works identically over SSH, in a container, and on
- // a headless machine. Fall back to the local-listener flow otherwise.
- reg, regErr := client.FetchRegistration(ctx, cfg.URL)
- _, discErr := client.Discover(ctx, cfg.URL)
- if regErr == nil && reg.DeviceFlow {
- tok, err := client.LoginDevice(ctx, cfg.URL, !opts.NoBrowser)
- if err != nil {
- return err
+ // The device grant is the only login flow: it needs no callback listener,
+ // no browser on this machine, and no client secret, so it works
+ // identically on a laptop, over SSH, in a container, and headless.
+ eps, discErr := client.Discover(ctx, cfg.URL)
+ if discErr == nil && eps.DeviceAuthorization != "" {
+ reg, regErr := client.FetchRegistration(ctx, cfg.URL)
+ if regErr != nil {
+ return fmt.Errorf("server publishes no collector registration; set LARD_COLLECTOR_CLIENT_ID there")
}
- cfg.OAuth = &client.OAuthToken{
- AccessToken: tok.AccessToken,
- RefreshToken: tok.RefreshToken,
- Expiry: tok.Expiry,
- CallbackPort: 0, // brokered: no local callback involved
- }
- cfg.Token = ""
- return nil
- }
- if discErr == nil {
- tok, err := client.Login(ctx, cfg.URL, port, !opts.NoBrowser)
+ tok, err := client.LoginDevice(ctx, cfg.URL, reg.ClientID, reg.Scopes, !opts.NoBrowser)
if err != nil {
return err
}
@@ -209,12 +237,9 @@ func authenticate(ctx context.Context, cfg *client.Config, opts Options) error {
AccessToken: tok.AccessToken,
RefreshToken: tok.RefreshToken,
Expiry: tok.Expiry,
- CallbackPort: port,
+ ClientID: reg.ClientID,
}
cfg.Token = ""
- if tok.RefreshToken == "" {
- fmt.Fprintln(os.Stderr, "note: no refresh token issued; you will need to log in again when this expires")
- }
return nil
}
--
2.51.2