diff --git a/README.md b/README.md index e65213d..89619a1 100644 --- a/README.md +++ b/README.md @@ -11,16 +11,48 @@ The canonical repo for this is hosted on tangled over at [`dunkirk.sh/lard`](htt echo 'HYPER_API_KEY=sk-hyper-...' > .env go run ./cmd/lard # listens on :7477 -# client: backfill crush sessions -LARD_URL=http://localhost:7477 go run ./cmd/lard-client backfill --root ~/code +# client: point it at the server, then load everything you have +lard-client login # asks for the url, opens your browser +lard-client backfill --root ~/code -# update it -LARD_URL=http://localhost:7477 go run ./cmd/lard-client daemon --interval 5m +# keep it fed in the background (macOS) +lard-client service install +``` + +`login` asks where the server lives, then runs a browser flow against whatever +auth server lard names, so there is no token to copy. It always prints the raw +authorization URL, so you can paste it into a browser on another machine. On a +headless box pass `--url` and `--token` and it never prompts. + +when the server brokers the login (the default once `LARD_COLLECTOR_CLIENT_ID` +is set) the collector opens no ports at all. you get a url on the server and it +polls until you finish, so ssh, containers, and headless boxes all work the same +way with nothing to forward. + +after that nothing needs poking: the agent syncs on an interval, and the server +consolidates itself once uploads go quiet. -# run a consolidation pass -LARD_URL=http://localhost:7477 go run ./cmd/lard-client consolidate +## client + +``` +lard-client login [--url URL] [--token TOKEN] [--root DIR...] +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 +lard-client daemon [--interval 5m] # sync in a loop, for non-macOS init +lard-client service install|uninstall|status +lard-client consolidate # force a pass now ``` +config lives at `~/.config/lard/client.json` (mode 0600, holds the token). +`LARD_URL` and `LARD_TOKEN` override it. + +`service install` writes a launchd agent that syncs on an interval and survives +reboots. It refuses to install if it cannot reach the server, since a silent +background failure is the worst outcome. Logs go to +`~/Library/Logs/lard-client.log`. Linux is not wired up yet: run +`lard-client daemon` under systemd. + ## Interfaces MCP (for agents) at `POST /mcp`: `get_context`, `memory_list`, `memory_read`, `memory_write`, `memory_append`, `memory_delete`. add to crush: @@ -67,10 +99,19 @@ paths are `profile`, `areas/`, `topics/`, `people/`. | `LARD_OAUTH_CLIENT_IDS` | | comma list of client ids allowed to call lard | | `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_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 | | `LARD_MODEL` | `deepseek-v4-flash` | consolidation model | -client env: `LARD_URL`, `LARD_TOKEN`. +client env: `LARD_URL`, `LARD_TOKEN` (both override `~/.config/lard/client.json`). + +consolidation is automatic: an ingest starts a quiet timer, and the pass runs +once uploads stop. bursts collapse into one pass, and a machine uploading +continuously still gets consolidated at `LARD_CONSOLIDATE_MAX_WAIT`. ## auth @@ -92,6 +133,55 @@ set `LARD_OAUTH_CLIENT_IDS` or `LARD_OAUTH_USERS`. indiko mints tokens for every app you sign into, so without an allowlist any one of them can read all your memory. lard warns at boot if you skip it. +### collector registration + +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`. + +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. + +### 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]): + +``` +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 +``` + +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 +``` + +that url is built from `LARD_PUBLIC_URL`, so set it to the server's external +address before registering. + +pending sessions live in memory for 10 minutes, device codes are single use, and +polling faster than once a second gets `slow_down`. + +[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 `POST /register` will not connect. use a client that accepts a configured client id. indiko also rejects a client id whose host differs from the redirect diff --git a/cmd/lard-client/main.go b/cmd/lard-client/main.go index 14d35e2..6cbe056 100644 --- a/cmd/lard-client/main.go +++ b/cmd/lard-client/main.go @@ -1,154 +1,378 @@ -// lard-client is the edge collector. Two modes: +// lard-client is the edge collector: it finds Crush session databases on this +// machine and uploads them to a central lard. // -// lard-client backfill --root ~/code --root ~/code/charm -// Scan for every .crush/crush.db under the roots and upload all -// sessions ever. Idempotent; safe to re-run. -// -// lard-client daemon [--workspace .] [--interval 5m] -// Periodically collect new/changed sessions and upload. Run per -// machine; with --workspace it watches a single repo, otherwise it -// rescans the roots each tick. +// First run asks where the server lives and walks a browser login, so setup is +// `lard-client login` with nothing else to look up. Every prompt has a flag +// equivalent, for scripts and headless machines. package main import ( "context" - "flag" "fmt" "log/slog" "os" - "os/signal" - "syscall" + "path/filepath" + "strings" "time" + "charm.land/fang/v2" + "github.com/spf13/cobra" + "github.com/taciturnaxolotl/lard/internal/client" "github.com/taciturnaxolotl/lard/internal/dotenv" + "github.com/taciturnaxolotl/lard/internal/service" + "github.com/taciturnaxolotl/lard/internal/setup" ) +// version is overridden at build time via -ldflags. +var version = "dev" + func main() { - if err := run(os.Args[1:]); err != nil { - fmt.Fprintln(os.Stderr, "lard-client:", err) + dotenv.LoadDefault() + if err := fang.Execute(context.Background(), rootCmd(), fang.WithVersion(version)); err != nil { os.Exit(1) } } -func run(args []string) error { - dotenv.LoadDefault() - if len(args) == 0 { - usage() - return fmt.Errorf("subcommand required: backfill | daemon | sync") - } - switch args[0] { - case "backfill": - return runSync(args[1:], true, false) - case "sync": - return runSync(args[1:], false, false) - case "daemon": - return runDaemon(args[1:]) - case "consolidate": - ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) - defer stop() - return uploader().Consolidate(ctx) - default: - usage() - return fmt.Errorf("unknown subcommand %q", args[0]) +func rootCmd() *cobra.Command { + root := &cobra.Command{ + Use: "lard-client", + Short: "Collect Crush sessions and send them to lard", + Long: `lard-client finds Crush session databases on this machine and uploads +them to a central lard server, which turns them into durable memory. + +Start with 'lard-client login', then 'lard-client backfill'.`, + SilenceUsage: true, } + root.AddCommand( + loginCmd(), + backfillCmd(), + syncCmd(), + daemonCmd(), + consolidateCmd(), + serviceCmd(), + statusCmd(), + ) + return root } -func usage() { - fmt.Fprintf(os.Stderr, `usage: - lard-client backfill --root DIR [--root DIR...] [--dry-run] - lard-client sync [--workspace DIR...] - lard-client daemon [--workspace DIR...] [--root DIR...] [--interval 5m] +// --- login --- + +func loginCmd() *cobra.Command { + var opts setup.Options + cmd := &cobra.Command{ + Use: "login", + Short: "Connect this machine to a lard server", + Long: `Point the collector at a server and get credentials. -env: - LARD_URL central service base URL (default http://localhost:7477) - LARD_TOKEN bearer token if the service requires auth -`) +With no flags this asks for the server URL, then opens your browser to +authorize. On a headless machine pass --url and --token instead.`, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + cfg, err := setup.Run(cmd.Context(), opts) + if err != nil { + return err + } + printConnected(cmd.Context(), cfg, opts.CallbackPort) + return nil + }, + } + f := cmd.Flags() + 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") + return cmd } -func uploader() *client.Uploader { - base := os.Getenv("LARD_URL") - if base == "" { - base = "http://localhost:7477" +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.") + } } - return client.NewUploader(base, os.Getenv("LARD_TOKEN")) + fmt.Printf("Saved %s\n\nNext: lard-client backfill\n", client.DefaultConfigPath()) } -type rootFlags []string +// --- collection --- + +func backfillCmd() *cobra.Command { + var roots []string + var dryRun bool + cmd := &cobra.Command{ + Use: "backfill", + Short: "Upload every session ever recorded", + Long: `Scan for Crush databases and upload all sessions, not just new ones. -func (r *rootFlags) String() string { return fmt.Sprint([]string(*r)) } -func (r *rootFlags) Set(v string) error { - *r = append(*r, v) - return nil +Idempotent: safe to run again, and safe to interrupt.`, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + return doSync(cmd.Context(), true, dryRun, configuredRoots(roots), nil) + }, + } + cmd.Flags().StringSliceVar(&roots, "root", nil, "directory to scan (repeatable; defaults to the saved config)") + cmd.Flags().BoolVar(&dryRun, "dry-run", false, "collect but do not upload") + return cmd } -func runSync(args []string, full, _ bool) error { - fs := flag.NewFlagSet("sync", flag.ContinueOnError) - var roots, workspaces rootFlags +func syncCmd() *cobra.Command { + var roots, workspaces []string var dryRun bool - fs.Var(&roots, "root", "root to scan for .crush dirs (repeatable)") - fs.Var(&workspaces, "workspace", "explicit workspace with .crush (repeatable)") - fs.BoolVar(&dryRun, "dry-run", false, "collect but do not upload") - if err := fs.Parse(args); err != nil { - return err + cmd := &cobra.Command{ + Use: "sync", + Short: "Upload sessions changed since the last run", + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + r := roots + if len(workspaces) == 0 { + r = configuredRoots(roots) + } + return doSync(cmd.Context(), false, dryRun, r, workspaces) + }, } - if len(roots) == 0 && len(workspaces) == 0 { - home, _ := os.UserHomeDir() - roots = rootFlags{home + "/code"} + f := cmd.Flags() + f.StringSliceVar(&roots, "root", nil, "directory to scan (repeatable)") + f.StringSliceVar(&workspaces, "workspace", nil, "single repo to sync (repeatable)") + f.BoolVar(&dryRun, "dry-run", false, "collect but do not upload") + return cmd +} + +func daemonCmd() *cobra.Command { + var roots, workspaces []string + var interval time.Duration + cmd := &cobra.Command{ + Use: "daemon", + Short: "Sync on an interval in the foreground", + Long: `Sync repeatedly until interrupted. + +On macOS prefer 'lard-client service install', which lets launchd own the +schedule. Use this under systemd or in a container.`, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + r := roots + if len(workspaces) == 0 { + r = configuredRoots(roots) + } + slog.Info("lard-client daemon", "interval", interval, "roots", r, "workspaces", workspaces) + tick := func() { + if err := doSync(cmd.Context(), false, false, r, workspaces); err != nil { + slog.Error("sync", "error", err) + } + } + tick() + t := time.NewTicker(interval) + defer t.Stop() + for { + select { + case <-cmd.Context().Done(): + return nil + case <-t.C: + tick() + } + } + }, } - return doSync(full, dryRun, roots, workspaces) + f := cmd.Flags() + f.StringSliceVar(&roots, "root", nil, "directory to scan (repeatable)") + f.StringSliceVar(&workspaces, "workspace", nil, "single repo to sync (repeatable)") + f.DurationVar(&interval, "interval", 5*time.Minute, "sync interval") + return cmd } -func doSync(full, dryRun bool, roots, workspaces []string) error { - ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) - defer stop() - st, err := client.LoadState(client.DefaultStatePath()) - if err != nil { - return err +func consolidateCmd() *cobra.Command { + return &cobra.Command{ + Use: "consolidate", + Short: "Ask the server to consolidate now", + Long: `Force a consolidation pass. + +The server normally does this on its own once uploads go quiet, so this is +only needed to skip the wait.`, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + up, err := uploader(cmd.Context()) + if err != nil { + return err + } + return up.Consolidate(cmd.Context()) + }, } - return client.Sync(ctx, uploader(), st, client.SyncOpts{ - Workspaces: workspaces, - Roots: roots, - Full: full, - Collector: client.Hostname(), - DryRun: dryRun, - }) } -func runDaemon(args []string) error { - fs := flag.NewFlagSet("daemon", flag.ContinueOnError) - var roots, workspaces rootFlags - var interval time.Duration - fs.Var(&roots, "root", "root to scan for .crush dirs (repeatable)") - fs.Var(&workspaces, "workspace", "explicit workspace with .crush (repeatable)") - fs.DurationVar(&interval, "interval", 5*time.Minute, "sync interval") - if err := fs.Parse(args); err != nil { - return err +// --- status --- + +func statusCmd() *cobra.Command { + return &cobra.Command{ + Use: "status", + Short: "Show the configured server and background agent", + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + cfg, err := client.LoadConfig(client.DefaultConfigPath()) + if err != nil { + return err + } + fmt.Printf("server: %s\n", cfg.URL) + fmt.Printf("auth: %s\n", cfg.AuthMode()) + fmt.Printf("roots: %s\n", strings.Join(configuredRoots(nil), ", ")) + if id, err := cfg.Verify(cmd.Context()); err != nil { + fmt.Printf("reach: %v\n", err) + } else if id != "" { + fmt.Printf("reach: ok, as %s\n", id) + } else { + fmt.Println("reach: ok") + } + if service.Supported() { + installed, loaded, detail, serr := service.Status() + switch { + case serr != nil: + fmt.Printf("agent: %v\n", serr) + case !installed: + fmt.Println("agent: not installed (lard-client service install)") + case !loaded: + fmt.Println("agent: installed but not loaded") + default: + fmt.Printf("agent: %s\n", detail) + } + } + return nil + }, } - if len(roots) == 0 && len(workspaces) == 0 { - home, _ := os.UserHomeDir() - roots = rootFlags{home + "/code"} +} + +// --- service --- + +func serviceCmd() *cobra.Command { + cmd := &cobra.Command{ + Use: "service", + Short: "Manage the background sync agent", } + var roots []string + var interval time.Duration + install := &cobra.Command{ + Use: "install", + Short: "Install and start the background agent", + Long: `Install a launchd agent that syncs on an interval and survives reboots. - ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) - defer stop() +Refuses to install if the server cannot be reached, since a background agent +that fails silently is worse than none.`, + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, _ []string) error { + cfg, err := client.LoadConfig(client.DefaultConfigPath()) + if err != nil { + return err + } + if _, err := cfg.Verify(cmd.Context()); err != nil { + return fmt.Errorf("%w\nrun 'lard-client login' first", err) + } + bin, err := os.Executable() + if err != nil { + return err + } + if bin, err = filepath.EvalSymlinks(bin); err != nil { + return err + } + r := configuredRoots(roots) + path, err := service.Install(service.Options{Binary: bin, Interval: interval, Roots: r}) + if err != nil { + return err + } + fmt.Printf("Installed %s\nSyncing every %s from %s\nLogs: ~/Library/Logs/lard-client.log\n", + path, interval, strings.Join(r, ", ")) + return nil + }, + } + install.Flags().StringSliceVar(&roots, "root", nil, "directory to scan (repeatable; defaults to the saved config)") + install.Flags().DurationVar(&interval, "interval", 5*time.Minute, "sync interval") - slog.Info("lard-client daemon", "interval", interval, "roots", roots, "workspaces", workspaces) - // Tick immediately, then on interval. - tick := func() { - if err := doSync(false, false, roots, workspaces); err != nil { - slog.Error("sync", "error", err) - } + uninstall := &cobra.Command{ + Use: "uninstall", + Short: "Stop and remove the background agent", + Args: cobra.NoArgs, + RunE: func(*cobra.Command, []string) error { + if err := service.Uninstall(); err != nil { + return err + } + fmt.Println("Uninstalled.") + return nil + }, } - tick() - t := time.NewTicker(interval) - defer t.Stop() - for { - select { - case <-ctx.Done(): + status := &cobra.Command{ + Use: "status", + Short: "Report whether the agent is running", + Args: cobra.NoArgs, + RunE: func(*cobra.Command, []string) error { + installed, loaded, detail, err := service.Status() + if err != nil { + return err + } + switch { + case !installed: + fmt.Println("Not installed; run 'lard-client service install'.") + case !loaded: + fmt.Println("Installed but not loaded; run 'lard-client service install' again.") + default: + fmt.Println(detail) + } return nil - case <-t.C: - tick() - } + }, + } + cmd.AddCommand(install, uninstall, status) + return cmd +} + +// --- shared --- + +// uploader builds an Uploader from the saved config, refreshing an expired +// OAuth token first so a background run never fails on a stale credential. +func uploader(ctx context.Context) (*client.Uploader, error) { + path := client.DefaultConfigPath() + cfg, err := client.LoadConfig(path) + if err != nil { + return nil, err + } + tok, err := cfg.Bearer(ctx, path) + if err != nil { + return nil, err + } + return client.NewUploader(cfg.URL, tok), nil +} + +// configuredRoots resolves which directories to scan: explicit flags, then the +// saved config, then ~/code. +func configuredRoots(flagRoots []string) []string { + if len(flagRoots) > 0 { + return flagRoots + } + if cfg, err := client.LoadConfig(client.DefaultConfigPath()); err == nil && len(cfg.Roots) > 0 { + return cfg.Roots } + home, _ := os.UserHomeDir() + return []string{filepath.Join(home, "code")} +} + +func doSync(ctx context.Context, full, dryRun bool, roots, workspaces []string) error { + st, err := client.LoadState(client.DefaultStatePath()) + if err != nil { + return err + } + up, err := uploader(ctx) + if err != nil { + return err + } + return client.Sync(ctx, up, st, client.SyncOpts{ + Workspaces: workspaces, + Roots: roots, + Full: full, + Collector: client.Hostname(), + DryRun: dryRun, + }) } diff --git a/cmd/lard/main.go b/cmd/lard/main.go index 31729ad..4a1800c 100644 --- a/cmd/lard/main.go +++ b/cmd/lard/main.go @@ -3,6 +3,7 @@ package main import ( "context" + "encoding/json" "errors" "fmt" "log/slog" @@ -10,11 +11,13 @@ import ( "os" "os/signal" "path/filepath" + "strconv" "strings" "syscall" "time" "github.com/taciturnaxolotl/lard/internal/auth" + "github.com/taciturnaxolotl/lard/internal/collector" "github.com/taciturnaxolotl/lard/internal/dotenv" "github.com/taciturnaxolotl/lard/internal/httpapi" "github.com/taciturnaxolotl/lard/internal/llm" @@ -58,16 +61,59 @@ func run() error { } api := httpapi.New(st, llmClient) + // Consolidate on its own once uploads go quiet, so a remote collector + // feeding this server keeps memory current with nobody poking an endpoint. + if after := envDuration("LARD_CONSOLIDATE_AFTER", httpapi.DefaultConsolidateAfter); after > 0 { + api.EnableAutoConsolidate(after, envDuration("LARD_CONSOLIDATE_MAX_WAIT", httpapi.DefaultConsolidateMaxWait)) + defer api.StopAutoConsolidate() + } mcpSrv := mcpserver.New(api) cfg := auth.Config{ - Mode: auth.Mode(envOr("LARD_AUTH", string(auth.ModeNone))), - Token: os.Getenv("LARD_TOKEN"), - IndikoURL: envOr("LARD_INDIKO_URL", "https://indiko.dunkirk.sh"), - PublicURL: os.Getenv("LARD_PUBLIC_URL"), - AllowedClientIDs: envList("LARD_OAUTH_CLIENT_IDS"), - AllowedUsers: envList("LARD_OAUTH_USERS"), - RequiredScopes: envList("LARD_OAUTH_SCOPES"), + Mode: auth.Mode(envOr("LARD_AUTH", string(auth.ModeNone))), + Token: os.Getenv("LARD_TOKEN"), + IndikoURL: envOr("LARD_INDIKO_URL", "https://indiko.dunkirk.sh"), + PublicURL: os.Getenv("LARD_PUBLIC_URL"), + AllowedClientIDs: envList("LARD_OAUTH_CLIENT_IDS"), + AllowedUsers: envList("LARD_OAUTH_USERS"), + RequiredScopes: envList("LARD_OAUTH_SCOPES"), + CollectorClientID: os.Getenv("LARD_COLLECTOR_CLIENT_ID"), + } + + // The collector registration: what identity edge collectors adopt, and + // whether this server exchanges their codes for them. + 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"), + } + // 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 + 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") + } } for _, warn := range cfg.Validate() { slog.Warn("auth: " + warn) @@ -81,6 +127,14 @@ func run() error { mux.Handle(auth.PathProtectedResource, auth.ProtectedResourceMetadata(cfg)) 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") @@ -126,6 +180,24 @@ func envList(key string) []string { return out } +// envDuration reads a Go duration (e.g. "5m"). "off" or "0" disables the +// feature by returning zero; an unparseable value falls back to the default. +func envDuration(key string, fallback time.Duration) time.Duration { + raw := strings.TrimSpace(os.Getenv(key)) + if raw == "" { + return fallback + } + if raw == "off" || raw == "never" || raw == "0" { + return 0 + } + d, err := time.ParseDuration(raw) + if err != nil { + slog.Warn("ignoring unparseable duration", "key", key, "value", raw) + return fallback + } + return d +} + func defaultDBPath() string { if d, err := os.UserConfigDir(); err == nil { return filepath.Join(d, "lard", "lard.db") @@ -139,3 +211,54 @@ 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/go.mod b/go.mod index 565a68f..8531295 100644 --- a/go.mod +++ b/go.mod @@ -3,34 +3,68 @@ module github.com/taciturnaxolotl/lard go 1.26.5 require ( + charm.land/fang/v2 v2.0.1 charm.land/fantasy v0.38.1 + charm.land/huh/v2 v2.0.3 github.com/google/uuid v1.6.0 github.com/modelcontextprotocol/go-sdk v1.6.1 + github.com/pkg/browser v0.0.0-20240102092130-5ac0b6a4141c + github.com/spf13/cobra v1.10.2 + golang.org/x/oauth2 v0.36.0 + golang.org/x/term v0.45.0 modernc.org/sqlite v1.54.0 ) require ( + charm.land/bubbles/v2 v2.0.0 // indirect + charm.land/bubbletea/v2 v2.0.2 // indirect + charm.land/lipgloss/v2 v2.0.1 // indirect + github.com/atotto/clipboard v0.1.4 // indirect + github.com/catppuccin/go v0.3.0 // indirect + github.com/charmbracelet/colorprofile v0.4.2 // indirect + github.com/charmbracelet/ultraviolet v0.0.0-20260205113103-524a6607adb8 // indirect + github.com/charmbracelet/x/ansi v0.11.6 // indirect + github.com/charmbracelet/x/exp/charmtone v0.0.0-20250603201427-c31516f43444 // indirect + github.com/charmbracelet/x/exp/ordered v0.1.0 // indirect github.com/charmbracelet/x/exp/slice v0.0.0-20250904123553-b4e2667e5ad5 // indirect + github.com/charmbracelet/x/exp/strings v0.1.0 // indirect + github.com/charmbracelet/x/term v0.2.2 // indirect + github.com/charmbracelet/x/termios v0.1.1 // indirect + github.com/charmbracelet/x/windows v0.2.2 // indirect + github.com/clipperhouse/displaywidth v0.11.0 // indirect + github.com/clipperhouse/uax29/v2 v2.7.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/go-json-experiment/json v0.0.0-20260623181947-01eb4420fa68 // indirect github.com/go-viper/mapstructure/v2 v2.5.0 // indirect github.com/goccy/go-yaml v1.19.2 // indirect github.com/google/jsonschema-go v0.4.3 // indirect + github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/kaptinlin/jsonpointer v0.4.27 // indirect github.com/kaptinlin/jsonschema v0.9.3 // indirect + github.com/lucasb-eyer/go-colorful v1.3.0 // indirect github.com/mattn/go-isatty v0.0.20 // indirect + github.com/mattn/go-runewidth v0.0.20 // indirect + github.com/mitchellh/hashstructure/v2 v2.0.2 // indirect + github.com/muesli/cancelreader v0.2.2 // indirect + github.com/muesli/mango v0.1.0 // indirect + github.com/muesli/mango-cobra v1.2.0 // indirect + github.com/muesli/mango-pflag v0.1.0 // indirect + github.com/muesli/roff v0.1.0 // indirect github.com/ncruces/go-strftime v1.0.0 // indirect github.com/openai/openai-go/v3 v3.43.0 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + github.com/rivo/uniseg v0.4.7 // indirect github.com/segmentio/asm v1.1.3 // indirect github.com/segmentio/encoding v0.5.4 // indirect + github.com/spf13/pflag v1.0.9 // indirect github.com/tidwall/gjson v1.18.0 // indirect github.com/tidwall/match v1.1.1 // indirect github.com/tidwall/pretty v1.2.1 // indirect github.com/tidwall/sjson v1.2.5 // indirect + github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e // indirect github.com/yosida95/uritemplate/v3 v3.0.2 // indirect golang.org/x/net v0.57.0 // indirect - golang.org/x/oauth2 v0.36.0 // indirect + golang.org/x/sync v0.22.0 // indirect golang.org/x/sys v0.47.0 // indirect golang.org/x/text v0.40.0 // indirect modernc.org/libc v1.74.1 // indirect diff --git a/go.sum b/go.sum index 92eecba..d30be3d 100644 --- a/go.sum +++ b/go.sum @@ -1,7 +1,58 @@ +charm.land/bubbles/v2 v2.0.0 h1:tE3eK/pHjmtrDiRdoC9uGNLgpopOd8fjhEe31B/ai5s= +charm.land/bubbles/v2 v2.0.0/go.mod h1:rCHoleP2XhU8um45NTuOWBPNVHxnkXKTiZqcclL/qOI= +charm.land/bubbletea/v2 v2.0.2 h1:4CRtRnuZOdFDTWSff9r8QFt/9+z6Emubz3aDMnf/dx0= +charm.land/bubbletea/v2 v2.0.2/go.mod h1:3LRff2U4WIYXy7MTxfbAQ+AdfM3D8Xuvz2wbsOD9OHQ= +charm.land/fang/v2 v2.0.1 h1:zQCM8JQJ1JnQX/66B5jlCYBUxL2as5JXQZ2KJ6EL0mY= +charm.land/fang/v2 v2.0.1/go.mod h1:S1GmkpcvK+OB5w9caywUnJcsMew45Ot8FXqoz8ALrII= charm.land/fantasy v0.38.1 h1:U1HNOVQaCK/3zwDbwUYZoiTXSZsoxn0JpnrVfqC4O5k= charm.land/fantasy v0.38.1/go.mod h1:4iPCl1cQYjWQHdN/L9n0sjWpWYuMY/ci0SWYcy5bONc= +charm.land/huh/v2 v2.0.3 h1:2cJsMqEPwSywGHvdlKsJyQKPtSJLVnFKyFbsYZTlLkU= +charm.land/huh/v2 v2.0.3/go.mod h1:93eEveeeqn47MwiC3tf+2atZ2l7Is88rAtmZNZ8x9Wc= +charm.land/lipgloss/v2 v2.0.1 h1:6Xzrn49+Py1Um5q/wZG1gWgER2+7dUyZ9XMEufqPSys= +charm.land/lipgloss/v2 v2.0.1/go.mod h1:KjPle2Qd3YmvP1KL5OMHiHysGcNwq6u83MUjYkFvEkM= +github.com/MakeNowJust/heredoc v1.0.0 h1:cXCdzVdstXyiTqTvfqk9SDHpKNjxuom+DOlyEeQ4pzQ= +github.com/MakeNowJust/heredoc v1.0.0/go.mod h1:mG5amYoWBHf8vpLOuehzbGGw0EHxpZZ6lCpQ4fNJ8LE= +github.com/atotto/clipboard v0.1.4 h1:EH0zSVneZPSuFR11BlR9YppQTVDbh5+16AmcJi4g1z4= +github.com/atotto/clipboard v0.1.4/go.mod h1:ZY9tmq7sm5xIbd9bOK4onWV4S6X0u6GY7Vn0Yu86PYI= +github.com/aymanbagabas/go-udiff v0.4.1 h1:OEIrQ8maEeDBXQDoGCbbTTXYJMYRCRO1fnodZ12Gv5o= +github.com/aymanbagabas/go-udiff v0.4.1/go.mod h1:0L9PGwj20lrtmEMeyw4WKJ/TMyDtvAoK9bf2u/mNo3w= +github.com/catppuccin/go v0.3.0 h1:d+0/YicIq+hSTo5oPuRi5kOpqkVA5tAsU6dNhvRu+aY= +github.com/catppuccin/go v0.3.0/go.mod h1:8IHJuMGaUUjQM82qBrGNBv7LFq6JI3NnQCF6MOlZjpc= +github.com/charmbracelet/colorprofile v0.4.2 h1:BdSNuMjRbotnxHSfxy+PCSa4xAmz7szw70ktAtWRYrY= +github.com/charmbracelet/colorprofile v0.4.2/go.mod h1:0rTi81QpwDElInthtrQ6Ni7cG0sDtwAd4C4le060fT8= +github.com/charmbracelet/ultraviolet v0.0.0-20260205113103-524a6607adb8 h1:eyFRbAmexyt43hVfeyBofiGSEmJ7krjLOYt/9CF5NKA= +github.com/charmbracelet/ultraviolet v0.0.0-20260205113103-524a6607adb8/go.mod h1:SQpCTRNBtzJkwku5ye4S3HEuthAlGy2n9VXZnWkEW98= +github.com/charmbracelet/x/ansi v0.11.6 h1:GhV21SiDz/45W9AnV2R61xZMRri5NlLnl6CVF7ihZW8= +github.com/charmbracelet/x/ansi v0.11.6/go.mod h1:2JNYLgQUsyqaiLovhU2Rv/pb8r6ydXKS3NIttu3VGZQ= +github.com/charmbracelet/x/conpty v0.1.1 h1:s1bUxjoi7EpqiXysVtC+a8RrvPPNcNvAjfi4jxsAuEs= +github.com/charmbracelet/x/conpty v0.1.1/go.mod h1:OmtR77VODEFbiTzGE9G1XiRJAga6011PIm4u5fTNZpk= +github.com/charmbracelet/x/errors v0.0.0-20240508181413-e8d8b6e2de86 h1:JSt3B+U9iqk37QUU2Rvb6DSBYRLtWqFqfxf8l5hOZUA= +github.com/charmbracelet/x/errors v0.0.0-20240508181413-e8d8b6e2de86/go.mod h1:2P0UgXMEa6TsToMSuFqKFQR+fZTO9CNGUNokkPatT/0= +github.com/charmbracelet/x/exp/charmtone v0.0.0-20250603201427-c31516f43444 h1:IJDiTgVE56gkAGfq0lBEloWgkXMk4hl/bmuPoicI4R0= +github.com/charmbracelet/x/exp/charmtone v0.0.0-20250603201427-c31516f43444/go.mod h1:T9jr8CzFpjhFVHjNjKwbAD7KwBNyFnj2pntAO7F2zw0= +github.com/charmbracelet/x/exp/golden v0.0.0-20250806222409-83e3a29d542f h1:pk6gmGpCE7F3FcjaOEKYriCvpmIN4+6OS/RD0vm4uIA= +github.com/charmbracelet/x/exp/golden v0.0.0-20250806222409-83e3a29d542f/go.mod h1:IfZAMTHB6XkZSeXUqriemErjAWCCzT0LwjKFYCZyw0I= +github.com/charmbracelet/x/exp/ordered v0.1.0 h1:55/qLwjIh0gL0Vni+QAWk7T/qRVP6sBf+2agPBgnOFE= +github.com/charmbracelet/x/exp/ordered v0.1.0/go.mod h1:5UHwmG+is5THxMyCJHNPCn2/ecI07aKNrW+LcResjJ8= github.com/charmbracelet/x/exp/slice v0.0.0-20250904123553-b4e2667e5ad5 h1:DTSZxdV9qQagD4iGcAt9RgaRBZtJl01bfKgdLzUzUPI= github.com/charmbracelet/x/exp/slice v0.0.0-20250904123553-b4e2667e5ad5/go.mod h1:vI5nDVMWi6veaYH+0Fmvpbe/+cv/iJfMntdh+N0+Tms= +github.com/charmbracelet/x/exp/strings v0.1.0 h1:i69S2XI7uG1u4NLGeJPSYU++Nmjvpo9nwd6aoEm7gkA= +github.com/charmbracelet/x/exp/strings v0.1.0/go.mod h1:/ehtMPNh9K4odGFkqYJKpIYyePhdp1hLBRvyY4bWkH8= +github.com/charmbracelet/x/term v0.2.2 h1:xVRT/S2ZcKdhhOuSP4t5cLi5o+JxklsoEObBSgfgZRk= +github.com/charmbracelet/x/term v0.2.2/go.mod h1:kF8CY5RddLWrsgVwpw4kAa6TESp6EB5y3uxGLeCqzAI= +github.com/charmbracelet/x/termios v0.1.1 h1:o3Q2bT8eqzGnGPOYheoYS8eEleT5ZVNYNy8JawjaNZY= +github.com/charmbracelet/x/termios v0.1.1/go.mod h1:rB7fnv1TgOPOyyKRJ9o+AsTU/vK5WHJ2ivHeut/Pcwo= +github.com/charmbracelet/x/windows v0.2.2 h1:IofanmuvaxnKHuV04sC0eBy/smG6kIKrWG2/jYn2GuM= +github.com/charmbracelet/x/windows v0.2.2/go.mod h1:/8XtdKZzedat74NQFn0NGlGL4soHB0YQZrETF96h75k= +github.com/charmbracelet/x/xpty v0.1.3 h1:eGSitii4suhzrISYH50ZfufV3v085BXQwIytcOdFSsw= +github.com/charmbracelet/x/xpty v0.1.3/go.mod h1:poPYpWuLDBFCKmKLDnhBp51ATa0ooD8FhypRwEFtH3Y= +github.com/clipperhouse/displaywidth v0.11.0 h1:lBc6kY44VFw+TDx4I8opi/EtL9m20WSEFgwIwO+UVM8= +github.com/clipperhouse/displaywidth v0.11.0/go.mod h1:bkrFNkf81G8HyVqmKGxsPufD3JhNl3dSqnGhOoSD/o0= +github.com/clipperhouse/uax29/v2 v2.7.0 h1:+gs4oBZ2gPfVrKPthwbMzWZDaAFPGYK72F0NJv2v7Vk= +github.com/clipperhouse/uax29/v2 v2.7.0/go.mod h1:EFJ2TJMRUaplDxHKj1qAEhCtQPW2tJSwu5BF98AuoVM= +github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g= +github.com/creack/pty v1.1.24 h1:bJrF4RRfyJnbTJqzRLHzcGaZK1NeM5kTC9jGgovnR1s= +github.com/creack/pty v1.1.24/go.mod h1:08sCNb52WyoAwi2QDyzUCTgcvVFhUzewun7wtTfvcwE= github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM= github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= @@ -24,26 +75,53 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= +github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= +github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/kaptinlin/jsonpointer v0.4.27 h1:5FOnhlkqQ4/lvHudaAWS8HJCXjN4yAHSIGl7aPKHI0Q= github.com/kaptinlin/jsonpointer v0.4.27/go.mod h1:dfub/n58cWS32Dyf3AZsnKblSAgrz9PyOU76GDPpx8Q= github.com/kaptinlin/jsonschema v0.9.3 h1:uDVd3w4aXwO0tbycblKYvFofhl3hVuE31vOl5JkDXI0= github.com/kaptinlin/jsonschema v0.9.3/go.mod h1:LvtQ/mO0E1e/3c3DiOWTj05LTeot8cTv7DyQq7eJMQg= +github.com/lucasb-eyer/go-colorful v1.3.0 h1:2/yBRLdWBZKrf7gB40FoiKfAWYQ0lqNcbuQwVHXptag= +github.com/lucasb-eyer/go-colorful v1.3.0/go.mod h1:R4dSotOR9KMtayYi1e77YzuveK+i7ruzyGqttikkLy0= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/mattn/go-runewidth v0.0.20 h1:WcT52H91ZUAwy8+HUkdM3THM6gXqXuLJi9O3rjcQQaQ= +github.com/mattn/go-runewidth v0.0.20/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs= +github.com/mitchellh/hashstructure/v2 v2.0.2 h1:vGKWl0YJqUNxE8d+h8f6NJLcCJrgbhC4NcD46KavDd4= +github.com/mitchellh/hashstructure/v2 v2.0.2/go.mod h1:MG3aRVU/N29oo/V/IhBX8GR/zz4kQkprJgF2EVszyDE= github.com/modelcontextprotocol/go-sdk v1.6.1 h1:0zOSupjKUxPKSocPT1Wtago+mUHU2/uZ4xSOY0FGReU= github.com/modelcontextprotocol/go-sdk v1.6.1/go.mod h1:kzm3kzFL1/+AziGOE0nUs3gvPoNxMCvkxokMkuFapXQ= +github.com/muesli/cancelreader v0.2.2 h1:3I4Kt4BQjOR54NavqnDogx/MIoWBFa0StPA8ELUXHmA= +github.com/muesli/cancelreader v0.2.2/go.mod h1:3XuTXfFS2VjM+HTLZY9Ak0l6eUKfijIfMUZ4EgX0QYo= +github.com/muesli/mango v0.1.0 h1:DZQK45d2gGbql1arsYA4vfg4d7I9Hfx5rX/GCmzsAvI= +github.com/muesli/mango v0.1.0/go.mod h1:5XFpbC8jY5UUv89YQciiXNlbi+iJgt29VDC5xbzrLL4= +github.com/muesli/mango-cobra v1.2.0 h1:DQvjzAM0PMZr85Iv9LIMaYISpTOliMEg+uMFtNbYvWg= +github.com/muesli/mango-cobra v1.2.0/go.mod h1:vMJL54QytZAJhCT13LPVDfkvCUJ5/4jNUKF/8NC2UjA= +github.com/muesli/mango-pflag v0.1.0 h1:UADqbYgpUyRoBja3g6LUL+3LErjpsOwaC9ywvBWe7Sg= +github.com/muesli/mango-pflag v0.1.0/go.mod h1:YEQomTxaCUp8PrbhFh10UfbhbQrM/xJ4i2PB8VTLLW0= +github.com/muesli/roff v0.1.0 h1:YD0lalCotmYuF5HhZliKWlIx7IEhiXeSfq7hNjFqGF8= +github.com/muesli/roff v0.1.0/go.mod h1:pjAHQM9hdUUwm/krAfrLGgJkXJ+YuhtsfZ42kieB2Ig= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= github.com/openai/openai-go/v3 v3.43.0 h1:C+MFVUMU3TJNgES+Ikt7HF8xcX7J0wynKeR9ST22hZM= github.com/openai/openai-go/v3 v3.43.0/go.mod h1:cdufnVK14cWcT9qA1rRtrXx4FTRsgbDPW7Ia7SS5cZo= +github.com/pkg/browser v0.0.0-20240102092130-5ac0b6a4141c h1:+mdjkGKdHQG3305AYmdv1U2eRNDiU2ErMBj1gwrq8eQ= +github.com/pkg/browser v0.0.0-20240102092130-5ac0b6a4141c/go.mod h1:7rwL4CYBLnjLxUqIJNnCWiEdr3bn6IUYi15bNlnbCCU= github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= +github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= +github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= +github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM= github.com/segmentio/asm v1.1.3 h1:WM03sfUOENvvKexOLp+pCqgb/WDjsi7EK8gIsICtzhc= github.com/segmentio/asm v1.1.3/go.mod h1:Ld3L4ZXGNcSLRg4JBsZ3//1+f/TjYl0Mzen/DQy1EJg= github.com/segmentio/encoding v0.5.4 h1:OW1VRern8Nw6ITAtwSZ7Idrl3MXCFwXHPgqESYfvNt0= github.com/segmentio/encoding v0.5.4/go.mod h1:HS1ZKa3kSN32ZHVZ7ZLPLXWvOVIiZtyJnO1gPH1sKt0= +github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU= +github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4= +github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY= +github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= @@ -56,8 +134,13 @@ github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4= github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY= github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28= +github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e h1:JVG44RsyaB9T2KIHavMF/ppJZNG9ZpyihvCd0w101no= +github.com/xo/terminfo v0.0.0-20220910002029-abceb7e1c41e/go.mod h1:RbqR21r5mrJuqunuUZ/Dhy/avygyECGrLceyNeo4LiM= github.com/yosida95/uritemplate/v3 v3.0.2 h1:Ed3Oyj9yrmi9087+NczuL5BwkIc4wvTb5zIM+UJPGz4= github.com/yosida95/uritemplate/v3 v3.0.2/go.mod h1:ILOh0sOhIJR3+L/8afwt/kE++YT040gmv5BQTMR2HP4= +go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= +golang.org/x/exp v0.0.0-20231006140011-7918f672742d h1:jtJma62tbqLibJ5sFQz8bKtEM8rJBtfilJ2qTU199MI= +golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQzzkH36rKxoZW1tw7ZJpeKx+hdo= golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= @@ -66,13 +149,17 @@ golang.org/x/oauth2 v0.36.0 h1:peZ/1z27fi9hUOFCAZaHyrpWG5lwe0RJEEEeH0ThlIs= golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/term v0.45.0 h1:NwWyBmoJCbfTHpxrWoZ9C6/VxOf7ic219I8xZZFdrf0= +golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w= golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= modernc.org/cc/v4 v4.29.0 h1:CXgwL8cvxmyzBQZzbSl/6xFtMCryb6u8IOqDci39cgc= diff --git a/internal/auth/auth.go b/internal/auth/auth.go index 5c10189..a159aae 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -44,6 +44,9 @@ const ( const ( PathProtectedResource = "/.well-known/oauth-protected-resource" PathAuthServer = "/.well-known/oauth-authorization-server" + // PathCollector is the unauthenticated prefix serving the collector OAuth + // registration and, for confidential clients, the code exchange. + PathCollector = "/auth/collector" ) // Config holds auth settings. @@ -66,6 +69,25 @@ type Config struct { AllowedUsers []string // RequiredScopes are scopes every token must carry. Empty means none. RequiredScopes []string + // CollectorClientID is the OAuth client this server tells collectors to + // use. It is trusted implicitly: publishing an identity and then rejecting + // it would be a contradiction, and forcing the operator to repeat it in + // AllowedClientIDs is a step that only ever gets forgotten. + CollectorClientID string +} + +// clientAllowlist is the set of client ids accepted, including the collector +// registration this server hands out. +func (c Config) clientAllowlist() []string { + if c.CollectorClientID == "" { + return c.AllowedClientIDs + } + if len(c.AllowedClientIDs) == 0 { + // An explicit collector id is itself a restriction, so honor it as the + // whole allowlist rather than treating "no list" as "allow anything". + return []string{c.CollectorClientID} + } + return append(append([]string{}, c.AllowedClientIDs...), c.CollectorClientID) } // Validate reports configuration problems worth logging at boot. Bearer mode @@ -81,7 +103,7 @@ func (c Config) Validate() []string { if c.PublicURL == "" { warns = append(warns, "bearer auth has no LARD_PUBLIC_URL; OAuth discovery metadata will be incomplete") } - if len(c.AllowedClientIDs) == 0 && len(c.AllowedUsers) == 0 { + if len(c.clientAllowlist()) == 0 && len(c.AllowedUsers) == 0 { warns = append(warns, "bearer auth has no audience restriction: any app the user authorized with indiko can read all memory (set LARD_OAUTH_CLIENT_IDS or LARD_OAUTH_USERS)") } return warns @@ -110,7 +132,9 @@ func Middleware(cfg Config, next http.Handler) http.Handler { v := &verifier{cfg: cfg, cache: map[string]cacheEntry{}} return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { p := r.URL.Path - if p == "/healthz" || strings.HasPrefix(p, "/.well-known/") { + // Discovery and the collector registration must be reachable without + // credentials: they are how a caller learns to get credentials. + if p == "/healthz" || strings.HasPrefix(p, "/.well-known/") || strings.HasPrefix(p, PathCollector) { next.ServeHTTP(w, r) return } @@ -229,9 +253,9 @@ func (v *verifier) authorize(r *http.Request) (id Identity, status int, code, de // Denials are logged with the claims that failed. A 403 here is almost // always an allowlist typo, and the operator cannot see the token's real // client_id or "me" value any other way. - if !allowed(v.cfg.AllowedClientIDs, res.ClientID) { + if !allowed(v.cfg.clientAllowlist(), res.ClientID) { slog.Warn("auth: client not allowed", - "token_client_id", res.ClientID, "allowed", v.cfg.AllowedClientIDs) + "token_client_id", res.ClientID, "allowed", v.cfg.clientAllowlist()) return id, http.StatusForbidden, "invalid_token", "token was not issued for this resource" } if !allowed(v.cfg.AllowedUsers, res.subject()) { diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go index ce4ae08..0cca1d4 100644 --- a/internal/auth/auth_test.go +++ b/internal/auth/auth_test.go @@ -287,3 +287,68 @@ func TestHealthzBypassesAuth(t *testing.T) { t.Fatalf("want 200, got %d", w.Code) } } + +// The client id a server publishes to collectors must be accepted without the +// operator also repeating it in the allowlist. Publishing an identity and then +// rejecting it is the failure this guards. +func TestCollectorClientIDIsTrustedImplicitly(t *testing.T) { + srv := fakeIndiko(t, map[string]any{ + "active": true, "me": "https://indiko.example.com/u/kieran", + "client_id": "ikc_collector", "scope": "profile", + }) + cfg := bearerConfig(srv.URL) + cfg.CollectorClientID = "ikc_collector" + h := Middleware(cfg, okHandler()) + if w := do(h, "Bearer good"); w.Code != http.StatusOK { + t.Fatalf("want 200 for the published collector id, got %d", w.Code) + } +} + +// An explicit collector id is itself a restriction, so it must not widen access +// to every client the user has ever authorized. +func TestCollectorClientIDNarrowsWhenNoOtherAllowlist(t *testing.T) { + srv := fakeIndiko(t, map[string]any{ + "active": true, "me": "https://indiko.example.com/u/kieran", + "client_id": "https://some-other-app.example.com/", "scope": "profile", + }) + cfg := bearerConfig(srv.URL) + cfg.CollectorClientID = "ikc_collector" + h := Middleware(cfg, okHandler()) + if w := do(h, "Bearer good"); w.Code != http.StatusForbidden { + t.Fatalf("want 403 for a foreign client, got %d", w.Code) + } +} + +func TestCollectorClientIDAddsToExistingAllowlist(t *testing.T) { + srv := fakeIndiko(t, map[string]any{ + "active": true, "me": "https://indiko.example.com/u/kieran", + "client_id": "ikc_mcp", "scope": "profile", + }) + cfg := bearerConfig(srv.URL) + cfg.AllowedClientIDs = []string{"ikc_mcp"} + cfg.CollectorClientID = "ikc_collector" + h := Middleware(cfg, okHandler()) + if w := do(h, "Bearer good"); w.Code != http.StatusOK { + t.Fatalf("want 200, got %d", w.Code) + } +} + +// Setting only a collector id counts as an audience restriction, so the +// "anyone can read your memory" warning must not fire. +func TestValidateAcceptsCollectorIDAsAudience(t *testing.T) { + cfg := Config{Mode: ModeBearer, IndikoURL: "https://i", PublicURL: "https://l", CollectorClientID: "ikc_x"} + if warns := cfg.Validate(); len(warns) != 0 { + t.Fatalf("want no warnings, got %v", warns) + } +} + +// The collector registration must be reachable without credentials, since it +// is how a collector learns which client to be. +func TestCollectorPathBypassesAuth(t *testing.T) { + h := Middleware(Config{Mode: ModeToken, Token: "s3cret"}, okHandler()) + w := httptest.NewRecorder() + h.ServeHTTP(w, httptest.NewRequest(http.MethodGet, PathCollector, nil)) + if w.Code != http.StatusOK { + t.Fatalf("want 200, got %d", w.Code) + } +} diff --git a/internal/client/client.go b/internal/client/client.go index 925587b..b9eb58c 100644 --- a/internal/client/client.go +++ b/internal/client/client.go @@ -12,6 +12,7 @@ import ( "time" "github.com/taciturnaxolotl/lard/internal/types" + "github.com/taciturnaxolotl/lard/internal/xdg" ) // Uploader pushes session batches to the central lard service. @@ -148,10 +149,8 @@ func (s *State) Save() error { return os.Rename(tmp, s.path) } -// DefaultStatePath is ~/.config/lard/client-state.json. +// DefaultStatePath is ~/.local/share/lard/client-state.json. State, not +// config: it is a sync watermark the user never edits by hand. func DefaultStatePath() string { - if d, err := os.UserConfigDir(); err == nil { - return filepath.Join(d, "lard", "client-state.json") - } - return "client-state.json" + return xdg.DataPath("client-state.json") } diff --git a/internal/client/config.go b/internal/client/config.go new file mode 100644 index 0000000..b9339bc --- /dev/null +++ b/internal/client/config.go @@ -0,0 +1,185 @@ +// Package client also owns the collector's own configuration: where the +// central lard lives and how to authenticate to it. +package client + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "os" + "path/filepath" + "strings" + "time" + + "github.com/taciturnaxolotl/lard/internal/xdg" +) + +// Config is the collector's persisted setup. It exists so a background daemon +// needs no environment at all: launchd and systemd both make per-service env +// vars awkward, and a file the user can inspect beats an invisible plist. +type Config struct { + URL string `json:"url"` + Roots []string `json:"roots,omitempty"` + // Token is a static shared secret, for a server running LARD_AUTH=token + // or a headless box where no browser is available. + Token string `json:"token,omitempty"` + // OAuth holds the browser-login credentials, used when the server runs + // LARD_AUTH=bearer. Preferred over Token: nothing to copy by hand, and it + // carries the same identity as the rest of the user's tooling. + OAuth *OAuthToken `json:"oauth,omitempty"` +} + +// OAuthToken is the persisted result of `lard-client login`. +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"` +} + +// expired reports whether the access token is gone or about to lapse. The +// minute of slack keeps a long upload from dying mid-flight. +func (t *OAuthToken) expired() bool { + if t == nil || t.AccessToken == "" { + return true + } + if t.Expiry.IsZero() { + return false + } + return time.Now().After(t.Expiry.Add(-time.Minute)) +} + +// DefaultConfigPath is ~/.config/lard/client.json. +func DefaultConfigPath() string { + return xdg.ConfigPath("client.json") +} + +// LoadConfig reads the config file, tolerating absence, then lets the +// environment override it. Env wins so a one-off run can point somewhere else +// without editing the file. +func LoadConfig(path string) (*Config, error) { + cfg := &Config{} + b, err := os.ReadFile(path) + if err != nil && !os.IsNotExist(err) { + return nil, err + } + if err == nil { + if err := json.Unmarshal(b, cfg); err != nil { + return nil, fmt.Errorf("parse %s: %w", path, err) + } + } + if v := os.Getenv("LARD_URL"); v != "" { + cfg.URL = v + } + if v := os.Getenv("LARD_TOKEN"); v != "" { + cfg.Token = v + } + if cfg.URL == "" { + cfg.URL = "http://localhost:7477" + } + cfg.URL = strings.TrimRight(cfg.URL, "/") + return cfg, nil +} + +// Save writes the config atomically with owner-only permissions, since it +// holds a bearer token. +func (c *Config) Save(path string) error { + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return err + } + b, err := json.MarshalIndent(c, "", " ") + if err != nil { + return err + } + tmp := path + ".tmp" + if err := os.WriteFile(tmp, append(b, '\n'), 0o600); err != nil { + return err + } + return os.Rename(tmp, path) +} + +// Verify checks that the configured URL and token actually work, so setup +// fails loudly at install time rather than silently in a background daemon. +// It returns the identity the server reports, when it reports one. +func (c *Config) Verify(ctx context.Context) (string, error) { + ctx, cancel := context.WithTimeout(ctx, 10*time.Second) + defer cancel() + req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.URL+"/whoami", nil) + if err != nil { + return "", err + } + if tok, err := c.Bearer(ctx, DefaultConfigPath()); err == nil && tok != "" { + req.Header.Set("authorization", "Bearer "+tok) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + return "", fmt.Errorf("cannot reach %s: %w", c.URL, err) + } + defer resp.Body.Close() + switch { + case resp.StatusCode == http.StatusUnauthorized: + return "", errors.New("server rejected the token (401); check LARD_TOKEN matches the server's") + case resp.StatusCode == http.StatusForbidden: + return "", errors.New("token is valid but not allowed for this server (403); check the server's allowlist") + case resp.StatusCode == http.StatusNotFound: + // An older server without /whoami. Reaching it at all is enough. + return "", nil + case resp.StatusCode != http.StatusOK: + return "", fmt.Errorf("unexpected status %d from %s/whoami", resp.StatusCode, c.URL) + } + var body struct { + Subject string `json:"subject"` + ClientID string `json:"clientId"` + } + _ = json.NewDecoder(resp.Body).Decode(&body) + if body.Subject != "" { + return body.Subject, nil + } + return body.ClientID, nil +} + +// Bearer returns the token to send, refreshing an expired OAuth access token +// and persisting the new one. It is the single place that decides between +// OAuth and a static secret, so callers never branch on auth mode. +// +// A refresh failure is fatal by design: silently falling back to no +// credentials would turn an expired login into a stream of 401s in a log file +// nobody reads. +func (c *Config) Bearer(ctx context.Context, path string) (string, error) { + if c.OAuth != nil && c.OAuth.AccessToken != "" { + if !c.OAuth.expired() { + return c.OAuth.AccessToken, nil + } + tok, err := RefreshToken(ctx, c.URL, c.OAuth.RefreshToken, c.OAuth.CallbackPort) + if err != nil { + return "", err + } + c.OAuth.AccessToken = tok.AccessToken + c.OAuth.Expiry = tok.Expiry + // Indiko does not rotate refresh tokens, but honor one if it appears. + if tok.RefreshToken != "" { + c.OAuth.RefreshToken = tok.RefreshToken + } + if err := c.Save(path); err != nil { + return "", fmt.Errorf("save refreshed token: %w", err) + } + return c.OAuth.AccessToken, nil + } + return c.Token, nil +} + +// AuthMode names how this config authenticates, for status output. +func (c *Config) AuthMode() string { + switch { + case c.OAuth != nil && c.OAuth.AccessToken != "": + return "oauth" + case c.Token != "": + return "token" + default: + return "none" + } +} diff --git a/internal/client/device.go b/internal/client/device.go new file mode 100644 index 0000000..997d6ba --- /dev/null +++ b/internal/client/device.go @@ -0,0 +1,190 @@ +package client + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "strings" + "time" + + "github.com/pkg/browser" + "golang.org/x/oauth2" +) + +// Poll outcomes, matching RFC 8628's error codes. +const ( + errAuthorizationPending = "authorization_pending" + errSlowDown = "slow_down" + errExpiredToken = "expired_token" + 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. +// +// 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) + if err != nil { + return nil, err + } + printDevicePrompt(auth, openBrowser) + + interval := time.Duration(auth.Interval) * time.Second + if interval <= 0 { + interval = 2 * time.Second + } + deadline := time.Now().Add(time.Duration(auth.ExpiresIn) * time.Second) + if auth.ExpiresIn <= 0 { + deadline = time.Now().Add(10 * time.Minute) + } + + for { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-time.After(interval): + } + 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) + if err != nil { + return nil, err + } + switch code { + case "": + fmt.Println("Connected.") + return tok, nil + 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 + case errExpiredToken: + return nil, errors.New("this login expired; run 'lard-client login' again") + case errAccessDenied: + return nil, errors.New("authorization was declined") + default: + return nil, fmt.Errorf("authorization failed: %s", code) + } + } +} + +// 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...") +} + +func startDevice(ctx context.Context, serverURL 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("{}")) + 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("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) + } + if auth.DeviceCode == "" || auth.VerificationURI == "" { + return nil, fmt.Errorf("starting authorization: server returned status %d", resp.StatusCode) + } + 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 + } + 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)) + if err != nil { + return nil, "", err + } + req.Header.Set("content-type", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + // A transient network blip should not abandon a login the user is + // halfway through, so treat it as "keep waiting". + return nil, errAuthorizationPending, nil + } + 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("polling: unreadable response (status %d)", resp.StatusCode) + } + if out.Error != "" { + return nil, out.Error, nil + } + if out.AccessToken == "" { + return nil, "", fmt.Errorf("polling: server returned status %d", resp.StatusCode) + } + return &oauth2.Token{ + AccessToken: out.AccessToken, + RefreshToken: out.RefreshToken, + Expiry: out.Expiry, + TokenType: "Bearer", + }, "", nil +} diff --git a/internal/client/oauth.go b/internal/client/oauth.go new file mode 100644 index 0000000..b28b2b2 --- /dev/null +++ b/internal/client/oauth.go @@ -0,0 +1,488 @@ +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 +} + +// 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. +func FetchRegistration(ctx context.Context, serverURL string) (*Registration, error) { + ctx, cancel := context.WithTimeout(ctx, 10*time.Second) + defer cancel() + var reg Registration + if err := getJSON(ctx, strings.TrimRight(serverURL, "/")+"/auth/collector", ®); err != nil { + return nil, err + } + if reg.ClientID == "" { + return nil, errors.New("server published an empty client id") + } + return ®, nil +} + +// endpoints are the OAuth endpoints discovered from a lard server. +type endpoints struct { + Issuer string + Authorization string + Token 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. +func Discover(ctx context.Context, serverURL string) (*endpoints, error) { + ctx, cancel := context.WithTimeout(ctx, 15*time.Second) + defer cancel() + base := strings.TrimRight(serverURL, "/") + + var prm struct { + AuthorizationServers []string `json:"authorization_servers"` + } + if err := getJSON(ctx, base+"/.well-known/oauth-protected-resource", &prm); err != nil { + return nil, fmt.Errorf("discover protected resource: %w", err) + } + if len(prm.AuthorizationServers) == 0 { + return nil, errors.New("server published no authorization servers") + } + as := strings.TrimRight(prm.AuthorizationServers[0], "/") + + var meta struct { + Issuer string `json:"issuer"` + Authorization string `json:"authorization_endpoint"` + Token string `json:"token_endpoint"` + } + if err := getJSON(ctx, as+"/.well-known/oauth-authorization-server", &meta); err != nil { + return nil, fmt.Errorf("discover authorization server: %w", err) + } + 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 +} + +func getJSON(ctx context.Context, url string, out any) error { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return err + } + req.Header.Set("accept", "application/json") + resp, err := http.DefaultClient.Do(req) + if err != nil { + return err + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return fmt.Errorf("%s returned %d", url, resp.StatusCode) + } + 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) { + 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 + } + 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{ + AuthURL: eps.Authorization, + TokenURL: eps.Token, + AuthStyle: oauth2.AuthStyleInParams, + }, + } + tok, err := cfg.TokenSource(ctx, &oauth2.Token{RefreshToken: refresh}).Token() + if err != nil { + return nil, fmt.Errorf("refresh token rejected (re-run 'lard-client login'): %w", err) + } + 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}) + 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/refresh", 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 refresh: %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 + } + return "this-host" +} diff --git a/internal/client/oauth_test.go b/internal/client/oauth_test.go new file mode 100644 index 0000000..100325f --- /dev/null +++ b/internal/client/oauth_test.go @@ -0,0 +1,151 @@ +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 new file mode 100644 index 0000000..d157562 --- /dev/null +++ b/internal/collector/collector.go @@ -0,0 +1,312 @@ +// Package collector serves the OAuth registration that edge collectors use to +// authenticate. +// +// 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. +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. + 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 + } + 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"` +} + +// Handler serves the collector registration, the brokered device login, and +// (for confidential clients) the code exchange. +type Handler struct { + cfg Config + dev DeviceConfig + prefix string + devices *deviceStore +} + +// 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()} +} + +// Register serves GET: the registration a collector should adopt. +func (h *Handler) Register(w http.ResponseWriter, r *http.Request) { + if !h.cfg.Configured() { + http.Error(w, `{"error":"no collector registration configured"}`, http.StatusNotFound) + return + } + 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}, + }) + 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 new file mode 100644 index 0000000..170f1d1 --- /dev/null +++ b/internal/collector/collector_test.go @@ -0,0 +1,187 @@ +package collector + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestRegisterNotFoundWhenUnconfigured(t *testing.T) { + w := httptest.NewRecorder() + New(Config{}, DeviceConfig{}, "/auth/collector").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") + w := httptest.NewRecorder() + h.Register(w, httptest.NewRequest(http.MethodGet, "/auth/collector", nil)) + if w.Code != http.StatusOK { + t.Fatalf("want 200, got %d", w.Code) + } + var reg Registration + if err := json.Unmarshal(w.Body.Bytes(), ®); err != nil { + t.Fatal(err) + } + 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") + } +} + +func TestPublicClientDoesNotRequestServerExchange(t *testing.T) { + h := New(Config{ClientID: "https://app.example.com/"}, DeviceConfig{}, "/auth/collector") + 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) + } +} diff --git a/internal/collector/device.go b/internal/collector/device.go new file mode 100644 index 0000000..70d8afc --- /dev/null +++ b/internal/collector/device.go @@ -0,0 +1,222 @@ +package collector + +import ( + "crypto/rand" + "encoding/base32" + "errors" + "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 +} + +func (s *deviceSession) done() bool { return s.Token != nil || s.Err != "" } + +// 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 +} + +var errNoSession = errors.New("no such authorization request") diff --git a/internal/collector/device_test.go b/internal/collector/device_test.go new file mode 100644 index 0000000..febd0b7 --- /dev/null +++ b/internal/collector/device_test.go @@ -0,0 +1,313 @@ +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/httpapi/autoconsolidate.go b/internal/httpapi/autoconsolidate.go new file mode 100644 index 0000000..be2ee5c --- /dev/null +++ b/internal/httpapi/autoconsolidate.go @@ -0,0 +1,110 @@ +package httpapi + +import ( + "context" + "log/slog" + "sync" + "time" +) + +// Consolidation cadence. Sessions arrive in bursts (a laptop wakes up and +// uploads an afternoon's work), so consolidating on every ingest would burn +// API calls on a queue that is still filling. Instead each ingest resets a +// short quiet timer, and the pass runs once the uploads stop. +const ( + // DefaultConsolidateAfter is how long the queue must be quiet. + DefaultConsolidateAfter = 5 * time.Minute + // DefaultConsolidateMaxWait caps the total delay, so a machine uploading + // continuously still gets consolidated instead of deferring forever. + DefaultConsolidateMaxWait = 30 * time.Minute +) + +// autoConsolidator runs a consolidation pass once ingests go quiet. +// +// It guarantees only one pass runs at a time. A trigger arriving mid-pass is +// remembered and starts a fresh wait afterwards, so sessions uploaded during a +// long pass are never dropped on the floor. +type autoConsolidator struct { + after time.Duration + maxWait time.Duration + run func(context.Context) error + + mu sync.Mutex + timer *time.Timer + deadline time.Time // hard cap for the current burst + running bool + pending bool // a trigger arrived while a pass was running +} + +func newAutoConsolidator(after, maxWait time.Duration, run func(context.Context) error) *autoConsolidator { + if after <= 0 { + after = DefaultConsolidateAfter + } + if maxWait < after { + maxWait = after + } + return &autoConsolidator{after: after, maxWait: maxWait, run: run} +} + +// Trigger notes that new sessions landed and (re)starts the quiet timer. +func (a *autoConsolidator) Trigger() { + a.mu.Lock() + defer a.mu.Unlock() + if a.running { + a.pending = true + return + } + now := time.Now() + if a.deadline.IsZero() { + a.deadline = now.Add(a.maxWait) + } + // Wait for quiet, but never past the burst deadline. + wait := a.after + if left := time.Until(a.deadline); left < wait { + wait = max(left, 0) + } + if a.timer != nil { + a.timer.Stop() + } + a.timer = time.AfterFunc(wait, a.fire) + slog.Debug("consolidate: scheduled", "in", wait) +} + +// Stop cancels any pending pass. Safe to call more than once. +func (a *autoConsolidator) Stop() { + a.mu.Lock() + defer a.mu.Unlock() + if a.timer != nil { + a.timer.Stop() + a.timer = nil + } +} + +func (a *autoConsolidator) fire() { + a.mu.Lock() + if a.running { + a.pending = true + a.mu.Unlock() + return + } + a.running = true + a.timer = nil + a.deadline = time.Time{} + a.mu.Unlock() + + slog.Info("consolidate: starting scheduled pass") + if err := a.run(context.Background()); err != nil { + slog.Error("consolidate: scheduled pass failed", "error", err) + } + + a.mu.Lock() + a.running = false + again := a.pending + a.pending = false + a.mu.Unlock() + + // Sessions arrived mid-pass; wait out another quiet period for them. + if again { + a.Trigger() + } +} diff --git a/internal/httpapi/autoconsolidate_test.go b/internal/httpapi/autoconsolidate_test.go new file mode 100644 index 0000000..a201632 --- /dev/null +++ b/internal/httpapi/autoconsolidate_test.go @@ -0,0 +1,129 @@ +package httpapi + +import ( + "context" + "sync/atomic" + "testing" + "time" +) + +func TestAutoConsolidateWaitsForQuiet(t *testing.T) { + var runs atomic.Int32 + a := newAutoConsolidator(60*time.Millisecond, time.Second, func(context.Context) error { + runs.Add(1) + return nil + }) + defer a.Stop() + + // A burst of ingests should collapse into a single pass. + for range 5 { + a.Trigger() + time.Sleep(15 * time.Millisecond) + } + if got := runs.Load(); got != 0 { + t.Fatalf("ran %d times during the burst; should have waited", got) + } + time.Sleep(150 * time.Millisecond) + if got := runs.Load(); got != 1 { + t.Fatalf("runs = %d, want exactly 1", got) + } +} + +// A machine uploading continuously must not defer consolidation forever. +func TestAutoConsolidateHonorsMaxWait(t *testing.T) { + var runs atomic.Int32 + a := newAutoConsolidator(50*time.Millisecond, 120*time.Millisecond, func(context.Context) error { + runs.Add(1) + return nil + }) + defer a.Stop() + + deadline := time.Now().Add(300 * time.Millisecond) + for time.Now().Before(deadline) { + a.Trigger() + time.Sleep(20 * time.Millisecond) + } + if got := runs.Load(); got == 0 { + t.Fatal("never ran despite continuous triggers; max wait was not honored") + } +} + +// Sessions arriving mid-pass must get their own pass afterwards, not be lost. +func TestAutoConsolidateReschedulesTriggerDuringRun(t *testing.T) { + var runs atomic.Int32 + started := make(chan struct{}, 4) + a := newAutoConsolidator(20*time.Millisecond, time.Second, func(context.Context) error { + started <- struct{}{} + runs.Add(1) + time.Sleep(60 * time.Millisecond) + return nil + }) + defer a.Stop() + + a.Trigger() + <-started // first pass is now running + a.Trigger() + + time.Sleep(250 * time.Millisecond) + if got := runs.Load(); got < 2 { + t.Fatalf("runs = %d, want at least 2 (the mid-pass trigger was dropped)", got) + } +} + +func TestAutoConsolidateOnlyOneAtATime(t *testing.T) { + var concurrent, maxSeen atomic.Int32 + a := newAutoConsolidator(10*time.Millisecond, time.Second, func(context.Context) error { + n := concurrent.Add(1) + for { + m := maxSeen.Load() + if n <= m || maxSeen.CompareAndSwap(m, n) { + break + } + } + time.Sleep(30 * time.Millisecond) + concurrent.Add(-1) + return nil + }) + defer a.Stop() + + for range 10 { + a.Trigger() + time.Sleep(5 * time.Millisecond) + } + time.Sleep(200 * time.Millisecond) + if got := maxSeen.Load(); got > 1 { + t.Fatalf("saw %d concurrent passes, want 1", got) + } +} + +// A failing pass must not wedge the scheduler. +func TestAutoConsolidateRecoversFromError(t *testing.T) { + var runs atomic.Int32 + a := newAutoConsolidator(20*time.Millisecond, time.Second, func(context.Context) error { + runs.Add(1) + return context.DeadlineExceeded + }) + defer a.Stop() + + a.Trigger() + time.Sleep(80 * time.Millisecond) + a.Trigger() + time.Sleep(80 * time.Millisecond) + if got := runs.Load(); got != 2 { + t.Fatalf("runs = %d, want 2", got) + } +} + +func TestAutoConsolidateStopCancelsPending(t *testing.T) { + var runs atomic.Int32 + a := newAutoConsolidator(50*time.Millisecond, time.Second, func(context.Context) error { + runs.Add(1) + return nil + }) + a.Trigger() + a.Stop() + time.Sleep(120 * time.Millisecond) + if got := runs.Load(); got != 0 { + t.Fatalf("runs = %d after Stop, want 0", got) + } +} diff --git a/internal/httpapi/httpapi.go b/internal/httpapi/httpapi.go index 6651fe0..0966224 100644 --- a/internal/httpapi/httpapi.go +++ b/internal/httpapi/httpapi.go @@ -14,6 +14,7 @@ import ( "strings" "time" + "github.com/taciturnaxolotl/lard/internal/auth" "github.com/taciturnaxolotl/lard/internal/llm" "github.com/taciturnaxolotl/lard/internal/pipeline" "github.com/taciturnaxolotl/lard/internal/store" @@ -26,6 +27,7 @@ type Server struct { registry *pipeline.Registry llm *llm.Client mux *http.ServeMux + auto *autoConsolidator } // New builds the HTTP server. llmClient may be nil if consolidation is never @@ -36,6 +38,28 @@ func New(st *store.Store, llmClient *llm.Client) *Server { return s } +// EnableAutoConsolidate makes the server consolidate on its own once uploads +// go quiet, so memory stays current without anyone calling /consolidate. No-op +// without an LLM client, since a pass would fail anyway. +func (s *Server) EnableAutoConsolidate(after, maxWait time.Duration) { + if s.llm == nil { + slog.Warn("auto-consolidate disabled: no LLM client") + return + } + s.auto = newAutoConsolidator(after, maxWait, func(ctx context.Context) error { + _, err := s.Consolidator().Run(ctx, 0) + return err + }) + slog.Info("auto-consolidate enabled", "quiet_period", after, "max_wait", maxWait) +} + +// StopAutoConsolidate cancels any pending scheduled pass. +func (s *Server) StopAutoConsolidate() { + if s.auto != nil { + s.auto.Stop() + } +} + func (s *Server) Handler() http.Handler { return s.mux } func (s *Server) Registry() *pipeline.Registry { return s.registry } func (s *Server) Store() *store.Store { return s.st } @@ -46,6 +70,11 @@ func (s *Server) Consolidator() *pipeline.Consolidator { func (s *Server) routes() { s.mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, r *http.Request) { writeJSON(w, 200, map[string]string{"status": "ok"}) }) + // Authenticated echo, so a client can prove its credentials work before + // installing itself as a background service. /healthz bypasses auth and + // therefore cannot answer that question. + s.mux.HandleFunc("GET /whoami", s.handleWhoami) + // Context bundle (session start). s.mux.HandleFunc("GET /context", s.handleContext) @@ -307,6 +336,18 @@ func writeConflict(w http.ResponseWriter, current *types.Subject) { // --- ingest & consolidate --- +// handleWhoami reports the authenticated caller. Reaching this handler at all +// means the credentials are good; the body says which identity they mapped to. +func (s *Server) handleWhoami(w http.ResponseWriter, r *http.Request) { + out := map[string]any{"authenticated": true} + if id, ok := auth.IdentityFrom(r.Context()); ok { + out["subject"] = id.Subject + out["clientId"] = id.ClientID + out["scopes"] = id.Scopes + } + writeJSON(w, 200, out) +} + func (s *Server) handleIngest(w http.ResponseWriter, r *http.Request) { var req types.IngestRequest if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, 64<<20)).Decode(&req); err != nil { @@ -325,6 +366,10 @@ func (s *Server) handleIngest(w http.ResponseWriter, r *http.Request) { } } } + // New work landed: start (or extend) the quiet period before consolidating. + if n > 0 && s.auto != nil { + s.auto.Trigger() + } writeJSON(w, 200, map[string]int{"ingested": n}) } diff --git a/internal/service/service.go b/internal/service/service.go new file mode 100644 index 0000000..3add404 --- /dev/null +++ b/internal/service/service.go @@ -0,0 +1,195 @@ +// Package service installs the collector as a background agent so it keeps +// running across logins and reboots without the user writing plist XML. +// +// macOS only for now. Linux (systemd user units) is the obvious next step, and +// the Install/Uninstall/Status shape is deliberately platform-neutral so +// adding it means one more file rather than reshaping callers. +package service + +import ( + "fmt" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "time" +) + +// Label is the launchd job label, also used for the plist filename. +const Label = "sh.dunkirk.lard-client" + +// Options describe the agent to install. +type Options struct { + // Binary is the absolute path to lard-client. + Binary string + // Interval between sync passes. + Interval time.Duration + // Roots to scan for .crush directories. Empty means the client's own + // default (~/code). + Roots []string + // LogDir holds stdout/stderr. Defaults to ~/Library/Logs. + LogDir string +} + +// Supported reports whether service management works on this platform. +func Supported() bool { return runtime.GOOS == "darwin" } + +// PlistPath is where the LaunchAgent definition lives. +func PlistPath() (string, error) { + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + return filepath.Join(home, "Library", "LaunchAgents", Label+".plist"), nil +} + +// Install writes the LaunchAgent and loads it. Re-running is safe: the old job +// is unloaded first, so this doubles as "upgrade to the current binary". +func Install(opts Options) (plistPath string, err error) { + if !Supported() { + return "", fmt.Errorf("service install is only implemented for macOS; run 'lard-client daemon' under your init system instead") + } + if opts.Binary == "" { + return "", fmt.Errorf("binary path required") + } + if opts.Interval <= 0 { + opts.Interval = 5 * time.Minute + } + if opts.LogDir == "" { + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + opts.LogDir = filepath.Join(home, "Library", "Logs") + } + if err := os.MkdirAll(opts.LogDir, 0o755); err != nil { + return "", err + } + path, err := PlistPath() + if err != nil { + return "", err + } + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + return "", err + } + // Unload any previous version first; ignore the error since a missing job + // is the normal first-install case. + _ = unload(path) + if err := os.WriteFile(path, []byte(renderPlist(opts)), 0o644); err != nil { + return "", err + } + if err := load(path); err != nil { + return path, err + } + return path, nil +} + +// Uninstall stops the agent and removes its definition. +func Uninstall() error { + if !Supported() { + return fmt.Errorf("service management is only implemented for macOS") + } + path, err := PlistPath() + if err != nil { + return err + } + if _, err := os.Stat(path); os.IsNotExist(err) { + return fmt.Errorf("not installed (no %s)", path) + } + _ = unload(path) + return os.Remove(path) +} + +// Status reports whether the agent is installed and currently loaded. +func Status() (installed, loaded bool, detail string, err error) { + if !Supported() { + return false, false, "", fmt.Errorf("service management is only implemented for macOS") + } + path, err := PlistPath() + if err != nil { + return false, false, "", err + } + if _, statErr := os.Stat(path); statErr != nil { + return false, false, "", nil + } + out, listErr := exec.Command("launchctl", "list", Label).CombinedOutput() + if listErr != nil { + return true, false, "", nil + } + return true, true, parseLaunchctlList(string(out)), nil +} + +// parseLaunchctlList pulls the PID and last exit status out of launchctl's +// plist-ish output, which is more useful to a human than the raw dump. +func parseLaunchctlList(out string) string { + var pid, status string + for _, line := range strings.Split(out, "\n") { + line = strings.TrimSpace(line) + switch { + case strings.HasPrefix(line, `"PID" = `): + pid = strings.Trim(strings.TrimPrefix(line, `"PID" = `), ";") + case strings.HasPrefix(line, `"LastExitStatus" = `): + status = strings.Trim(strings.TrimPrefix(line, `"LastExitStatus" = `), ";") + } + } + switch { + case pid != "": + return "running, pid " + pid + case status != "" && status != "0": + return "not running, last exit status " + status + default: + return "loaded, waiting for next interval" + } +} + +func load(path string) error { + if out, err := exec.Command("launchctl", "load", "-w", path).CombinedOutput(); err != nil { + return fmt.Errorf("launchctl load: %w: %s", err, strings.TrimSpace(string(out))) + } + return nil +} + +func unload(path string) error { + return exec.Command("launchctl", "unload", "-w", path).Run() +} + +// renderPlist builds the LaunchAgent definition. StartInterval handles the +// schedule, so the daemon runs one pass and exits rather than sleeping in a +// loop: launchd is a better timekeeper than a long-lived process that a laptop +// suspend can stall. +func renderPlist(opts Options) string { + args := []string{opts.Binary, "sync"} + for _, r := range opts.Roots { + args = append(args, "--root", r) + } + var b strings.Builder + b.WriteString(`` + "\n") + b.WriteString(`` + "\n") + b.WriteString(`` + "\n\n") + fmt.Fprintf(&b, " Label\n %s\n", Label) + b.WriteString(" ProgramArguments\n \n") + for _, a := range args { + fmt.Fprintf(&b, " %s\n", escapeXML(a)) + } + b.WriteString(" \n") + fmt.Fprintf(&b, " StartInterval\n %d\n", int(opts.Interval.Seconds())) + // Also run once at load, so install gives immediate feedback. + b.WriteString(" RunAtLoad\n \n") + fmt.Fprintf(&b, " StandardOutPath\n %s\n", + escapeXML(filepath.Join(opts.LogDir, "lard-client.log"))) + fmt.Fprintf(&b, " StandardErrorPath\n %s\n", + escapeXML(filepath.Join(opts.LogDir, "lard-client.log"))) + // Keep the job quiet in Activity Monitor and off the critical path. + b.WriteString(" ProcessType\n Background\n") + b.WriteString(" LowPriorityIO\n \n") + b.WriteString("\n\n") + return b.String() +} + +func escapeXML(s string) string { + s = strings.ReplaceAll(s, "&", "&") + s = strings.ReplaceAll(s, "<", "<") + s = strings.ReplaceAll(s, ">", ">") + return s +} diff --git a/internal/setup/setup.go b/internal/setup/setup.go new file mode 100644 index 0000000..9b00f8d --- /dev/null +++ b/internal/setup/setup.go @@ -0,0 +1,247 @@ +// Package setup is the collector's interactive first-run flow: ask where the +// server is, authenticate, and save the result. +// +// It exists so a new machine needs no flags and no copied secrets. Every +// prompt has a non-interactive equivalent, so scripts and headless boxes are +// never forced through a TUI. +package setup + +import ( + "context" + "errors" + "fmt" + "net/url" + "os" + "strings" + + "charm.land/huh/v2" + "golang.org/x/term" + + "github.com/taciturnaxolotl/lard/internal/client" +) + +// Interactive reports whether we can prompt: a TTY on both ends, and no CI +// marker. Anything else must fail with a message instead of hanging on input +// that will never arrive. +func Interactive() bool { + if os.Getenv("CI") != "" || os.Getenv("LARD_NO_INTERACTIVE") != "" { + return false + } + return term.IsTerminal(int(os.Stdin.Fd())) && term.IsTerminal(int(os.Stdout.Fd())) +} + +// 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 +} + +// Run resolves a working configuration, prompting where needed, and saves it. +func Run(ctx context.Context, opts Options) (*client.Config, error) { + path := client.DefaultConfigPath() + cfg, err := client.LoadConfig(path) + if err != nil { + return nil, err + } + if opts.URL != "" { + cfg.URL = normalizeURL(opts.URL) + } + if len(opts.Roots) > 0 { + cfg.Roots = opts.Roots + } + + // LoadConfig defaults the URL to localhost, which is a fine fallback but a + // bad thing to silently adopt on a machine whose server is elsewhere. Treat + // an unedited default as "not yet configured" and ask. + if opts.URL == "" && needsURL(path, cfg) { + if !Interactive() { + return nil, errors.New("no server configured: pass --url https://lard.example.com") + } + if err := askURL(cfg); err != nil { + return nil, err + } + } + if cfg.URL == "" { + return nil, errors.New("no server URL") + } + + if err := authenticate(ctx, cfg, opts); err != nil { + return nil, err + } + if _, err := cfg.Verify(ctx); err != nil { + return nil, err + } + if err := cfg.Save(path); err != nil { + return nil, err + } + return cfg, nil +} + +// needsURL reports whether the URL is still unset in any meaningful sense. +func needsURL(path string, cfg *client.Config) bool { + if _, err := os.Stat(path); err == nil { + return false // an existing file is the user's choice, default or not + } + return os.Getenv("LARD_URL") == "" +} + +func askURL(cfg *client.Config) error { + value := cfg.URL + if value == "http://localhost:7477" { + value = "" // don't pre-fill a guess the user probably doesn't want + } + // An explicit width is required, not cosmetic: bubbles' placeholder + // rendering sizes a buffer from the input width and panics when that + // width is unset. + form := huh.NewForm( + huh.NewGroup( + huh.NewInput(). + Title("Where does lard live?"). + Description("The base URL of your central lard server."). + Placeholder("https://lard.example.com"). + Value(&value). + Validate(validateURL), + ), + ).WithWidth(formWidth()) + if err := form.Run(); err != nil { + return err + } + cfg.URL = normalizeURL(value) + return nil +} + +// formWidth picks a readable form width that fits the terminal. +func formWidth() int { + const fallback = 64 + w, _, err := term.GetSize(int(os.Stdout.Fd())) + if err != nil || w <= 0 { + return fallback + } + return min(w-4, fallback) +} + +func validateURL(s string) error { + s = strings.TrimSpace(s) + if s == "" { + return errors.New("required") + } + u, err := url.Parse(normalizeURL(s)) + if err != nil { + return errors.New("not a valid URL") + } + if u.Host == "" { + return errors.New("needs a host, e.g. lard.example.com") + } + return nil +} + +// normalizeURL fills in a scheme and trims the trailing slash, so a user can +// type a bare hostname and get something that works. +func normalizeURL(s string) string { + s = strings.TrimSpace(s) + if s == "" { + return "" + } + if !strings.Contains(s, "://") { + // Assume TLS for a real host, plain HTTP for a local one. + if strings.HasPrefix(s, "localhost") || strings.HasPrefix(s, "127.0.0.1") { + s = "http://" + s + } else { + s = "https://" + s + } + } + return strings.TrimRight(s, "/") +} + +// authenticate obtains credentials for cfg.URL, choosing the lightest path +// that works: an explicit token, then existing credentials that still verify, +// then the browser flow, then a typed token. +func authenticate(ctx context.Context, cfg *client.Config, opts Options) error { + if opts.Token != "" { + cfg.Token = opts.Token + cfg.OAuth = nil + return nil + } + // Already have something that works? Don't make the user re-authorize. + if cfg.AuthMode() != "none" { + if _, err := cfg.Verify(ctx); err == nil { + return nil + } + } + // 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 + } + 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) + if err != nil { + return err + } + cfg.OAuth = &client.OAuthToken{ + AccessToken: tok.AccessToken, + RefreshToken: tok.RefreshToken, + Expiry: tok.Expiry, + CallbackPort: port, + } + 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 + } + + // No OAuth on offer, so the server wants a shared secret. + if !Interactive() { + return fmt.Errorf("%s needs a token and does not offer OAuth; pass --token", cfg.URL) + } + var token string + form := huh.NewForm( + huh.NewGroup( + huh.NewInput(). + Title("Access token"). + Description(fmt.Sprintf("%s does not offer a browser login.\nPaste the server's LARD_TOKEN.", cfg.URL)). + EchoMode(huh.EchoModePassword). + Value(&token). + Validate(func(s string) error { + if strings.TrimSpace(s) == "" { + return errors.New("required") + } + return nil + }), + ), + ).WithWidth(formWidth()) + if err := form.Run(); err != nil { + return err + } + cfg.Token = strings.TrimSpace(token) + cfg.OAuth = nil + return nil +} diff --git a/internal/setup/setup_test.go b/internal/setup/setup_test.go new file mode 100644 index 0000000..2977e40 --- /dev/null +++ b/internal/setup/setup_test.go @@ -0,0 +1,46 @@ +package setup + +import "testing" + +// A user types a bare hostname; the flow must produce something that works. +func TestNormalizeURLInfersScheme(t *testing.T) { + cases := map[string]string{ + "lard.example.com": "https://lard.example.com", + "https://lard.example.com/": "https://lard.example.com", + "http://lard.example.com": "http://lard.example.com", + " lard.example.com ": "https://lard.example.com", + "lard.example.com/": "https://lard.example.com", + // Local addresses almost never have TLS, so assume plain HTTP. + "localhost:7477": "http://localhost:7477", + "127.0.0.1:7477": "http://127.0.0.1:7477", + "": "", + } + for in, want := range cases { + if got := normalizeURL(in); got != want { + t.Errorf("normalizeURL(%q) = %q, want %q", in, got, want) + } + } +} + +func TestValidateURL(t *testing.T) { + good := []string{"lard.example.com", "https://lard.example.com", "localhost:7477"} + for _, s := range good { + if err := validateURL(s); err != nil { + t.Errorf("validateURL(%q) = %v, want nil", s, err) + } + } + bad := []string{"", " "} + for _, s := range bad { + if err := validateURL(s); err == nil { + t.Errorf("validateURL(%q) = nil, want an error", s) + } + } +} + +// The form must never be sized zero: bubbles' placeholder rendering allocates +// from the width and panics on a negative one. +func TestFormWidthIsPositive(t *testing.T) { + if w := formWidth(); w <= 0 { + t.Fatalf("formWidth() = %d, want > 0", w) + } +} diff --git a/internal/xdg/xdg.go b/internal/xdg/xdg.go new file mode 100644 index 0000000..eef742c --- /dev/null +++ b/internal/xdg/xdg.go @@ -0,0 +1,47 @@ +// Package xdg resolves lard's config directory. +// +// Go's os.UserConfigDir returns ~/Library/Application Support on macOS, which +// is right for GUI apps and wrong for a CLI: every other tool in a terminal +// keeps its dotfiles in ~/.config. This package follows the XDG basedir spec +// on every platform so the path is predictable and easy to type. +package xdg + +import ( + "os" + "path/filepath" +) + +// ConfigDir returns lard's config directory, honoring XDG_CONFIG_HOME. +func ConfigDir() string { + if d := os.Getenv("XDG_CONFIG_HOME"); d != "" { + return filepath.Join(d, "lard") + } + if home, err := os.UserHomeDir(); err == nil { + return filepath.Join(home, ".config", "lard") + } + // Last resort: the working directory, so nothing silently vanishes. + return "lard" +} + +// ConfigPath joins name onto the config directory. +func ConfigPath(name string) string { + return filepath.Join(ConfigDir(), name) +} + +// DataDir returns lard's data directory, honoring XDG_DATA_HOME. The database +// and memory files live here rather than next to config, since they are state +// a user backs up rather than settings a user edits. +func DataDir() string { + if d := os.Getenv("XDG_DATA_HOME"); d != "" { + return filepath.Join(d, "lard") + } + if home, err := os.UserHomeDir(); err == nil { + return filepath.Join(home, ".local", "share", "lard") + } + return "lard" +} + +// DataPath joins name onto the data directory. +func DataPath(name string) string { + return filepath.Join(DataDir(), name) +}