From 3ca0c71b249eec1fe9735ff6a8f1df40b5ace6bb Mon Sep 17 00:00:00 2001 From: dawn Date: Mon, 6 Jul 2026 19:26:57 +0300 Subject: [PATCH] spindle/mill: authenticate executors with per-executor tokens Registers executors in the db and resolves identity from the token hash in the Authorization header rather than the client-supplied Hello. Adds `spindle mill executor add/list/revoke` and seeds the dev SharedSecret as a bootstrap identity. --- cmd/spindle/main.go | 120 +++++++++ .../gen/spindle/agent/v1/spindle.agent.v1.rs | 103 ++++---- spindle/config/config.go | 3 + spindle/db/db.go | 7 + spindle/db/mill_tokens.go | 111 +++++++++ spindle/db/mill_tokens_test.go | 235 ++++++++++++++++++ spindle/mill/auth_test.go | 32 +++ spindle/mill/executor/executor.go | 11 - spindle/mill/executor/reserved.go | 4 +- spindle/mill/handler.go | 42 +++- spindle/mill/integration_test.go | 155 ++++++++++++ spindle/mill/mill.go | 28 ++- spindle/mill/mill_test.go | 33 +++ spindle/mill/proto/gen/mill.pb.go | 44 +--- spindle/mill/proto/spindle/mill/v1/mill.proto | 11 +- spindle/mill/session.go | 31 ++- spindle/mill/token.go | 29 +++ spindle/server.go | 3 + 18 files changed, 869 insertions(+), 133 deletions(-) create mode 100644 spindle/db/mill_tokens.go create mode 100644 spindle/db/mill_tokens_test.go create mode 100644 spindle/mill/token.go diff --git a/cmd/spindle/main.go b/cmd/spindle/main.go index c377a7ca..302e414f 100644 --- a/cmd/spindle/main.go +++ b/cmd/spindle/main.go @@ -2,12 +2,17 @@ package main import ( "context" + "fmt" "log/slog" "os" + "text/tabwriter" + "time" "github.com/urfave/cli/v3" tlog "tangled.org/core/log" "tangled.org/core/spindle" + "tangled.org/core/spindle/db" + "tangled.org/core/spindle/mill" ) func main() { @@ -16,6 +21,7 @@ func main() { Usage: "spindle continuous integration runner", Commands: []*cli.Command{ Command(), + millCommand(), }, DefaultCommand: "run", } @@ -41,3 +47,117 @@ func Command() *cli.Command { }, } } + +// millCommand groups mill-host administration. Executors allowed to join the +// mill are managed under `mill executor`. +func millCommand() *cli.Command { + dbFlag := &cli.StringFlag{ + Name: "db", + Usage: "path to the spindle sqlite db", + Value: "spindle.db", + Sources: cli.EnvVars("SPINDLE_SERVER_DB_PATH"), + } + openDB := func(ctx context.Context, cmd *cli.Command) (*db.DB, error) { + return db.Make(ctx, cmd.String("db")) + } + return &cli.Command{ + Name: "mill", + Usage: "mill host administration", + Commands: []*cli.Command{ + { + Name: "executor", + Usage: "manage executors allowed to join this mill", + Commands: []*cli.Command{ + { + Name: "add", + Usage: "register an executor and print its token", + ArgsUsage: "", + Flags: []cli.Flag{ + dbFlag, + &cli.DurationFlag{ + Name: "ttl", + Usage: "token lifetime (e.g. 720h); omit for no expiry", + }, + }, + Action: func(ctx context.Context, cmd *cli.Command) error { + name := cmd.Args().First() + if name == "" { + return fmt.Errorf("usage: spindle mill executor add ") + } + d, err := openDB(ctx, cmd) + if err != nil { + return err + } + token, err := mill.GenerateToken() + if err != nil { + return err + } + var expiresAt *time.Time + if ttl := cmd.Duration("ttl"); ttl > 0 { + exp := time.Now().UTC().Add(ttl) + expiresAt = &exp + } + if err := d.AddExecutorToken(name, mill.HashToken(token), expiresAt); err != nil { + return fmt.Errorf("registering executor %q: %w", name, err) + } + fmt.Println(token) + return nil + }, + }, + { + Name: "list", + Usage: "list registered executors", + Flags: []cli.Flag{dbFlag}, + Action: func(ctx context.Context, cmd *cli.Command) error { + d, err := openDB(ctx, cmd) + if err != nil { + return err + } + tokens, err := d.ListExecutorTokens() + if err != nil { + return err + } + w := tabwriter.NewWriter(os.Stdout, 0, 0, 3, ' ', 0) + fmt.Fprintln(w, "NAME\tCREATED\tEXPIRES") + for _, t := range tokens { + expires := "never" + if t.ExpiresAt != nil { + expires = t.ExpiresAt.Format(time.RFC3339) + if time.Now().After(*t.ExpiresAt) { + expires += " (expired)" + } + } + fmt.Fprintf(w, "%s\t%s\t%s\n", t.Name, t.CreatedAt, expires) + } + return w.Flush() + }, + }, + { + Name: "revoke", + Usage: "revoke an executor's token", + ArgsUsage: "", + Flags: []cli.Flag{dbFlag}, + Action: func(ctx context.Context, cmd *cli.Command) error { + name := cmd.Args().First() + if name == "" { + return fmt.Errorf("usage: spindle mill executor revoke ") + } + d, err := openDB(ctx, cmd) + if err != nil { + return err + } + ok, err := d.RevokeExecutorToken(name) + if err != nil { + return err + } + if !ok { + return fmt.Errorf("no such executor identity %q", name) + } + return nil + }, + }, + }, + }, + }, + } +} diff --git a/shuttle/src/gen/spindle/agent/v1/spindle.agent.v1.rs b/shuttle/src/gen/spindle/agent/v1/spindle.agent.v1.rs index a274c2a2..04a12079 100644 --- a/shuttle/src/gen/spindle/agent/v1/spindle.agent.v1.rs +++ b/shuttle/src/gen/spindle/agent/v1/spindle.agent.v1.rs @@ -2,146 +2,145 @@ // This file is @generated by prost-build. #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct Hello { - #[prost(uint32, tag="1")] + #[prost(uint32, tag = "1")] pub protocol_version: u32, - #[prost(string, tag="2")] + #[prost(string, tag = "2")] pub agent_version: ::prost::alloc::string::String, - #[prost(string, tag="3")] + #[prost(string, tag = "3")] pub boot_id: ::prost::alloc::string::String, - #[prost(string, tag="4")] + #[prost(string, tag = "4")] pub nix_version: ::prost::alloc::string::String, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct Init { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub job_id: ::prost::alloc::string::String, - #[prost(string, repeated, tag="2")] + #[prost(string, repeated, tag = "2")] pub cache_trusted_public_keys: ::prost::alloc::vec::Vec<::prost::alloc::string::String>, - #[prost(uint32, tag="3")] + #[prost(uint32, tag = "3")] pub cache_read_proxy_port: u32, - #[prost(uint32, tag="4")] + #[prost(uint32, tag = "4")] pub cache_upload_proxy_port: u32, - #[prost(uint32, tag="5")] + #[prost(uint32, tag = "5")] pub dns_proxy_port: u32, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct ExecStart { - #[prost(string, repeated, tag="1")] + #[prost(string, repeated, tag = "1")] pub argv: ::prost::alloc::vec::Vec<::prost::alloc::string::String>, - #[prost(string, repeated, tag="2")] + #[prost(string, repeated, tag = "2")] pub env: ::prost::alloc::vec::Vec<::prost::alloc::string::String>, - #[prost(string, tag="3")] + #[prost(string, tag = "3")] pub cwd: ::prost::alloc::string::String, - #[prost(string, tag="4")] + #[prost(string, tag = "4")] pub user: ::prost::alloc::string::String, - #[prost(uint32, tag="5")] + #[prost(uint32, tag = "5")] pub timeout_seconds: u32, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct ExecStdout { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub data: ::prost::alloc::string::String, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct ExecStderr { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub data: ::prost::alloc::string::String, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct ExecExit { - #[prost(int32, tag="1")] + #[prost(int32, tag = "1")] pub exit_code: i32, - #[prost(string, tag="2")] + #[prost(string, tag = "2")] pub error: ::prost::alloc::string::String, /// set when the guest killed the step on its own timeout timer, so the host /// can classify it as a timeout rather than inferring failure from exit_code. - #[prost(bool, tag="3")] + #[prost(bool, tag = "3")] pub timed_out: bool, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct ActivateConfig { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub config_key: ::prost::alloc::string::String, - #[prost(string, tag="2")] + #[prost(string, tag = "2")] pub base_config_hash: ::prost::alloc::string::String, - #[prost(string, tag="3")] + #[prost(string, tag = "3")] pub user_config: ::prost::alloc::string::String, - #[prost(string, tag="4")] + #[prost(string, tag = "4")] pub toplevel: ::prost::alloc::string::String, - #[prost(uint32, tag="5")] + #[prost(uint32, tag = "5")] pub timeout_seconds: u32, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct ActivateConfigResult { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub config_key: ::prost::alloc::string::String, - #[prost(string, tag="2")] + #[prost(string, tag = "2")] pub toplevel: ::prost::alloc::string::String, - #[prost(string, tag="3")] + #[prost(string, tag = "3")] pub error: ::prost::alloc::string::String, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct BuiltPaths { - #[prost(string, repeated, tag="1")] + #[prost(string, repeated, tag = "1")] pub paths: ::prost::alloc::vec::Vec<::prost::alloc::string::String>, - #[prost(string, tag="2")] + #[prost(string, tag = "2")] pub reason: ::prost::alloc::string::String, } #[derive(Clone, Copy, PartialEq, Eq, Hash, ::prost::Message)] pub struct CacheDrain { - #[prost(uint32, tag="1")] + #[prost(uint32, tag = "1")] pub timeout_seconds: u32, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct CacheDrainResult { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub error: ::prost::alloc::string::String, - #[prost(uint32, tag="2")] + #[prost(uint32, tag = "2")] pub cache_queued: u32, - #[prost(uint32, tag="3")] + #[prost(uint32, tag = "3")] pub cache_active: u32, - #[prost(uint32, tag="4")] + #[prost(uint32, tag = "4")] pub cache_uploaded: u32, - #[prost(uint32, tag="5")] + #[prost(uint32, tag = "5")] pub cache_failed: u32, } #[derive(Clone, Copy, PartialEq, Eq, Hash, ::prost::Message)] -pub struct Poweroff { -} +pub struct Poweroff {} #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct PoweroffResult { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub error: ::prost::alloc::string::String, } #[derive(Clone, PartialEq, Eq, Hash, ::prost::Message)] pub struct Message { - #[prost(string, tag="1")] + #[prost(string, tag = "1")] pub id: ::prost::alloc::string::String, - #[prost(message, optional, tag="2")] + #[prost(message, optional, tag = "2")] pub hello: ::core::option::Option, - #[prost(message, optional, tag="3")] + #[prost(message, optional, tag = "3")] pub init: ::core::option::Option, - #[prost(message, optional, tag="4")] + #[prost(message, optional, tag = "4")] pub exec_start: ::core::option::Option, - #[prost(message, optional, tag="5")] + #[prost(message, optional, tag = "5")] pub exec_stdout: ::core::option::Option, - #[prost(message, optional, tag="6")] + #[prost(message, optional, tag = "6")] pub exec_stderr: ::core::option::Option, - #[prost(message, optional, tag="7")] + #[prost(message, optional, tag = "7")] pub exec_exit: ::core::option::Option, - #[prost(message, optional, tag="8")] + #[prost(message, optional, tag = "8")] pub activate_config: ::core::option::Option, - #[prost(message, optional, tag="9")] + #[prost(message, optional, tag = "9")] pub activate_config_result: ::core::option::Option, - #[prost(message, optional, tag="10")] + #[prost(message, optional, tag = "10")] pub built_paths: ::core::option::Option, - #[prost(message, optional, tag="11")] + #[prost(message, optional, tag = "11")] pub cache_drain: ::core::option::Option, - #[prost(message, optional, tag="12")] + #[prost(message, optional, tag = "12")] pub cache_drain_result: ::core::option::Option, - #[prost(message, optional, tag="13")] + #[prost(message, optional, tag = "13")] pub poweroff: ::core::option::Option, - #[prost(message, optional, tag="14")] + #[prost(message, optional, tag = "14")] pub poweroff_result: ::core::option::Option, } // @@protoc_insertion_point(module) diff --git a/spindle/config/config.go b/spindle/config/config.go index bd4aa3e5..a13559e4 100644 --- a/spindle/config/config.go +++ b/spindle/config/config.go @@ -133,6 +133,9 @@ func (c *Config) validate() error { if c.Mill.URL == "" { return fmt.Errorf("SPINDLE_ROLE=executor requires SPINDLE_MILL_URL (the mill to dial)") } + if c.Mill.SharedSecret == "" { + return fmt.Errorf("SPINDLE_ROLE=executor requires SPINDLE_MILL_SHARED_SECRET (its executor token)") + } default: return fmt.Errorf("unknown SPINDLE_ROLE %q (want standalone, mill, or executor)", c.Role) } diff --git a/spindle/db/db.go b/spindle/db/db.go index 8a74d6d2..f4a05c69 100644 --- a/spindle/db/db.go +++ b/spindle/db/db.go @@ -123,6 +123,13 @@ func Make(ctx context.Context, dbPath string) (*DB, error) { foreign key (pipeline_id) references pipelines(id) on delete cascade ); + create table if not exists mill_executors ( + name text primary key, + token_hash text not null, + created_at text not null default (strftime('%Y-%m-%dT%H:%M:%SZ', 'now')), + expires_at text + ); + create table if not exists migrations ( id integer primary key autoincrement, name text unique diff --git a/spindle/db/mill_tokens.go b/spindle/db/mill_tokens.go new file mode 100644 index 00000000..03b04d8e --- /dev/null +++ b/spindle/db/mill_tokens.go @@ -0,0 +1,111 @@ +package db + +import ( + "database/sql" + "time" +) + +// ExecutorToken is a registered mill executor identity. The raw token is never +// stored; token_hash is a hash of it (see mill.HashToken). ExpiresAt is nil for +// a token that never expires. +type ExecutorToken struct { + Name string + CreatedAt string + ExpiresAt *time.Time +} + +// AddExecutorToken registers a new executor identity. expiresAt is optional (nil +// means the token never expires). It fails if the name is already taken. +func (d *DB) AddExecutorToken(name, tokenHash string, expiresAt *time.Time) error { + _, err := d.Exec( + `insert into mill_executors (name, token_hash, expires_at) values (?, ?, ?)`, + name, tokenHash, expiryArg(expiresAt), + ) + return err +} + +// UpsertExecutorToken registers or replaces an executor identity's token with no +// expiry. Used for the dev bootstrap seed; production registers per-executor +// tokens via AddExecutorToken. +func (d *DB) UpsertExecutorToken(name, tokenHash string) error { + _, err := d.Exec( + `insert into mill_executors (name, token_hash, expires_at) values (?, ?, null) + on conflict(name) do update set token_hash = excluded.token_hash, expires_at = null`, + name, tokenHash, + ) + return err +} + +// ResolveExecutorToken returns the executor name a token hash is registered +// under. ok is false when no identity matches or the token has expired. +func (d *DB) ResolveExecutorToken(tokenHash string) (string, bool, error) { + var name string + var expires sql.NullString + err := d.QueryRow( + `select name, expires_at from mill_executors where token_hash = ?`, tokenHash, + ).Scan(&name, &expires) + if err == sql.ErrNoRows { + return "", false, nil + } + if err != nil { + return "", false, err + } + if exp, ok := parseExpiry(expires); ok && time.Now().After(exp) { + return "", false, nil + } + return name, true, nil +} + +// RevokeExecutorToken removes an executor identity, returning whether a row was +// deleted. +func (d *DB) RevokeExecutorToken(name string) (bool, error) { + res, err := d.Exec(`delete from mill_executors where name = ?`, name) + if err != nil { + return false, err + } + n, err := res.RowsAffected() + return n > 0, err +} + +// ListExecutorTokens returns the registered executor identities ordered by name. +func (d *DB) ListExecutorTokens() ([]ExecutorToken, error) { + rows, err := d.Query(`select name, created_at, expires_at from mill_executors order by name`) + if err != nil { + return nil, err + } + defer rows.Close() + + var out []ExecutorToken + for rows.Next() { + var t ExecutorToken + var expires sql.NullString + if err := rows.Scan(&t.Name, &t.CreatedAt, &expires); err != nil { + return nil, err + } + if exp, ok := parseExpiry(expires); ok { + t.ExpiresAt = &exp + } + out = append(out, t) + } + return out, rows.Err() +} + +// expiryArg renders an optional expiry as a sql argument (nil -> NULL). +func expiryArg(t *time.Time) any { + if t == nil { + return nil + } + return t.UTC().Format(time.RFC3339) +} + +// parseExpiry decodes a stored expiry column; ok is false for NULL or unparseable. +func parseExpiry(s sql.NullString) (time.Time, bool) { + if !s.Valid || s.String == "" { + return time.Time{}, false + } + t, err := time.Parse(time.RFC3339, s.String) + if err != nil { + return time.Time{}, false + } + return t, true +} diff --git a/spindle/db/mill_tokens_test.go b/spindle/db/mill_tokens_test.go new file mode 100644 index 00000000..4c6689b7 --- /dev/null +++ b/spindle/db/mill_tokens_test.go @@ -0,0 +1,235 @@ +package db + +import ( + "testing" + "time" +) + +func TestAddExecutorTokenRejectsDuplicateName(t *testing.T) { + d := newTestDB(t) + + if err := d.AddExecutorToken("exec-1", "hash-a", nil); err != nil { + t.Fatalf("AddExecutorToken: %v", err) + } + if err := d.AddExecutorToken("exec-1", "hash-b", nil); err == nil { + t.Fatal("AddExecutorToken re-registered an existing name; a duplicate must be rejected") + } + + name, ok, err := d.ResolveExecutorToken("hash-a") + if err != nil { + t.Fatalf("ResolveExecutorToken(hash-a): %v", err) + } + if !ok || name != "exec-1" { + t.Fatalf("ResolveExecutorToken(hash-a) = (%q, %v), want (exec-1, true)", name, ok) + } + if _, ok, _ := d.ResolveExecutorToken("hash-b"); ok { + t.Fatal("rejected duplicate's token resolved; the failed insert leaked a credential") + } +} + +func TestResolveExecutorTokenMissAndHit(t *testing.T) { + d := newTestDB(t) + + name, ok, err := d.ResolveExecutorToken("no-such-hash") + if err != nil { + t.Fatalf("ResolveExecutorToken(miss): %v", err) + } + if ok || name != "" { + t.Fatalf("ResolveExecutorToken(miss) = (%q, %v), want (\"\", false)", name, ok) + } + + if err := d.AddExecutorToken("exec-1", "hash-1", nil); err != nil { + t.Fatalf("AddExecutorToken: %v", err) + } + + name, ok, err = d.ResolveExecutorToken("hash-1") + if err != nil { + t.Fatalf("ResolveExecutorToken(hit): %v", err) + } + if !ok || name != "exec-1" { + t.Fatalf("ResolveExecutorToken(hash-1) = (%q, %v), want (exec-1, true)", name, ok) + } + + if _, ok, _ := d.ResolveExecutorToken("hash-unregistered"); ok { + t.Fatal("ResolveExecutorToken matched an unregistered hash") + } +} + +func TestUpsertExecutorTokenRotatesHash(t *testing.T) { + d := newTestDB(t) + + if err := d.AddExecutorToken("exec-1", "old-hash", nil); err != nil { + t.Fatalf("AddExecutorToken: %v", err) + } + if err := d.UpsertExecutorToken("exec-1", "new-hash"); err != nil { + t.Fatalf("UpsertExecutorToken(rotate): %v", err) + } + + name, ok, err := d.ResolveExecutorToken("new-hash") + if err != nil { + t.Fatalf("ResolveExecutorToken(new-hash): %v", err) + } + if !ok || name != "exec-1" { + t.Fatalf("ResolveExecutorToken(new-hash) = (%q, %v), want (exec-1, true)", name, ok) + } + if _, ok, _ := d.ResolveExecutorToken("old-hash"); ok { + t.Fatal("rotated-out token still resolves; UpsertExecutorToken did not replace the hash") + } + + if err := d.UpsertExecutorToken("exec-2", "hash-2"); err != nil { + t.Fatalf("UpsertExecutorToken(insert): %v", err) + } + name, ok, err = d.ResolveExecutorToken("hash-2") + if err != nil { + t.Fatalf("ResolveExecutorToken(hash-2): %v", err) + } + if !ok || name != "exec-2" { + t.Fatalf("ResolveExecutorToken(hash-2) = (%q, %v), want (exec-2, true)", name, ok) + } +} + +func TestRevokeExecutorToken(t *testing.T) { + d := newTestDB(t) + + if err := d.AddExecutorToken("exec-1", "hash-1", nil); err != nil { + t.Fatalf("AddExecutorToken: %v", err) + } + + deleted, err := d.RevokeExecutorToken("exec-1") + if err != nil { + t.Fatalf("RevokeExecutorToken: %v", err) + } + if !deleted { + t.Fatal("RevokeExecutorToken reported no deletion for an existing identity") + } + + if _, ok, _ := d.ResolveExecutorToken("hash-1"); ok { + t.Fatal("revoked token still resolves; revocation is not enforced") + } + + if deleted, err := d.RevokeExecutorToken("exec-1"); err != nil || deleted { + t.Fatalf("RevokeExecutorToken(already-gone) = (%v, %v), want (false, nil)", deleted, err) + } + + if deleted, err := d.RevokeExecutorToken("ghost"); err != nil || deleted { + t.Fatalf("RevokeExecutorToken(unknown) = (%v, %v), want (false, nil)", deleted, err) + } +} + +func TestListExecutorTokens(t *testing.T) { + d := newTestDB(t) + + tokens, err := d.ListExecutorTokens() + if err != nil { + t.Fatalf("ListExecutorTokens(empty): %v", err) + } + if len(tokens) != 0 { + t.Fatalf("ListExecutorTokens on empty table = %d rows, want 0", len(tokens)) + } + + // insert out of alphabetical order to prove the ordering is the query's. + for _, name := range []string{"charlie", "alice", "bob"} { + if err := d.AddExecutorToken(name, "hash-"+name, nil); err != nil { + t.Fatalf("AddExecutorToken(%s): %v", name, err) + } + } + + tokens, err = d.ListExecutorTokens() + if err != nil { + t.Fatalf("ListExecutorTokens: %v", err) + } + want := []string{"alice", "bob", "charlie"} + if len(tokens) != len(want) { + t.Fatalf("ListExecutorTokens = %d rows, want %d", len(tokens), len(want)) + } + for i := range want { + if tokens[i].Name != want[i] { + t.Fatalf("ListExecutorTokens[%d].Name = %q, want %q (ordered by name)", i, tokens[i].Name, want[i]) + } + } +} + +// expired tokens must fail closed (ok=false, no error). +func TestResolveExecutorTokenExpiry(t *testing.T) { + cases := []struct { + name string + offset time.Duration + noExpiry bool + wantOK bool + }{ + {"future expiry resolves", time.Hour, false, true}, + {"past expiry fails closed", -time.Hour, false, false}, + {"nil expiry never expires", 0, true, true}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + d := newTestDB(t) + + var expires *time.Time + if !tc.noExpiry { + exp := time.Now().Add(tc.offset) + expires = &exp + } + if err := d.AddExecutorToken("exec-1", "hash-1", expires); err != nil { + t.Fatalf("AddExecutorToken: %v", err) + } + + name, ok, err := d.ResolveExecutorToken("hash-1") + if err != nil { + t.Fatalf("ResolveExecutorToken: %v", err) + } + if ok != tc.wantOK { + t.Fatalf("ResolveExecutorToken ok = %v, want %v", ok, tc.wantOK) + } + if tc.wantOK && name != "exec-1" { + t.Fatalf("ResolveExecutorToken name = %q, want exec-1", name) + } + if !tc.wantOK && name != "" { + t.Fatalf("ResolveExecutorToken name = %q, want \"\" when failing closed", name) + } + }) + } +} + +// expiry storage is RFC3339 second-precision UTC, so round-trips compare to +// the second. +func TestListExecutorTokensSurfacesExpiry(t *testing.T) { + d := newTestDB(t) + + exp := time.Now().Add(24 * time.Hour) + if err := d.AddExecutorToken("expiring", "hash-exp", &exp); err != nil { + t.Fatalf("AddExecutorToken(expiring): %v", err) + } + if err := d.AddExecutorToken("forever", "hash-forever", nil); err != nil { + t.Fatalf("AddExecutorToken(forever): %v", err) + } + + tokens, err := d.ListExecutorTokens() + if err != nil { + t.Fatalf("ListExecutorTokens: %v", err) + } + + got := make(map[string]*time.Time, len(tokens)) + for _, tok := range tokens { + got[tok.Name] = tok.ExpiresAt + } + + e, present := got["forever"] + if !present { + t.Fatal("ListExecutorTokens omitted the non-expiring identity") + } + if e != nil { + t.Fatalf("forever.ExpiresAt = %v, want nil (never expires)", e) + } + + e, present = got["expiring"] + if !present { + t.Fatal("ListExecutorTokens omitted the expiring identity") + } + if e == nil { + t.Fatal("expiring.ExpiresAt = nil, want the stored expiry") + } + if e.Unix() != exp.Unix() { + t.Fatalf("expiring.ExpiresAt = %d (unix), want %d", e.Unix(), exp.Unix()) + } +} diff --git a/spindle/mill/auth_test.go b/spindle/mill/auth_test.go index 62182514..8ccab52d 100644 --- a/spindle/mill/auth_test.go +++ b/spindle/mill/auth_test.go @@ -25,6 +25,38 @@ func nopEncoder() scriptedEncoder { return scriptedEncoder(func(*millproto.Message) error { return nil }) } +func TestHashToken(t *testing.T) { + const raw = "super-secret-executor-token" + + if HashToken(raw) != HashToken(raw) { + t.Fatal("HashToken is not deterministic; the same token would stop authenticating") + } + if HashToken("token-a") == HashToken("token-b") { + t.Fatal("HashToken collided two distinct tokens") + } + if HashToken(raw) == raw { + t.Fatal("HashToken returned the raw token; a hash leak would expose a usable credential") + } +} + +func TestGenerateTokenDistinct(t *testing.T) { + const n = 100 + seen := make(map[string]struct{}, n) + for i := range n { + tok, err := GenerateToken() + if err != nil { + t.Fatalf("GenerateToken: %v", err) + } + if tok == "" { + t.Fatalf("GenerateToken returned an empty token on call %d", i) + } + if _, dup := seen[tok]; dup { + t.Fatalf("GenerateToken repeated a token after %d calls: %q", i, tok) + } + seen[tok] = struct{}{} + } +} + func TestAttachSessionRejectsSecondLiveSession(t *testing.T) { l := discardLogger() m := New(l, Config{ReconnectGrace: time.Minute}) diff --git a/spindle/mill/executor/executor.go b/spindle/mill/executor/executor.go index ac561136..429785bc 100644 --- a/spindle/mill/executor/executor.go +++ b/spindle/mill/executor/executor.go @@ -152,11 +152,8 @@ func (e *Executor) runSession(ctx context.Context) error { hello := &millproto.Message{Hello: &millv1.Hello{ ProtocolVersion: millproto.ProtocolVersion, - NodeId: e.nodeID, - Engines: e.engineNames(), Arch: runtime.GOARCH, Labels: e.labels, - LastOffset: e.relay.lastOffset(), }} if err := enc.Encode(hello); err != nil { return fmt.Errorf("send hello: %w", err) @@ -527,14 +524,6 @@ func (e *Executor) Drain() { e.pushSnapshot() } -func (e *Executor) engineNames() []string { - names := make([]string, 0, len(e.engines)) - for name := range e.engines { - names = append(names, name) - } - return names -} - func ttlDuration(secs uint32) time.Duration { if secs == 0 { return defaultReservationTTL diff --git a/spindle/mill/executor/reserved.go b/spindle/mill/executor/reserved.go index 30a6b669..c79dbcac 100644 --- a/spindle/mill/executor/reserved.go +++ b/spindle/mill/executor/reserved.go @@ -17,8 +17,8 @@ import ( // This is the only change to the execution path on an executor. type reservedEngine struct { models.Engine - slot engine.WorkflowSlot - once sync.Once + slot engine.WorkflowSlot + once sync.Once } // newReservedEngine returns a wrapper around inner that hands back slot exactly diff --git a/spindle/mill/handler.go b/spindle/mill/handler.go index fea599cd..f9c6f149 100644 --- a/spindle/mill/handler.go +++ b/spindle/mill/handler.go @@ -2,6 +2,7 @@ package mill import ( "net/http" + "strings" "github.com/gorilla/websocket" @@ -18,11 +19,10 @@ var upgrader = websocket.Upgrader{ // secret in the Authorization header, checked before the upgrade so a bad token // never opens a socket. func (m *Mill) HandleExecutorConn(w http.ResponseWriter, r *http.Request) { - if m.cfg.SharedSecret != "" { - if r.Header.Get("Authorization") != "Bearer "+m.cfg.SharedSecret { - http.Error(w, "unauthorized", http.StatusUnauthorized) - return - } + name, ok := m.authenticate(r) + if !ok { + http.Error(w, "unauthorized", http.StatusUnauthorized) + return } conn, err := upgrader.Upgrade(w, r, nil) @@ -52,13 +52,17 @@ func (m *Mill) HandleExecutorConn(w http.ResponseWriter, r *http.Request) { return } - sess := newSession(h.GetNodeId(), enc, m.l) + // identity comes from the authenticated token, never the client-supplied + // Hello node id, so a valid token can't impersonate another executor. + sess := newSession(name, enc, m.l) + sess.closeTransport = conn.Close + sess.labels = h.GetLabels() resume, ok := m.attachSession(sess) if !ok { - m.l.Warn("rejecting duplicate live executor session", "node", sess.nodeID) + m.l.Warn("rejecting duplicate live executor session", "node", name) return } - m.l.Info("executor connected", "node", sess.nodeID, "engines", h.GetEngines(), "arch", h.GetArch(), "resume", resume) + m.l.Info("executor connected", "node", sess.nodeID, "arch", h.GetArch(), "labels", h.GetLabels(), "resume", resume) if err := sess.send(&millproto.Message{Resume: &millv1.Resume{AckOffset: resume}}); err != nil { m.l.Error("fleet send resume failed", "err", err) @@ -70,3 +74,25 @@ func (m *Mill) HandleExecutorConn(w http.ResponseWriter, r *http.Request) { sess.readLoop(m, dec) m.detachSession(sess) } + +// authenticate resolves the executor identity from the pre-shared token in the +// Authorization header before the upgrade, so a bad token never opens a socket. +// The token is looked up by hash; an unknown or missing token is rejected +// (fail closed). +func (m *Mill) authenticate(r *http.Request) (string, bool) { + const prefix = "Bearer " + h := r.Header.Get("Authorization") + if !strings.HasPrefix(h, prefix) { + return "", false + } + token := strings.TrimPrefix(h, prefix) + if token == "" || m.db == nil { + return "", false + } + name, ok, err := m.db.ResolveExecutorToken(HashToken(token)) + if err != nil { + m.l.Error("executor token lookup failed", "err", err) + return "", false + } + return name, ok +} diff --git a/spindle/mill/integration_test.go b/spindle/mill/integration_test.go index 51890e03..8413d6a4 100644 --- a/spindle/mill/integration_test.go +++ b/spindle/mill/integration_test.go @@ -34,6 +34,9 @@ func TestEndToEndDummyJob(t *testing.T) { bn := notifier.New() mill := New(l, Config{LogDir: millDir, ReconnectGrace: time.Minute, BidTimeout: 2 * time.Second}) mill.Attach(bdb, &bn) + if err := bdb.AddExecutorToken("exec-1", HashToken("test-token"), nil); err != nil { + t.Fatalf("register executor token: %v", err) + } srv := httptest.NewServer(http.HandlerFunc(mill.HandleExecutorConn)) defer srv.Close() @@ -50,6 +53,7 @@ func TestEndToEndDummyJob(t *testing.T) { cfg.Server.Hostname = "exec-1" cfg.Mill.URL = wsURL cfg.Mill.Seats = 2 + cfg.Mill.SharedSecret = "test-token" engines := map[string]models.Engine{"dummy": dummy.New(l)} exec := executor.New(cfg, engines, edb, &en, l) @@ -84,6 +88,157 @@ func TestEndToEndDummyJob(t *testing.T) { } } +func TestExecutorConfiguredLabelsAreStoredOnSession(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + l := slog.New(slog.NewTextHandler(io.Discard, nil)) + + millDir := t.TempDir() + bdb, err := db.Make(ctx, filepath.Join(millDir, "mill.db")) + if err != nil { + t.Fatalf("mill db: %v", err) + } + bn := notifier.New() + mill := New(l, Config{LogDir: millDir, ReconnectGrace: time.Minute, BidTimeout: 2 * time.Second}) + mill.Attach(bdb, &bn) + if err := bdb.AddExecutorToken("exec-labels", HashToken("test-token"), nil); err != nil { + t.Fatalf("register executor token: %v", err) + } + + srv := httptest.NewServer(http.HandlerFunc(mill.HandleExecutorConn)) + defer srv.Close() + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + + execDir := t.TempDir() + edb, err := db.Make(ctx, filepath.Join(execDir, "exec.db")) + if err != nil { + t.Fatalf("exec db: %v", err) + } + en := notifier.New() + cfg := &config.Config{} + cfg.Server.LogDir = execDir + cfg.Server.Hostname = "exec-labels" + cfg.Mill.URL = wsURL + cfg.Mill.Seats = 2 + cfg.Mill.SharedSecret = "test-token" + cfg.Mill.Labels = []string{"linux", "arm64", "gpu"} + + engines := map[string]models.Engine{"dummy": dummy.New(l)} + exec := executor.New(cfg, engines, edb, &en, l) + go exec.Connect(ctx) + + if !waitForSessionLabels(t, mill, "exec-labels", []string{"linux", "arm64", "gpu"}) { + t.Fatal("mill session never stored executor labels from hello") + } +} + +func TestEndToEndDummyJobUsesRequiredLabelsAcrossExecutors(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + l := slog.New(slog.NewTextHandler(io.Discard, nil)) + + millDir := t.TempDir() + bdb, err := db.Make(ctx, filepath.Join(millDir, "mill.db")) + if err != nil { + t.Fatalf("mill db: %v", err) + } + bn := notifier.New() + mill := New(l, Config{LogDir: millDir, ReconnectGrace: time.Minute, BidTimeout: 2 * time.Second}) + mill.Attach(bdb, &bn) + if err := bdb.AddExecutorToken("exec-x86", HashToken("token-x86"), nil); err != nil { + t.Fatalf("register x86 executor token: %v", err) + } + if err := bdb.AddExecutorToken("exec-arm", HashToken("token-arm"), nil); err != nil { + t.Fatalf("register arm executor token: %v", err) + } + + srv := httptest.NewServer(http.HandlerFunc(mill.HandleExecutorConn)) + defer srv.Close() + wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + + startExecutor := func(name, token string, labels []string) { + t.Helper() + execDir := t.TempDir() + edb, err := db.Make(ctx, filepath.Join(execDir, "exec.db")) + if err != nil { + t.Fatalf("%s exec db: %v", name, err) + } + en := notifier.New() + cfg := &config.Config{} + cfg.Server.LogDir = execDir + cfg.Server.Hostname = name + cfg.Mill.URL = wsURL + cfg.Mill.Seats = 1 + cfg.Mill.SharedSecret = token + cfg.Mill.Labels = labels + + engines := map[string]models.Engine{"dummy": dummy.New(l)} + exec := executor.New(cfg, engines, edb, &en, l) + go exec.Connect(ctx) + } + startExecutor("exec-x86", "token-x86", []string{"linux/amd64", "kvm"}) + startExecutor("exec-arm", "token-arm", []string{"linux/arm64", "kvm"}) + + if !waitForSessionLabels(t, mill, "exec-x86", []string{"linux/amd64", "kvm"}) { + t.Fatal("x86 executor did not connect with labels") + } + if !waitForSessionLabels(t, mill, "exec-arm", []string{"linux/arm64", "kvm"}) { + t.Fatal("arm executor did not connect with labels") + } + + be := NewEngine("dummy", mill) + twf := tangled.Pipeline_Workflow{ + Name: "build-arm", + RunsOn: []string{"linux/arm64"}, + Raw: "steps:\n - name: hello\n command: echo hi\n", + } + wf, err := be.InitWorkflow(twf, tangled.Pipeline{}) + if err != nil { + t.Fatalf("InitWorkflow: %v", err) + } + wid := models.WorkflowId{PipelineId: models.PipelineId{Knot: "knot.test", Rkey: "rkey-arm"}, Name: "build-arm"} + + placeCtx, placeCancel := context.WithTimeout(ctx, 10*time.Second) + defer placeCancel() + + slot, err := mill.place(placeCtx, "dummy", wid, wf) + if err != nil { + t.Fatalf("place: %v", err) + } + defer slot.Release() + + lease := wf.Data.(*millWorkflowState).Lease + if lease == nil { + t.Fatal("place did not attach lease to workflow state") + } + if lease.nodeID != "exec-arm" { + t.Fatalf("placed on %q, want exec-arm", lease.nodeID) + } + + if err := mill.commitAndWait(placeCtx, wf, nil); err != nil { + t.Fatalf("commitAndWait: %v, want success", err) + } +} + +func waitForSessionLabels(t *testing.T, m *Mill, nodeID string, want []string) bool { + t.Helper() + deadline := time.Now().Add(3 * time.Second) + for time.Now().Before(deadline) { + m.mu.Lock() + sess := m.sessions[nodeID] + var got []string + if sess != nil { + got = append([]string(nil), sess.labels...) + } + m.mu.Unlock() + if sameStringMultiset(got, want) { + return true + } + time.Sleep(10 * time.Millisecond) + } + return false +} + func waitForStatus(t *testing.T, d *db.DB, wid models.WorkflowId, want string) bool { t.Helper() deadline := time.Now().Add(3 * time.Second) diff --git a/spindle/mill/mill.go b/spindle/mill/mill.go index da15622a..321cce3f 100644 --- a/spindle/mill/mill.go +++ b/spindle/mill/mill.go @@ -90,6 +90,15 @@ func (m *Mill) Attach(d *db.DB, n *notifier.Notifier) { m.mu.Unlock() } +// dev bootstrap: shared-secret identity, single-token fleet. production +// should register per-executor tokens instead. +func (m *Mill) SeedBootstrapToken() error { + if m.cfg.SharedSecret == "" || m.db == nil { + return nil + } + return m.db.UpsertExecutorToken(bootstrapTokenName, HashToken(m.cfg.SharedSecret)) +} + func (m *Mill) nextLeaseID() string { m.mu.Lock() m.leaseSeq++ @@ -122,17 +131,20 @@ func (m *Mill) attachSession(sess *millSession) (uint64, bool) { defer m.mu.Unlock() if old := m.sessions[sess.nodeID]; old != nil { - if !old.disconnected { + if !old.disconnected && time.Since(old.lastSeen) <= m.cfg.ReconnectGrace { // reject a second live session for the same identity so a valid token // can't hijack an in-flight executor. return 0, false } - // the old session is in its reconnect grace window: adopt its leases. if old.graceTimer != nil { old.graceTimer.Stop() } old.close() - m.l.Info("executor reconnected", "node", sess.nodeID) + if old.disconnected { + m.l.Info("executor reconnected", "node", sess.nodeID) + } else { + m.l.Warn("replacing silent executor session", "node", sess.nodeID) + } } m.sessions[sess.nodeID] = sess // wake commit retries waiting out this node's reconnect grace. @@ -140,6 +152,16 @@ func (m *Mill) attachSession(sess *millSession) (uint64, bool) { return m.nodeOffset[sess.nodeID], true } +func (m *Mill) touchSession(sess *millSession) bool { + m.mu.Lock() + defer m.mu.Unlock() + if m.sessions[sess.nodeID] != sess || sess.disconnected { + return false + } + sess.lastSeen = time.Now() + return true +} + func (m *Mill) detachSession(sess *millSession) { m.mu.Lock() if m.sessions[sess.nodeID] != sess { diff --git a/spindle/mill/mill_test.go b/spindle/mill/mill_test.go index 2a4f8130..61c37655 100644 --- a/spindle/mill/mill_test.go +++ b/spindle/mill/mill_test.go @@ -495,3 +495,36 @@ func TestFallbackBidGetsFullTimeout(t *testing.T) { t.Fatalf("bid winner = %+v, want fallback after silent incumbent times out", lease) } } + +func TestAttachSessionReplacesSilentIncumbentButRejectsActiveDuplicate(t *testing.T) { + m := New(discardLogger(), Config{ReconnectGrace: time.Minute}) + old := newSession("node-1", nopEncoder(), discardLogger()) + transportClosed := make(chan struct{}) + old.closeTransport = func() error { + close(transportClosed) + return nil + } + if _, ok := m.attachSession(old); !ok { + t.Fatal("first attach rejected") + } + m.mu.Lock() + old.lastSeen = time.Now().Add(-2 * m.cfg.ReconnectGrace) + m.mu.Unlock() + + replacement := newSession("node-1", nopEncoder(), discardLogger()) + if _, ok := m.attachSession(replacement); !ok { + t.Fatal("silent incumbent blocked authenticated replacement") + } + select { + case <-transportClosed: + default: + t.Fatal("replacing a silent incumbent did not close its transport") + } + m.mu.Lock() + replacement.lastSeen = time.Now().Add(-2 * m.cfg.ReconnectGrace) + m.mu.Unlock() + replacement.dispatch(m, &millproto.Message{NodeSnapshot: &millv1.NodeSnapshot{NodeId: "node-1"}}) + if _, ok := m.attachSession(newSession("node-1", nopEncoder(), discardLogger())); ok { + t.Fatal("active replacement did not reject a duplicate session") + } +} diff --git a/spindle/mill/proto/gen/mill.pb.go b/spindle/mill/proto/gen/mill.pb.go index ef072acd..9cbf7223 100644 --- a/spindle/mill/proto/gen/mill.pb.go +++ b/spindle/mill/proto/gen/mill.pb.go @@ -76,17 +76,10 @@ func (RejectClass) EnumDescriptor() ([]byte, []int) { type Hello struct { state protoimpl.MessageState `protogen:"open.v1"` ProtocolVersion uint32 `protobuf:"varint,1,opt,name=protocol_version,json=protocolVersion,proto3" json:"protocol_version,omitempty"` - // stable across reconnects (e.g. hostname / DID); the mill keys sessions on - // this so a brief blip reattaches the same node rather than creating a new one. - NodeId string `protobuf:"bytes,2,opt,name=node_id,json=nodeId,proto3" json:"node_id,omitempty"` - // engine names this node can run ("microvm", "nixery"). - Engines []string `protobuf:"bytes,3,rep,name=engines,proto3" json:"engines,omitempty"` // GOARCH of the node, retained as an informational trait for logs. - Arch string `protobuf:"bytes,4,opt,name=arch,proto3" json:"arch,omitempty"` - // the highest relay offset the executor believes it has sent; a resume hint. - LastOffset uint64 `protobuf:"varint,5,opt,name=last_offset,json=lastOffset,proto3" json:"last_offset,omitempty"` + Arch string `protobuf:"bytes,2,opt,name=arch,proto3" json:"arch,omitempty"` // opaque operator-defined labels used for runs_on matching. - Labels []string `protobuf:"bytes,6,rep,name=labels,proto3" json:"labels,omitempty"` + Labels []string `protobuf:"bytes,3,rep,name=labels,proto3" json:"labels,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -128,20 +121,6 @@ func (x *Hello) GetProtocolVersion() uint32 { return 0 } -func (x *Hello) GetNodeId() string { - if x != nil { - return x.NodeId - } - return "" -} - -func (x *Hello) GetEngines() []string { - if x != nil { - return x.Engines - } - return nil -} - func (x *Hello) GetArch() string { if x != nil { return x.Arch @@ -149,13 +128,6 @@ func (x *Hello) GetArch() string { return "" } -func (x *Hello) GetLastOffset() uint64 { - if x != nil { - return x.LastOffset - } - return 0 -} - func (x *Hello) GetLabels() []string { if x != nil { return x.Labels @@ -1175,15 +1147,11 @@ var File_spindle_mill_v1_mill_proto protoreflect.FileDescriptor const file_spindle_mill_v1_mill_proto_rawDesc = "" + "\n" + - "\x1aspindle/mill/v1/mill.proto\x12\x0fspindle.mill.v1\x1a\x1bbuf/validate/validate.proto\"\xbb\x01\n" + + "\x1aspindle/mill/v1/mill.proto\x12\x0fspindle.mill.v1\x1a\x1bbuf/validate/validate.proto\"^\n" + "\x05Hello\x12)\n" + - "\x10protocol_version\x18\x01 \x01(\rR\x0fprotocolVersion\x12 \n" + - "\anode_id\x18\x02 \x01(\tB\a\xbaH\x04r\x02\x10\x01R\x06nodeId\x12\x18\n" + - "\aengines\x18\x03 \x03(\tR\aengines\x12\x12\n" + - "\x04arch\x18\x04 \x01(\tR\x04arch\x12\x1f\n" + - "\vlast_offset\x18\x05 \x01(\x04R\n" + - "lastOffset\x12\x16\n" + - "\x06labels\x18\x06 \x03(\tR\x06labels\"'\n" + + "\x10protocol_version\x18\x01 \x01(\rR\x0fprotocolVersion\x12\x12\n" + + "\x04arch\x18\x02 \x01(\tR\x04arch\x12\x16\n" + + "\x06labels\x18\x03 \x03(\tR\x06labels\"'\n" + "\x06Resume\x12\x1d\n" + "\n" + "ack_offset\x18\x01 \x01(\x04R\tackOffset\"\xae\x01\n" + diff --git a/spindle/mill/proto/spindle/mill/v1/mill.proto b/spindle/mill/proto/spindle/mill/v1/mill.proto index 45762b95..330ef89c 100644 --- a/spindle/mill/proto/spindle/mill/v1/mill.proto +++ b/spindle/mill/proto/spindle/mill/v1/mill.proto @@ -10,17 +10,10 @@ option go_package = "tangled.org/core/spindle/mill/proto/gen;millv1"; // carries the static traits of the node plus a resume hint. message Hello { uint32 protocol_version = 1; - // stable across reconnects (e.g. hostname / DID); the mill keys sessions on - // this so a brief blip reattaches the same node rather than creating a new one. - string node_id = 2 [(buf.validate.field).string.min_len = 1]; - // engine names this node can run ("microvm", "nixery"). - repeated string engines = 3; // GOARCH of the node, retained as an informational trait for logs. - string arch = 4; - // the highest relay offset the executor believes it has sent; a resume hint. - uint64 last_offset = 5; + string arch = 2; // opaque operator-defined labels used for runs_on matching. - repeated string labels = 6; + repeated string labels = 3; } // Resume is the mill's reply to Hello. The executor replays every buffered diff --git a/spindle/mill/session.go b/spindle/mill/session.go index 39c568c3..710575ed 100644 --- a/spindle/mill/session.go +++ b/spindle/mill/session.go @@ -18,16 +18,18 @@ var errSessionClosed = errors.New("mill: executor session closed") // it, so a single reader goroutine demuxes by message type and correlates // request/response by lease id. We never hold a lock across a decode. type millSession struct { - nodeID string - labels []string - enc messageEncoder - l *slog.Logger + nodeID string + labels []string + enc messageEncoder + l *slog.Logger + closeTransport func() error // snapshot, disconnected and graceTimer are guarded by Mill.mu, not the // session mutex below (the fleet ranks across sessions under its own lock). snapshot *millv1.NodeSnapshot disconnected bool graceTimer *time.Timer + lastSeen time.Time mu sync.Mutex pending map[string]chan *millproto.Message // lease_id -> response waiter @@ -42,11 +44,12 @@ type messageEncoder interface { func newSession(nodeID string, enc messageEncoder, l *slog.Logger) *millSession { return &millSession{ - nodeID: nodeID, - enc: enc, - l: l, - pending: make(map[string]chan *millproto.Message), - closed: make(chan struct{}), + nodeID: nodeID, + enc: enc, + l: l, + pending: make(map[string]chan *millproto.Message), + closed: make(chan struct{}), + lastSeen: time.Now(), } } @@ -55,7 +58,12 @@ func (s *millSession) send(msg *millproto.Message) error { } func (s *millSession) close() { - s.closeOnce.Do(func() { close(s.closed) }) + s.closeOnce.Do(func() { + close(s.closed) + if s.closeTransport != nil { + _ = s.closeTransport() + } + }) } // await registers a one-shot waiter for the next response addressed to leaseID, @@ -124,6 +132,9 @@ func (s *millSession) readLoop(m *Mill, dec *millproto.Decoder) { } func (s *millSession) dispatch(m *Mill, msg *millproto.Message) { + if !m.touchSession(s) { + return + } switch { case msg.GetNodeSnapshot() != nil: m.onSnapshot(s, msg.GetNodeSnapshot()) diff --git a/spindle/mill/token.go b/spindle/mill/token.go new file mode 100644 index 00000000..27326515 --- /dev/null +++ b/spindle/mill/token.go @@ -0,0 +1,29 @@ +package mill + +import ( + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/hex" +) + +// bootstrapTokenName is the identity the mill seeds cfg.SharedSecret under so a +// single-secret dev fleet works without minting. Production registers per-executor +// identities with `spindle mill executor add`. +const bootstrapTokenName = "dev-bootstrap" + +func GenerateToken() (string, error) { + var b [32]byte + if _, err := rand.Read(b[:]); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(b[:]), nil +} + +// HashToken hashes a raw token for storage/lookup. The mill only ever persists +// and compares hashes, so a DB leak never exposes a usable token, and lookup by +// hash avoids a secret-dependent comparison. +func HashToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} diff --git a/spindle/server.go b/spindle/server.go index 70134277..ecff63b0 100644 --- a/spindle/server.go +++ b/spindle/server.go @@ -420,6 +420,9 @@ func Run(ctx context.Context) error { // but the mill's db/notifier are created inside New. m.Attach(s.DB(), s.Notifier()) s.mill = m + if err := m.SeedBootstrapToken(); err != nil { + return fmt.Errorf("seeding mill bootstrap token: %w", err) + } } if cfg.Role == config.RoleExecutor { s.exec = executor.New(cfg, engines, s.DB(), s.Notifier(), log.SubLogger(logger, "executor")) -- 2.51.2