mirror of
https://github.com/juanfont/headscale.git
synced 2026-08-07 15:58:45 +09:00
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:
@@ -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()),
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user