hscontrol/api/v1: implement pre-auth key endpoints

CreatePreAuthKey (with ACL tag validation), ListPreAuthKeys, ExpirePreAuthKey,
DeletePreAuthKey over the state layer, with HTTP-parity tests.
This commit is contained in:
Kristoffer Dalby
2026-06-17 15:29:17 +00:00
parent df6876a570
commit fc26c447ec
3 changed files with 269 additions and 0 deletions
+40
View File
@@ -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()),
+117
View File
@@ -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
}
@@ -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)
}
}