From fc26c447ec27b88dc190ed79964ab1239f797187 Mon Sep 17 00:00:00 2001 From: Kristoffer Dalby Date: Wed, 17 Jun 2026 15:29:17 +0000 Subject: [PATCH] hscontrol/api/v1: implement pre-auth key endpoints CreatePreAuthKey (with ACL tag validation), ListPreAuthKeys, ExpirePreAuthKey, DeletePreAuthKey over the state layer, with HTTP-parity tests. --- hscontrol/api/v1/convert.go | 40 ++++++ hscontrol/api/v1/preauthkeys.go | 117 ++++++++++++++++++ .../servertest/apiv1_preauthkeys_test.go | 112 +++++++++++++++++ 3 files changed, 269 insertions(+) create mode 100644 hscontrol/api/v1/preauthkeys.go create mode 100644 hscontrol/servertest/apiv1_preauthkeys_test.go diff --git a/hscontrol/api/v1/convert.go b/hscontrol/api/v1/convert.go index 49194445..0aa378ea 100644 --- a/hscontrol/api/v1/convert.go +++ b/hscontrol/api/v1/convert.go @@ -44,6 +44,46 @@ func optTime(ts *timestamppb.Timestamp) oas.OptDateTime { return oas.NewOptDateTime(ts.AsTime()) } +func optBool(b bool) oas.OptBool { + if !b { + return oas.OptBool{} + } + + return oas.NewOptBool(b) +} + +// strs normalises an empty slice to nil so it is omitted from the response +// rather than emitted as an empty array. +func strs(s []string) []string { + if len(s) == 0 { + return nil + } + + return s +} + +func optUser(u *v1.User) oas.OptUser { + if u == nil { + return oas.OptUser{} + } + + return oas.NewOptUser(oasUser(u)) +} + +func oasPreAuthKey(k *v1.PreAuthKey) oas.PreAuthKey { + return oas.PreAuthKey{ + User: optUser(k.GetUser()), + ID: optUint64(k.GetId()), + Key: optString(k.GetKey()), + Reusable: optBool(k.GetReusable()), + Ephemeral: optBool(k.GetEphemeral()), + Used: optBool(k.GetUsed()), + Expiration: optTime(k.GetExpiration()), + CreatedAt: optTime(k.GetCreatedAt()), + AclTags: strs(k.GetAclTags()), + } +} + func oasAPIKey(k *v1.ApiKey) oas.ApiKey { return oas.ApiKey{ ID: optUint64(k.GetId()), diff --git a/hscontrol/api/v1/preauthkeys.go b/hscontrol/api/v1/preauthkeys.go new file mode 100644 index 00000000..1d6e5189 --- /dev/null +++ b/hscontrol/api/v1/preauthkeys.go @@ -0,0 +1,117 @@ +package apiv1 + +import ( + "cmp" + "context" + "errors" + "slices" + "strings" + "time" + + oas "github.com/juanfont/headscale/gen/api/v1" + "github.com/juanfont/headscale/hscontrol/types" +) + +// CreatePreAuthKey creates a pre-auth key for a user. +func (s *Server) CreatePreAuthKey( + _ context.Context, + req *oas.CreatePreAuthKeyReq, +) (*oas.CreatePreAuthKeyOK, error) { + var expiration time.Time + if v, ok := req.Expiration.Get(); ok { + expiration = v + } + + for _, tag := range req.AclTags { + err := validateTag(tag) + if err != nil { + return nil, badRequest(err.Error()) + } + } + + var userID *types.UserID + + if req.User.Or(0) != 0 { + user, err := s.state.GetUserByID(types.UserID(req.User.Or(0))) + if err != nil { + return nil, mapStateError(err) + } + + userID = user.TypedID() + } + + preAuthKey, err := s.state.CreatePreAuthKey( + userID, + req.Reusable.Or(false), + req.Ephemeral.Or(false), + &expiration, + req.AclTags, + ) + if err != nil { + return nil, mapStateError(err) + } + + return &oas.CreatePreAuthKeyOK{ + PreAuthKey: oas.NewOptPreAuthKey(oasPreAuthKey(preAuthKey.Proto())), + }, nil +} + +// ListPreAuthKeys lists all pre-auth keys, sorted by id. +func (s *Server) ListPreAuthKeys(_ context.Context) (*oas.ListPreAuthKeysOK, error) { + keys, err := s.state.ListPreAuthKeys() + if err != nil { + return nil, mapStateError(err) + } + + slices.SortFunc(keys, func(a, b types.PreAuthKey) int { return cmp.Compare(a.ID, b.ID) }) + + out := make([]oas.PreAuthKey, len(keys)) + for i := range keys { + out[i] = oasPreAuthKey(keys[i].Proto()) + } + + return &oas.ListPreAuthKeysOK{PreAuthKeys: out}, nil +} + +// ExpirePreAuthKey expires a pre-auth key. +func (s *Server) ExpirePreAuthKey(_ context.Context, req *oas.ExpirePreAuthKeyReq) error { + err := s.state.ExpirePreAuthKey(req.ID.Or(0)) + if err != nil { + return mapStateError(err) + } + + return nil +} + +// DeletePreAuthKey deletes a pre-auth key. +func (s *Server) DeletePreAuthKey(_ context.Context, params oas.DeletePreAuthKeyParams) error { + err := s.state.DeletePreAuthKey(params.ID.Or(0)) + if err != nil { + return mapStateError(err) + } + + return nil +} + +var ( + errTagPrefix = errors.New("tag must start with the string 'tag:'") + errTagLowercase = errors.New("tag should be lowercase") + errTagSpaces = errors.New("tags must not contain spaces") +) + +// validateTag enforces the ACL tag format ("tag:" prefix, lowercase, no spaces). +func validateTag(tag string) error { + if !strings.HasPrefix(tag, "tag:") { + return errTagPrefix + } + + if strings.ToLower(tag) != tag { + return errTagLowercase + } + + if len(strings.Fields(tag)) > 1 { + return errTagSpaces + } + + return nil +} diff --git a/hscontrol/servertest/apiv1_preauthkeys_test.go b/hscontrol/servertest/apiv1_preauthkeys_test.go new file mode 100644 index 00000000..d7ab08a0 --- /dev/null +++ b/hscontrol/servertest/apiv1_preauthkeys_test.go @@ -0,0 +1,112 @@ +package servertest_test + +import ( + "context" + "net/http" + "testing" + "time" + + apiv1 "github.com/juanfont/headscale/gen/api/v1" + "github.com/juanfont/headscale/hscontrol/types" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAPIv1_CreatePreAuthKey(t *testing.T) { + srv, client := apiClient(t) + ctx := context.Background() + u := srv.CreateUser(t, "alice") + + resp, err := client.CreatePreAuthKey(ctx, &apiv1.CreatePreAuthKeyReq{ + User: apiv1.NewOptUint64(uint64(u.ID)), + Reusable: apiv1.NewOptBool(true), + AclTags: []string{"tag:test"}, + }) + require.NoError(t, err) + + pak := resp.PreAuthKey.Value + assert.NotEmpty(t, pak.Key.Value) + assert.True(t, pak.Reusable.Value) + assert.Equal(t, []string{"tag:test"}, pak.AclTags) + assert.Equal(t, uint64(u.ID), pak.User.Value.ID.Value) + + // Invalid tag format is a 400. + _, err = client.CreatePreAuthKey(ctx, &apiv1.CreatePreAuthKeyReq{ + User: apiv1.NewOptUint64(uint64(u.ID)), + AclTags: []string{"badtag"}, + }) + requireProblem(t, err, http.StatusBadRequest) +} + +func TestAPIv1_ListPreAuthKeys(t *testing.T) { + srv, client := apiClient(t) + ctx := context.Background() + u := srv.CreateUser(t, "alice") + + srv.CreatePreAuthKey(t, types.UserID(u.ID)) + srv.CreatePreAuthKey(t, types.UserID(u.ID)) + + resp, err := client.ListPreAuthKeys(ctx) + require.NoError(t, err) + require.Len(t, resp.PreAuthKeys, 2) + assert.LessOrEqual(t, resp.PreAuthKeys[0].ID.Value, resp.PreAuthKeys[1].ID.Value) +} + +func TestAPIv1_ExpirePreAuthKey(t *testing.T) { + srv, client := apiClient(t) + ctx := context.Background() + u := srv.CreateUser(t, "alice") + + resp, err := client.CreatePreAuthKey(ctx, &apiv1.CreatePreAuthKeyReq{ + User: apiv1.NewOptUint64(uint64(u.ID)), + Reusable: apiv1.NewOptBool(true), + }) + require.NoError(t, err) + + id := resp.PreAuthKey.Value.ID.Value + + require.NoError(t, client.ExpirePreAuthKey(ctx, &apiv1.ExpirePreAuthKeyReq{ + ID: apiv1.NewOptUint64(id), + })) + + keys, err := srv.State().ListPreAuthKeys() + require.NoError(t, err) + + found := false + + for _, k := range keys { + if k.ID == id { + found = true + + require.NotNil(t, k.Expiration) + assert.True(t, k.Expiration.Before(time.Now()), "key should be expired") + } + } + + require.True(t, found, "expired key should still be listed") +} + +func TestAPIv1_DeletePreAuthKey(t *testing.T) { + srv, client := apiClient(t) + ctx := context.Background() + u := srv.CreateUser(t, "alice") + + resp, err := client.CreatePreAuthKey(ctx, &apiv1.CreatePreAuthKeyReq{ + User: apiv1.NewOptUint64(uint64(u.ID)), + Reusable: apiv1.NewOptBool(true), + }) + require.NoError(t, err) + + id := resp.PreAuthKey.Value.ID.Value + + require.NoError(t, client.DeletePreAuthKey(ctx, apiv1.DeletePreAuthKeyParams{ + ID: apiv1.NewOptUint64(id), + })) + + keys, err := srv.State().ListPreAuthKeys() + require.NoError(t, err) + + for _, k := range keys { + assert.NotEqual(t, id, k.ID) + } +}