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(), "It may have expired, or already been used. Start again with "+ - "lard-client login.

"+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