diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..b175375 --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +temp-*.env +plan.md +.claude/settings.local.json diff --git a/client.go b/client.go index db19928..c45a633 100644 --- a/client.go +++ b/client.go @@ -37,7 +37,7 @@ type kvRequest struct { // Pre-computed ed25519 signing keys (via NewPairEd25519() or any other means) // must be passed to a new client, with the expectation that the public key // be made available to the server to facilitate authentication. -// see: WriteRegistry() for details +// see: FileRegistry.Register() for details func NewClient(serverURL, keyPub, keyPriv string) (*Client, error) { rsaPublic, rsaPrivate, err := newPairRSA(Defaults.BitsizeRSA) if err != nil { diff --git a/example/testreg.yml b/example/testreg.yml deleted file mode 100644 index 9492a35..0000000 --- a/example/testreg.yml +++ /dev/null @@ -1,5 +0,0 @@ -- name: SERVICE1 - keypub: | - -----BEGIN ED25519 PUBLIC KEY----- - aOnchd3XQTs0iNNJGbrf4uwknkhjCO4mmLSbnEMUSQo= - -----END ED25519 PUBLIC KEY----- diff --git a/locket.go b/locket.go index ec48a9f..3791898 100644 --- a/locket.go +++ b/locket.go @@ -20,6 +20,12 @@ type defaults struct { MaxClockSkew time.Duration // max client/server clock difference before a request is rejected } +// PathRegistry is the API endpoint for registry operations. +// - GET: list all entries +// - POST: upsert an entry (RegEntry JSON body) +// - DELETE: remove an entry (RegEntry JSON body with name) +var PathRegistry = "/locket/registry" + // map[serviceName]keyPrivateSigning type KeysPrivateSigning map[string]string diff --git a/locket_test.go b/locket_test.go index 5f19a86..f339a2f 100644 --- a/locket_test.go +++ b/locket_test.go @@ -1,11 +1,12 @@ package locket import ( + "context" "fmt" "net/http" "net/http/httptest" "os" - "path" + "path/filepath" "strings" "testing" @@ -17,24 +18,16 @@ func TestE2E(t *testing.T) { pub, priv, err := NewPairEd25519() require.NoError(t, err) - reg := []RegEntry{{ - Name: "SERVICE1", - KeyPub: pub, - }} - testReg := path.Join("example", "testreg.yml") - err = WriteRegistry(testReg, reg) + fileReg := FileRegistry{Path: filepath.Join(t.TempDir(), "registry.yml")} + err = fileReg.Upsert(RegEntry{Name: "SERVICE1", KeyPub: pub}) require.NoError(t, err) source := Dotenv{ ServiceSecrets: testServiceMap, - Path: path.Join("example", ".env"), + Path: filepath.Join("example", ".env"), } - registry, err := ReadRegistryFile(testReg) - require.NoError(t, err) - require.NotNil(t, registry) - require.Greater(t, len(registry), 0) - server, err := NewServer(source, registry) + server, err := NewServer(context.Background(), source, fileReg, 0, nil) require.NoError(t, err) defer server.Close() @@ -51,9 +44,8 @@ func TestE2E(t *testing.T) { text := MarshalDotenv(envVars) t.Logf("formatted env vars: %s", text) // write the env vars to a file - f, err := os.CreateTemp("example", "temp-*.env") + f, err := os.CreateTemp(t.TempDir(), "temp-*.env") require.NoError(t, err) - defer os.Remove(f.Name()) _, err = f.WriteString(text) require.NoError(t, err) err = f.Close() diff --git a/registry.go b/registry.go index bbc2a53..5bd0fa7 100644 --- a/registry.go +++ b/registry.go @@ -3,8 +3,6 @@ package locket import ( "fmt" "os" - "path/filepath" - "strings" "gopkg.in/yaml.v3" ) @@ -12,127 +10,127 @@ import ( /* Registry is the process by which pre-computed signing keys (ed25519) are created before deploying either server or client. - - Public signing keys for all allowed services are provided to the server via .yml - - Public and private keys are provided to the client for signing requests. - - Separately, both client and server create encryption keys on startup. + - Public signing keys for all allowed services are provided to the server. + - Public and private keys are provided to the client for signing requests. + - Separately, both client and server create encryption keys on startup. */ // RegEntry is a single registry item, // representing a single client which -// the server should recognize and authorize +// the server should recognize and authorize. type RegEntry struct { - Name string `yaml:"name"` - KeyPub string `yaml:"keypub"` + Name string `yaml:"name" json:"name"` + KeyPub string `yaml:"keypub" json:"keypub"` } -// WriteRegistry creates a yaml file with a registry of allowed clients. -func WriteRegistry(path string, data []RegEntry) error { - f, err := os.Create(path) - if err != nil { - return fmt.Errorf("create file: %w", err) - } - defer f.Close() - - for i, item := range data { - data[i].Name = strings.TrimSuffix(filepath.Base(item.Name), ".env") - } +// Registry reads and writes authorized client entries. +// Implementations include FileRegistry (local YAML) and +// RemoteRegistry (HTTP API). +type Registry interface { + // Entries returns all authorized clients. + Entries() ([]RegEntry, error) + // Upsert inserts or updates a client entry by name. + Upsert(RegEntry) error + // Delete removes a client entry by name. + Delete(name string) error +} - b, err := yaml.Marshal(data) - if err != nil { - return fmt.Errorf("marshal: %w", err) - } - _, err = f.Write(b) - if err != nil { - return fmt.Errorf("write file: %w", err) - } - return nil +// FileRegistry is a Registry backed by a local YAML file. +type FileRegistry struct { + Path string } -// ReadRegistryFile turns a yaml file into a list of RegEntry -// for use in server authenticating client requests. -func ReadRegistryFile(filepath string) ([]RegEntry, error) { - f, err := os.ReadFile(filepath) +// Entries reads all authorized clients from the YAML file. +func (f FileRegistry) Entries() ([]RegEntry, error) { + b, err := os.ReadFile(f.Path) if err != nil { return nil, fmt.Errorf("read file: %w", err) } var out []RegEntry - err = yaml.Unmarshal(f, &out) + err = yaml.Unmarshal(b, &out) if err != nil { return nil, fmt.Errorf("unmarshal: %w", err) } return out, nil } -// UnmarshalRegistry turns a byte slice into a list of RegEntry -// for use in server authenticating client requests. -// Bytes format easier for embed.FS -func UnmarshalRegistry(bytes []byte) ([]RegEntry, error) { - var out []RegEntry - err := yaml.Unmarshal(bytes, &out) - if err != nil { - return nil, fmt.Errorf("unmarshal: %w", err) - } - return out, nil -} - -// Register reads the existing registry file, upserts service key, and rewrites. -// If the registry file does not exist, it will be created. -// If no registry for the named service exists, a new entry will be created. -// An existing entry for the named service will be updated with new public key. -// Each new call of Register will generate new key pair, returning: -// public key, private key, or any error. -func Register(name string, registryPath string) (string, string, error) { - publicKey, privateKey, err := NewPairEd25519() - if err != nil { - return "", "", fmt.Errorf("generate key pair: %w", err) - } - - // match the normalization WriteRegistry applies, so re-registering the same - // service updates its entry rather than appending a duplicate. - name = strings.TrimSuffix(filepath.Base(name), ".env") - - var registry []RegEntry - _, err = os.Stat(registryPath) - if err == nil { - registry, err = ReadRegistryFile(registryPath) +// Upsert inserts or updates a client entry in the YAML file. +// If the file does not exist, it is created. The name is stored verbatim +// (exact-match, case-sensitive) to match the RemoteRegistry / cloud contract; +// derive a clean service name before calling if needed. +func (f FileRegistry) Upsert(entry RegEntry) error { + var entries []RegEntry + _, err := os.Stat(f.Path) + switch { + case err == nil: + entries, err = f.Entries() if err != nil { - return "", "", fmt.Errorf("read registry file: %w", err) + return fmt.Errorf("read existing: %w", err) } - } else { - log.Debug( - "registry file does not exist (or other err); will create new one", - "registryPath", registryPath, - "statError", err, - ) + case !os.IsNotExist(err): + return fmt.Errorf("stat file: %w", err) } - // check if the service already exists in the registry replaced := false - for i, entry := range registry { - if entry.Name == name { - log.Debug("updating existing service in registry", - "service", name, - "publicKey", publicKey, - ) - registry[i].KeyPub = publicKey + for i, e := range entries { + if e.Name == entry.Name { + entries[i].KeyPub = entry.KeyPub replaced = true + break } } if !replaced { - log.Debug("adding new service to registry", - "service", name, - "publicKey", publicKey, - ) - registry = append(registry, RegEntry{ - Name: name, - KeyPub: publicKey, - }) + entries = append(entries, entry) } - // write the updated registry - err = WriteRegistry(registryPath, registry) + return f.write(entries) +} + +// Delete removes a client entry by name from the YAML file. +func (f FileRegistry) Delete(name string) error { + entries, err := f.Entries() + if err != nil { + return fmt.Errorf("read existing: %w", err) + } + filtered := entries[:0] + for _, e := range entries { + if e.Name != name { + filtered = append(filtered, e) + } + } + return f.write(filtered) +} + +// Register generates a new ed25519 signing keypair, upserts the +// public key into the YAML file, and returns the keypair. +func (f FileRegistry) Register(name string) (string, string, error) { + pub, priv, err := NewPairEd25519() + if err != nil { + return "", "", fmt.Errorf("generate key pair: %w", err) + } + err = f.Upsert(RegEntry{Name: name, KeyPub: pub}) + if err != nil { + return "", "", fmt.Errorf("upsert: %w", err) + } + return pub, priv, nil +} + +// write serializes entries to the YAML file, creating or +// truncating it as needed. +func (f FileRegistry) write(entries []RegEntry) error { + file, err := os.Create(f.Path) + if err != nil { + return fmt.Errorf("create file: %w", err) + } + defer file.Close() + + b, err := yaml.Marshal(entries) + if err != nil { + return fmt.Errorf("marshal: %w", err) + } + _, err = file.Write(b) if err != nil { - return "", "", fmt.Errorf("write updated registry: %w", err) + return fmt.Errorf("write file: %w", err) } - return publicKey, privateKey, nil + return nil } diff --git a/registry_remote.go b/registry_remote.go new file mode 100644 index 0000000..7248988 --- /dev/null +++ b/registry_remote.go @@ -0,0 +1,160 @@ +package locket + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/url" + "time" +) + +// RemoteRegistry is a Registry backed by an HTTP API. +// URL is the base URL (e.g. "http://api:8888") to which +// PathRegistry is appended for all operations. +// Token, if set, is sent as an X-Auth-Token header. +type RemoteRegistry struct { + URL string + Token string + Client *http.Client +} + +// client returns the configured HTTP client, or a default with timeout. +func (r RemoteRegistry) client() *http.Client { + if r.Client != nil { + return r.Client + } + return &http.Client{Timeout: 10 * time.Second} +} + +// endpoint returns the full URL to the registry API, +// safely joining the base URL and PathRegistry. It validates +// that r.URL is an absolute URL with scheme and host. +func (r RemoteRegistry) endpoint() (string, error) { + joined, err := url.JoinPath(r.URL, PathRegistry) + if err != nil { + return "", fmt.Errorf("join url: %w", err) + } + u, err := url.Parse(joined) + if err != nil { + return "", fmt.Errorf("parse url: %w", err) + } + if u.Scheme == "" || u.Host == "" { + return "", fmt.Errorf( + "invalid base URL %q: missing scheme or host", r.URL, + ) + } + return joined, nil +} + +// Entries fetches all authorized clients from the remote API. +func (r RemoteRegistry) Entries() ([]RegEntry, error) { + endpoint, err := r.endpoint() + if err != nil { + return nil, fmt.Errorf("endpoint: %w", err) + } + req, err := http.NewRequest(http.MethodGet, endpoint, nil) + if err != nil { + return nil, fmt.Errorf("new request: %w", err) + } + r.setHeaders(req) + + resp, err := r.client().Do(req) + if err != nil { + return nil, fmt.Errorf("do request: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return nil, fmt.Errorf("status %s", resp.Status) + } + + var entries []RegEntry + err = json.NewDecoder(resp.Body).Decode(&entries) + if err != nil { + return nil, fmt.Errorf("decode: %w", err) + } + return entries, nil +} + +// Upsert creates or updates an authorized client via the remote API. +func (r RemoteRegistry) Upsert(entry RegEntry) error { + b, err := json.Marshal(entry) + if err != nil { + return fmt.Errorf("marshal: %w", err) + } + + endpoint, err := r.endpoint() + if err != nil { + return fmt.Errorf("endpoint: %w", err) + } + req, err := http.NewRequest( + http.MethodPost, endpoint, bytes.NewReader(b), + ) + if err != nil { + return fmt.Errorf("new request: %w", err) + } + r.setHeaders(req) + req.Header.Set("Content-Type", "application/json") + + resp, err := r.client().Do(req) + if err != nil { + return fmt.Errorf("do request: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return fmt.Errorf("status %s", resp.Status) + } + return nil +} + +// Delete removes an authorized client by name via the remote API. +func (r RemoteRegistry) Delete(name string) error { + b, err := json.Marshal(RegEntry{Name: name}) + if err != nil { + return fmt.Errorf("marshal: %w", err) + } + + endpoint, err := r.endpoint() + if err != nil { + return fmt.Errorf("endpoint: %w", err) + } + req, err := http.NewRequest( + http.MethodDelete, endpoint, bytes.NewReader(b), + ) + if err != nil { + return fmt.Errorf("new request: %w", err) + } + r.setHeaders(req) + req.Header.Set("Content-Type", "application/json") + + resp, err := r.client().Do(req) + if err != nil { + return fmt.Errorf("do request: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + return fmt.Errorf("status %s", resp.Status) + } + return nil +} + +// Register generates a new ed25519 signing keypair, upserts the +// public key via the remote API, and returns the keypair. +func (r RemoteRegistry) Register(name string) (string, string, error) { + pub, priv, err := NewPairEd25519() + if err != nil { + return "", "", fmt.Errorf("generate key pair: %w", err) + } + err = r.Upsert(RegEntry{Name: name, KeyPub: pub}) + if err != nil { + return "", "", fmt.Errorf("upsert: %w", err) + } + return pub, priv, nil +} + +// setHeaders applies auth headers to the request. +func (r RemoteRegistry) setHeaders(req *http.Request) { + if r.Token != "" { + req.Header.Set("X-Auth-Token", r.Token) + } +} diff --git a/registry_remote_test.go b/registry_remote_test.go new file mode 100644 index 0000000..3cfaadb --- /dev/null +++ b/registry_remote_test.go @@ -0,0 +1,126 @@ +package locket + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestRemoteRegistryEntries(t *testing.T) { + want := []RegEntry{ + {Name: "svc1", KeyPub: "pub1"}, + {Name: "svc2", KeyPub: "pub2"}, + } + + srv := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodGet, r.Method) + require.Equal(t, PathRegistry, r.URL.Path) + require.Equal(t, "tok", r.Header.Get("X-Auth-Token")) + w.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(w).Encode(want)) + }, + )) + defer srv.Close() + + // Trailing slash on URL; JoinPath should normalize. + reg := RemoteRegistry{URL: srv.URL + "/", Token: "tok"} + got, err := reg.Entries() + require.NoError(t, err) + require.Equal(t, want, got) +} + +func TestRemoteRegistryUpsert(t *testing.T) { + want := RegEntry{Name: "svc1", KeyPub: "pub1"} + + srv := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodPost, r.Method) + require.Equal(t, PathRegistry, r.URL.Path) + require.Equal(t, + "application/json", r.Header.Get("Content-Type"), + ) + + body, err := io.ReadAll(r.Body) + require.NoError(t, err) + var got RegEntry + require.NoError(t, json.Unmarshal(body, &got)) + require.Equal(t, want, got) + + // Return 201 to confirm 2xx acceptance. + w.WriteHeader(http.StatusCreated) + }, + )) + defer srv.Close() + + reg := RemoteRegistry{URL: srv.URL} + require.NoError(t, reg.Upsert(want)) +} + +func TestRemoteRegistryDelete(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodDelete, r.Method) + require.Equal(t, PathRegistry, r.URL.Path) + + body, err := io.ReadAll(r.Body) + require.NoError(t, err) + var got RegEntry + require.NoError(t, json.Unmarshal(body, &got)) + require.Equal(t, "svc1", got.Name) + + // Return 204 to confirm 2xx acceptance. + w.WriteHeader(http.StatusNoContent) + }, + )) + defer srv.Close() + + reg := RemoteRegistry{URL: srv.URL} + require.NoError(t, reg.Delete("svc1")) +} + +func TestRemoteRegistryRegister(t *testing.T) { + var seen RegEntry + + srv := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodPost, r.Method) + require.NoError(t, json.NewDecoder(r.Body).Decode(&seen)) + w.WriteHeader(http.StatusOK) + }, + )) + defer srv.Close() + + reg := RemoteRegistry{URL: srv.URL} + pub, priv, err := reg.Register("svc1") + require.NoError(t, err) + require.NotEmpty(t, pub) + require.NotEmpty(t, priv) + require.Equal(t, "svc1", seen.Name) + require.Equal(t, pub, seen.KeyPub) +} + +func TestRemoteRegistryInvalidBaseURL(t *testing.T) { + reg := RemoteRegistry{URL: "api:8888"} + _, err := reg.Entries() + require.Error(t, err) + require.ErrorContains(t, err, "missing scheme or host") +} + +func TestRemoteRegistryNon2xx(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + http.Error(w, "boom", http.StatusInternalServerError) + }, + )) + defer srv.Close() + + reg := RemoteRegistry{URL: srv.URL} + _, err := reg.Entries() + require.Error(t, err) + require.ErrorContains(t, err, "500") +} diff --git a/registry_test.go b/registry_test.go index f74335a..7fcb023 100644 --- a/registry_test.go +++ b/registry_test.go @@ -1,8 +1,6 @@ package locket import ( - "os" - "path" "path/filepath" "testing" @@ -20,13 +18,15 @@ var testRegistryItems = []RegEntry{ }, } -var testExampleReg = path.Join("example", "registry.yml") +func TestFileRegistryReadWrite(t *testing.T) { + reg := FileRegistry{Path: filepath.Join(t.TempDir(), "registry.yml")} -func TestReadWrite(t *testing.T) { - err := WriteRegistry(testExampleReg, testRegistryItems) - require.NoError(t, err) + for _, item := range testRegistryItems { + err := reg.Upsert(item) + require.NoError(t, err) + } - items, err := ReadRegistryFile(testExampleReg) + items, err := reg.Entries() require.NoError(t, err) require.NotNil(t, items) require.Equal(t, testRegistryItems, items) @@ -35,61 +35,86 @@ func TestReadWrite(t *testing.T) { } } -// TestRegisterNoDuplicate is the regression test for Register name -// normalization: re-registering a service (including with a .env suffix that -// WriteRegistry strips) must update the existing entry rather than append a -// duplicate. -func TestRegisterNoDuplicate(t *testing.T) { - reg := filepath.Join(t.TempDir(), "registry.yml") +// TestFileRegistryUpsertDedup verifies Upsert dedups on the exact name +// (verbatim, case-sensitive) per the cloud registry contract: re-registering +// the same name updates the entry in place, while a different name (e.g. one +// with a .env suffix) is a distinct entry — Upsert does not normalize. +func TestFileRegistryUpsertDedup(t *testing.T) { + reg := FileRegistry{Path: filepath.Join(t.TempDir(), "registry.yml")} - _, _, err := Register("svc.env", reg) + pub1, _, err := reg.Register("svc") require.NoError(t, err) - pub2, _, err := Register("svc.env", reg) + // re-registering the exact name updates in place (no duplicate) + pub2, _, err := reg.Register("svc") require.NoError(t, err) - // the plain name normalizes to the same entry too - _, _, err = Register("svc", reg) + // a different, un-normalized name is a distinct entry + _, _, err = reg.Register("svc.env") require.NoError(t, err) - entries, err := ReadRegistryFile(reg) + entries, err := reg.Entries() require.NoError(t, err) - require.Len(t, entries, 1, "re-registering the same service must not duplicate") - require.Equal(t, "svc", entries[0].Name) - // last write wins on the key - require.NotEqual(t, pub2, entries[0].KeyPub) + require.Len(t, entries, 2, + "exact-name re-register dedups; a distinct name does not") + + byName := make(map[string]string, len(entries)) + for _, e := range entries { + byName[e.Name] = e.KeyPub + } + require.Contains(t, byName, "svc") + require.Contains(t, byName, "svc.env") + require.Equal(t, pub2, byName["svc"], "last write wins for the exact name") + require.NotEqual(t, pub1, pub2) } -func TestRegister(t *testing.T) { - testRegistry := path.Join("example", "test-registry.yml") +func TestFileRegistryRegister(t *testing.T) { + reg := FileRegistry{Path: filepath.Join(t.TempDir(), "registry.yml")} + services := []string{"service A", "service B", "service C"} - // no file exists + for _, service := range services { - pub, priv, err := Register(service, testRegistry) + pub, priv, err := reg.Register(service) require.NoError(t, err) require.NotEmpty(t, pub) require.NotEmpty(t, priv) t.Logf("Registered %q", service) - registry, err := ReadRegistryFile(testRegistry) + entries, err := reg.Entries() require.NoError(t, err) - for _, item := range registry { - t.Logf("%s\n%s", item.Name, item.KeyPub) + for _, e := range entries { + t.Logf("%s\n%s", e.Name, e.KeyPub) } } - // replace/update/upsert service A with new key - pub, priv, err := Register(services[0], testRegistry) + + // upsert service A with new key + pub, priv, err := reg.Register(services[0]) require.NoError(t, err) require.NotEmpty(t, pub) require.NotEmpty(t, priv) t.Logf("Updated %q", services[0]) - registry, err := ReadRegistryFile(testRegistry) + + entries, err := reg.Entries() require.NoError(t, err) - for _, item := range registry { - t.Logf("%s\n%s", item.Name, item.KeyPub) + for _, e := range entries { + t.Logf("%s\n%s", e.Name, e.KeyPub) } - // cleanup file - err = os.Remove(testRegistry) +} + +func TestFileRegistryDelete(t *testing.T) { + reg := FileRegistry{Path: filepath.Join(t.TempDir(), "registry.yml")} + + require.NoError(t, reg.Upsert(RegEntry{Name: "a", KeyPub: "k1"})) + require.NoError(t, reg.Upsert(RegEntry{Name: "b", KeyPub: "k2"})) + + entries, err := reg.Entries() + require.NoError(t, err) + require.Len(t, entries, 2) + + require.NoError(t, reg.Delete("a")) + + entries, err = reg.Entries() require.NoError(t, err) - t.Logf("Removed test registry file: %s", testRegistry) + require.Len(t, entries, 1) + require.Equal(t, "b", entries[0].Name) } diff --git a/server.go b/server.go index 32c3cd8..62513f3 100644 --- a/server.go +++ b/server.go @@ -1,6 +1,7 @@ package locket import ( + "context" "encoding/json" "fmt" "net" @@ -12,12 +13,33 @@ import ( "github.com/google/uuid" ) -type Server struct { - secrets map[string]Secrets // secret k/v pairs - registry []RegEntry // registered services - keyRsaPublic string // encryption public key - keyRsaPrivate string // encryption private key - seen *nonceCache // request nonces seen within the replay window +// AllowRequestFunc decides whether an HTTP request is permitted. +// Return nil to allow, or an error to deny with 403. +type AllowRequestFunc func(r *http.Request) error + +// AllowCIDR returns an AllowRequestFunc that permits requests from +// the given CIDR range only. This is the default policy when +// none is provided to NewServer. +func AllowCIDR(cidr string) AllowRequestFunc { + _, network, cidrErr := net.ParseCIDR(cidr) + + return func(r *http.Request) error { + if cidrErr != nil { + return fmt.Errorf("parse CIDR: %w", cidrErr) + } + ip, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + return fmt.Errorf("parse remote addr: %w", err) + } + clientIP := net.ParseIP(ip) + if clientIP == nil { + return fmt.Errorf("parse ip: %q", ip) + } + if !network.Contains(clientIP) { + return fmt.Errorf("IP %s not in %s", ip, cidr) + } + return nil + } } // nonceCache tracks request nonces so the server can reject exact replays @@ -84,22 +106,71 @@ func (c *nonceCache) close() { c.stopOnce.Do(func() { close(c.stop) }) } -// kvResponse is the server's response to the client's request, -// containing the encrypted secret value. +// Server serves secrets to authenticated clients over HTTP. +// It validates client requests against a Registry of authorized +// signing keys, refreshing the registry on a configurable interval. +type Server struct { + secrets map[string]Secrets + reg Registry + entries []RegEntry + mu sync.RWMutex + allow AllowRequestFunc + keyRsaPublic string + keyRsaPrivate string + seen *nonceCache + cancel context.CancelFunc // stops the registry poll goroutine +} + +// kvResponse is the server's encrypted secret response. type kvResponse struct { Payload string `json:"payload"` } -// NewServer sets up a new secrets server when provided source options -// and registry of allowed services (and their public signing keys), -// expected to be read from file or embed before calling NewServer(). -func NewServer(opts source, registry []RegEntry) (*Server, error) { +// NewServer creates a Server, loading secrets from the given source +// and authorized clients from the given Registry. If pollInterval +// is positive, the server refreshes its registry in the background. +// If allow is nil, AllowCIDR(Defaults.AllowCIDR) is used. +func NewServer( + ctx context.Context, + opts source, + reg Registry, + pollInterval time.Duration, + allow AllowRequestFunc, +) (*Server, error) { + if ctx == nil { + ctx = context.Background() + } + if reg == nil { + return nil, fmt.Errorf("registry must not be nil") + } + rsaPublic, rsaPrivate, err := newPairRSA(Defaults.BitsizeRSA) if err != nil { return nil, fmt.Errorf("generate key pair: %w", err) } - server := Server{ - registry: registry, + + entries, err := reg.Entries() + if err != nil { + return nil, fmt.Errorf("initial registry fetch: %w", err) + } + log.Info("registry loaded", "entries", len(entries)) + + if allow == nil { + // validate up front so a bad Defaults.AllowCIDR fails fast here + // instead of silently rejecting every request with 403 later. + if _, _, err := net.ParseCIDR(Defaults.AllowCIDR); err != nil { + return nil, fmt.Errorf( + "invalid default allow CIDR %q: %w", + Defaults.AllowCIDR, err, + ) + } + allow = AllowCIDR(Defaults.AllowCIDR) + } + + server := &Server{ + reg: reg, + entries: entries, + allow: allow, keyRsaPublic: rsaPublic, keyRsaPrivate: rsaPrivate, seen: newNonceCache(Defaults.MaxClockSkew), @@ -112,38 +183,88 @@ func NewServer(opts source, registry []RegEntry) (*Server, error) { return nil, fmt.Errorf("load env: %w", err) } server.secrets = secrets - return &server, nil case Dotenv: if len(opts.ServiceSecrets) == 0 { - return nil, fmt.Errorf("at least one service required to load *.env file") + return nil, fmt.Errorf( + "at least one service required to load *.env file", + ) } if opts.Path == "" { - return nil, fmt.Errorf("at least one path required to *.env file") + return nil, fmt.Errorf( + "opts.Path must be set to the .env file path", + ) } secrets, err := opts.Load() if err != nil { return nil, fmt.Errorf("load .env file: %w", err) } server.secrets = secrets - return &server, nil case Onepass: secrets, err := opts.Load() if err != nil { return nil, fmt.Errorf("load onepass: %w", err) } server.secrets = secrets - return &server, nil default: return nil, fmt.Errorf("invalid source") } + + if pollInterval > 0 { + // derive a child context so Close can stop polling independently of + // the caller's context (which may be context.Background()). + pollCtx, cancel := context.WithCancel(ctx) + server.cancel = cancel + go server.poll(pollCtx, pollInterval) + } + + return server, nil } -// Close releases the server's background resources (the nonce-cache sweeper). -// The Server must not be used after Close. +// Close releases the server's background resources: the registry poll +// goroutine (if any) and the nonce-cache sweeper. The Server must not be used +// after Close. func (s *Server) Close() { + if s.cancel != nil { + s.cancel() + } s.seen.close() } +// poll refreshes the registry on a fixed interval until ctx is cancelled. +func (s *Server) poll(ctx context.Context, interval time.Duration) { + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + entries, err := s.reg.Entries() + if err != nil { + log.Error("registry poll failed", "error", err) + continue + } + s.mu.Lock() + s.entries = entries + s.mu.Unlock() + log.Debug("registry refreshed", "entries", len(entries)) + } + } +} + +// registrySnapshot returns a point-in-time copy of the registry. +func (s *Server) registrySnapshot() []RegEntry { + s.mu.RLock() + defer s.mu.RUnlock() + out := make([]RegEntry, len(s.entries)) + copy(out, s.entries) + return out +} + +// Handler is the HTTP handler for the locket secret server. +// GET returns the server's RSA public encryption key. +// POST accepts an encrypted, signed secret request and returns +// the encrypted secret value. func (s *Server) Handler(w http.ResponseWriter, r *http.Request) { id := uuid.New().String() log.Info("received request", @@ -154,172 +275,182 @@ func (s *Server) Handler(w http.ResponseWriter, r *http.Request) { ) switch r.Method { case http.MethodOptions: - w.Header().Set("Allow", fmt.Sprintf("%s, %s", http.MethodGet, http.MethodPost)) + w.Header().Set("Allow", fmt.Sprintf( + "%s, %s", http.MethodGet, http.MethodPost, + )) return case http.MethodGet: w.Header().Set("Content-Type", "text/plain") w.Write([]byte(s.keyRsaPublic)) return case http.MethodPost: - var request kvRequest - // log.Debug("decoding request", "body", r.Body) - err := json.NewDecoder(r.Body).Decode(&request) - if err != nil { - http.Error(w, "bad request", http.StatusBadRequest) - return - } - log.Debug("request", - "payload", request.Payload, - "client_pubkey", request.ClientPubKey, - "signature", request.PayloadSignature, + s.handlePost(w, r, id) + default: + log.Warn("method not allowed", + "method", r.Method, "request_id", id, + "ip", r.RemoteAddr, ) - payload, err := decryptRSA(s.keyRsaPrivate, request.Payload) - if err != nil { - log.Error("decrypt payload", "request_id", id, "error", err) - http.Error(w, "bad request", http.StatusBadRequest) - return - } - log.Debug("request payload decrypted", "request_id", id) + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + } +} - // require from CIDR range DefaultAllowCIDR - ip, _, err := net.SplitHostPort(r.RemoteAddr) - if err != nil { - log.Error("split host port", "request_id", id, "error", err) - // maybe a 5xx, but probably only because of malformed host addr - http.Error(w, "bad request", http.StatusBadRequest) - return - } - clientIP := net.ParseIP(ip) - _, cidr, err := net.ParseCIDR(Defaults.AllowCIDR) - if err != nil { - log.Error("parse CIDR", "request_id", id, "error", err) - http.Error(w, "bad request", http.StatusBadRequest) - return - } - if !cidr.Contains(clientIP) { - log.Warn("IP rejected", - "request_id", id, - "ip", r.RemoteAddr, - "allowCIDR", Defaults.AllowCIDR, - ) - http.Error(w, "forbidden", http.StatusForbidden) - return - } - log.Debug("IP allowed", +// handlePost processes an encrypted secret request. +func (s *Server) handlePost( + w http.ResponseWriter, r *http.Request, id string, +) { + if err := s.allow(r); err != nil { + log.Warn("request denied", "request_id", id, "ip", r.RemoteAddr, - "allowCIDR", Defaults.AllowCIDR, + "error", err, ) + http.Error(w, "forbidden", http.StatusForbidden) + return + } - // a nonce is required to detect replays - if request.Nonce == "" { - log.Warn("request missing nonce", "request_id", id) - http.Error(w, "bad request", http.StatusBadRequest) - return - } + const maxBody = 1 << 20 + r.Body = http.MaxBytesReader(w, r.Body, maxBody) - // reject stale or future-dated requests to bound replay - skew := time.Since(time.Unix(request.Timestamp, 0)) - if skew < 0 { - skew = -skew - } - if skew > Defaults.MaxClockSkew { - log.Warn("request timestamp outside allowed window", - "request_id", id, - "skew", skew, - "max", Defaults.MaxClockSkew, + var request kvRequest + err := json.NewDecoder(r.Body).Decode(&request) + if err != nil { + if _, ok := err.(*http.MaxBytesError); ok { + http.Error(w, + "request entity too large", + http.StatusRequestEntityTooLarge, ) - http.Error(w, "forbidden", http.StatusForbidden) return } + http.Error(w, "bad request", http.StatusBadRequest) + return + } + log.Debug("request", + "payload", request.Payload, + "client_pubkey", request.ClientPubKey, + "signature", request.PayloadSignature, + "request_id", id, + ) - // verify signature against registry; the signed message binds the - // client pubkey, timestamp, and nonce so a captured request cannot be - // replayed with a substituted ClientPubKey to redirect the secret. - var matches bool - var verifiedService string - message := requestMessage(payload, request.ClientPubKey, request.Timestamp, request.Nonce) - log.Debug("verifying signature", "request_id", id) - for _, svc := range s.registry { - match, err := verifyEd25519(svc.KeyPub, message, request.PayloadSignature) - if err != nil { - log.Error("verify signature", "request_id", id, "error", err) - continue - } - if match { - matches = true - verifiedService = svc.Name - break - } - } - if !matches { - log.Error("signature mismatch", "request_id", id) - http.Error(w, "forbidden", http.StatusForbidden) - return - } else { - log.Debug("signature verified", - "service", verifiedService, - "request_id", id, - ) - } - - // reject replays: a nonce is valid only until a replay could no longer - // pass the freshness check above. Checked after signature verification - // so unauthenticated requests cannot fill the cache. - expiry := time.Unix(request.Timestamp, 0).Add(Defaults.MaxClockSkew) - if s.seen.observe(request.Nonce, expiry) { - log.Warn("replayed request rejected", - "service", verifiedService, - "request_id", id, - ) - http.Error(w, "forbidden", http.StatusForbidden) - return - } + payload, err := decryptRSA(s.keyRsaPrivate, request.Payload) + if err != nil { + log.Error("decrypt payload", + "request_id", id, "error", err, + ) + http.Error(w, "bad request", http.StatusBadRequest) + return + } + log.Debug("request payload decrypted", "request_id", id) - log.Debug("secrets for service", "service", verifiedService, "secrets_qty", len(s.secrets)) - secrets, ok := s.secrets[strings.ToLower(verifiedService)] - if !ok { - log.Warn("service not found in registry, check case sensitivity (expects lower)", - "service_requesting", verifiedService, - "request_id", id, - ) - http.Error(w, "forbidden", http.StatusForbidden) - return - } + // a nonce is required to detect replays + if request.Nonce == "" { + log.Warn("request missing nonce", "request_id", id) + http.Error(w, "bad request", http.StatusBadRequest) + return + } - value, ok := secrets[payload] - if !ok { - log.Warn("secret not found", "service", verifiedService, "key", payload, "request_id", id) - http.Error(w, "not found", http.StatusNotFound) - return - } + // reject stale or future-dated requests to bound replay + skew := time.Since(time.Unix(request.Timestamp, 0)) + if skew < 0 { + skew = -skew + } + if skew > Defaults.MaxClockSkew { + log.Warn("request timestamp outside allowed window", + "request_id", id, + "skew", skew, + "max", Defaults.MaxClockSkew, + ) + http.Error(w, "forbidden", http.StatusForbidden) + return + } - ecryptedSecret, err := encryptRSA(request.ClientPubKey, value) + // verify signature against registry; the signed message binds the client + // pubkey, timestamp, and nonce so a captured request cannot be replayed + // with a substituted ClientPubKey to redirect the secret. + registry := s.registrySnapshot() + var verifiedService string + message := requestMessage( + payload, request.ClientPubKey, request.Timestamp, request.Nonce, + ) + for _, svc := range registry { + match, err := verifyEd25519( + svc.KeyPub, message, request.PayloadSignature, + ) if err != nil { - log.Error("encrypt secret", "request_id", id, "error", err) - http.Error(w, "server error", http.StatusInternalServerError) - return + log.Error("verify signature", + "request_id", id, "error", err, + ) + continue } - response := kvResponse{ - Payload: ecryptedSecret, + if match { + verifiedService = svc.Name + break } - // header must be set before the body is written to take effect - w.Header().Set("Content-Type", "application/json") - err = json.NewEncoder(w).Encode(response) - if err != nil { - log.Error("encode response", "request_id", id, "error", err) - http.Error(w, "bad request", http.StatusBadRequest) - return - } - log.Info("sending secret", + } + if verifiedService == "" { + log.Error("signature mismatch", "request_id", id) + http.Error(w, "forbidden", http.StatusForbidden) + return + } + log.Debug("signature verified", + "service", verifiedService, "request_id", id, + ) + + // reject replays: a nonce is valid only until a replay could no longer + // pass the freshness check above. Checked after signature verification + // so unauthenticated requests cannot fill the cache. + expiry := time.Unix(request.Timestamp, 0).Add(Defaults.MaxClockSkew) + if s.seen.observe(request.Nonce, expiry) { + log.Warn("replayed request rejected", "service", verifiedService, - "name", payload, - "ip", r.RemoteAddr, "request_id", id, ) - default: - log.Warn("method not allowed", "method", r.Method, "request_id", id, "ip", r.RemoteAddr) - http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + http.Error(w, "forbidden", http.StatusForbidden) + return + } + + secrets, ok := s.secrets[strings.ToLower(verifiedService)] + if !ok { + log.Warn("service not found, check case (expects lower)", + "service", verifiedService, + "request_id", id, + ) + http.Error(w, "forbidden", http.StatusForbidden) + return + } + + value, ok := secrets[payload] + if !ok { + log.Warn("secret not found", + "service", verifiedService, + "key", payload, + "request_id", id, + ) + http.Error(w, "not found", http.StatusNotFound) + return + } + + encrypted, err := encryptRSA(request.ClientPubKey, value) + if err != nil { + log.Error("encrypt secret", + "request_id", id, "error", err, + ) + http.Error(w, "server error", http.StatusInternalServerError) + return } + + w.Header().Set("Content-Type", "application/json") + err = json.NewEncoder(w).Encode(kvResponse{Payload: encrypted}) + if err != nil { + log.Error("encode response", + "request_id", id, "error", err, + ) + return + } + log.Info("sending secret", + "service", verifiedService, + "name", payload, + "ip", r.RemoteAddr, + "request_id", id, + ) } diff --git a/server_test.go b/server_test.go index 51a9762..c7c0414 100644 --- a/server_test.go +++ b/server_test.go @@ -2,10 +2,13 @@ package locket import ( "bytes" + "context" "encoding/json" "io" "net/http" "net/http/httptest" + "path/filepath" + "sync" "testing" "time" @@ -20,18 +23,20 @@ const ( // newTestServer builds a server backed by the example .env and a single // registered service, returning the running test server, the underlying // *Server (for its encryption pubkey), and the service's ed25519 signing keys. -// No files are written, so example/testreg.yml is left untouched. +// The registry lives in a temp dir so no tracked fixtures are touched. func newTestServer(t *testing.T) (*httptest.Server, *Server, string) { t.Helper() pub, priv, err := NewPairEd25519() require.NoError(t, err) - registry := []RegEntry{{Name: "SERVICE1", KeyPub: pub}} + reg := FileRegistry{Path: filepath.Join(t.TempDir(), "registry.yml")} + require.NoError(t, reg.Upsert(RegEntry{Name: "SERVICE1", KeyPub: pub})) + source := Dotenv{ Path: testEnvFile, ServiceSecrets: testServiceMap, } - server, err := NewServer(source, registry) + server, err := NewServer(context.Background(), source, reg, 0, nil) require.NoError(t, err) t.Cleanup(server.Close) @@ -136,15 +141,17 @@ func TestHandlerRejectsPubkeySubstitution(t *testing.T) { // bug: a fully valid request from outside the allowed CIDR must be blocked and // must not leak the secret. func TestHandlerRejectsOutOfCIDR(t *testing.T) { - ts, server, signingPriv := newTestServer(t) - clientPub, clientPriv, err := newPairRSA(Defaults.BitsizeRSA) - require.NoError(t, err) - // test requests originate from 127.0.0.1; exclude it from the allowlist. + // The allow policy is captured at construction, so set this before + // newTestServer builds the server. prev := Defaults.AllowCIDR Defaults.AllowCIDR = "10.0.0.0/24" t.Cleanup(func() { Defaults.AllowCIDR = prev }) + ts, server, signingPriv := newTestServer(t) + clientPub, clientPriv, err := newPairRSA(Defaults.BitsizeRSA) + require.NoError(t, err) + req := craftRequest(t, server.keyRsaPublic, signingPriv, testSecretName, clientPub, time.Now().Unix()) resp, body := postRequest(t, ts.URL, req) require.Equal(t, http.StatusForbidden, resp.StatusCode) @@ -169,6 +176,51 @@ func TestHandlerRejectsReplay(t *testing.T) { assertSecretNotLeaked(t, body2, clientPriv) } +// countingRegistry records how many times Entries is called, to verify the +// poll goroutine's lifecycle. +type countingRegistry struct { + mu sync.Mutex + count int +} + +func (c *countingRegistry) Entries() ([]RegEntry, error) { + c.mu.Lock() + c.count++ + c.mu.Unlock() + return nil, nil +} +func (c *countingRegistry) Upsert(RegEntry) error { return nil } +func (c *countingRegistry) Delete(string) error { return nil } +func (c *countingRegistry) Register(string) (string, string, error) { + return "", "", nil +} +func (c *countingRegistry) calls() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.count +} + +// TestServerCloseStopsPoll is the regression test for the lifecycle gap: Close +// must stop the registry poll goroutine even when the server was created with a +// non-cancelable context. +func TestServerCloseStopsPoll(t *testing.T) { + reg := &countingRegistry{} + source := Dotenv{Path: testEnvFile, ServiceSecrets: testServiceMap} + interval := 5 * time.Millisecond + + server, err := NewServer(context.Background(), source, reg, interval, nil) + require.NoError(t, err) + + time.Sleep(40 * time.Millisecond) + require.Greater(t, reg.calls(), 1, "poll should have run several times") + + server.Close() + time.Sleep(20 * time.Millisecond) // let any in-flight tick finish + stopped := reg.calls() + time.Sleep(40 * time.Millisecond) // several more intervals + require.Equal(t, stopped, reg.calls(), "poll must not run after Close") +} + // TestHandlerRejectsStaleTimestamp is the regression test for the replay // window: a request whose signed timestamp is outside MaxClockSkew is rejected // even though the signature itself is valid.