diff --git a/hscontrol/api/v1/apikeys.go b/hscontrol/api/v1/apikeys.go new file mode 100644 index 00000000..5b52f864 --- /dev/null +++ b/hscontrol/api/v1/apikeys.go @@ -0,0 +1,108 @@ +package apiv1 + +import ( + "cmp" + "context" + "slices" + "time" + + oas "github.com/juanfont/headscale/gen/api/v1" + "github.com/juanfont/headscale/hscontrol/types" +) + +// CreateApiKey creates an API key and returns the full secret. This is the only +// time the secret is exposed. +func (s *Server) CreateApiKey( + _ context.Context, + req *oas.CreateApiKeyReq, +) (*oas.CreateApiKeyOK, error) { + var expiration time.Time + if v, ok := req.Expiration.Get(); ok { + expiration = v + } + + key, _, err := s.state.CreateAPIKey(&expiration) + if err != nil { + return nil, mapStateError(err) + } + + return &oas.CreateApiKeyOK{ApiKey: oas.NewOptString(key)}, nil +} + +// ListApiKeys lists all API keys (masked), sorted by id. +func (s *Server) ListApiKeys(_ context.Context) (*oas.ListApiKeysOK, error) { + keys, err := s.state.ListAPIKeys() + if err != nil { + return nil, mapStateError(err) + } + + slices.SortFunc(keys, func(a, b types.APIKey) int { return cmp.Compare(a.ID, b.ID) }) + + out := make([]oas.ApiKey, len(keys)) + for i := range keys { + out[i] = oasAPIKey(keys[i].Proto()) + } + + return &oas.ListApiKeysOK{ApiKeys: out}, nil +} + +// ExpireApiKey expires an API key identified by id or prefix. +func (s *Server) ExpireApiKey(_ context.Context, req *oas.ExpireApiKeyReq) error { + key, apiErr := s.apiKeyByIDOrPrefix(req.ID.Or(0), req.Prefix.Or("")) + if apiErr != nil { + return apiErr + } + + err := s.state.ExpireAPIKey(key) + if err != nil { + return mapStateError(err) + } + + return nil +} + +// DeleteApiKey deletes an API key identified by prefix (or id). +func (s *Server) DeleteApiKey(_ context.Context, params oas.DeleteApiKeyParams) error { + key, apiErr := s.apiKeyByIDOrPrefix(params.ID.Or(0), params.Prefix) + if apiErr != nil { + return apiErr + } + + err := s.state.DestroyAPIKey(*key) + if err != nil { + return mapStateError(err) + } + + return nil +} + +// apiKeyByIDOrPrefix looks up an API key by exactly one of id or prefix. +// Providing neither or both is a 400, matching the gRPC contract. +func (s *Server) apiKeyByIDOrPrefix( + id uint64, + prefix string, +) (*types.APIKey, *oas.ErrorStatusCode) { + hasID := id != 0 + hasPrefix := prefix != "" + + switch { + case hasID && hasPrefix: + return nil, badRequest("provide either id or prefix, not both") + case hasID: + key, err := s.state.GetAPIKeyByID(id) + if err != nil { + return nil, mapStateError(err) + } + + return key, nil + case hasPrefix: + key, err := s.state.GetAPIKey(prefix) + if err != nil { + return nil, mapStateError(err) + } + + return key, nil + default: + return nil, badRequest("must provide id or prefix") + } +} diff --git a/hscontrol/api/v1/convert.go b/hscontrol/api/v1/convert.go index f0788527..49194445 100644 --- a/hscontrol/api/v1/convert.go +++ b/hscontrol/api/v1/convert.go @@ -44,6 +44,16 @@ func optTime(ts *timestamppb.Timestamp) oas.OptDateTime { return oas.NewOptDateTime(ts.AsTime()) } +func oasAPIKey(k *v1.ApiKey) oas.ApiKey { + return oas.ApiKey{ + ID: optUint64(k.GetId()), + Prefix: optString(k.GetPrefix()), + Expiration: optTime(k.GetExpiration()), + CreatedAt: optTime(k.GetCreatedAt()), + LastSeen: optTime(k.GetLastSeen()), + } +} + func oasUser(u *v1.User) oas.User { return oas.User{ ID: optUint64(u.GetId()), diff --git a/hscontrol/servertest/apiv1_apikeys_test.go b/hscontrol/servertest/apiv1_apikeys_test.go new file mode 100644 index 00000000..3ee2cccb --- /dev/null +++ b/hscontrol/servertest/apiv1_apikeys_test.go @@ -0,0 +1,87 @@ +package servertest_test + +import ( + "context" + "net/http" + "testing" + "time" + + apiv1 "github.com/juanfont/headscale/gen/api/v1" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAPIv1_CreateApiKey(t *testing.T) { + _, client := apiClient(t) + ctx := context.Background() + + resp, err := client.CreateApiKey(ctx, &apiv1.CreateApiKeyReq{ + Expiration: apiv1.NewOptDateTime(time.Now().Add(time.Hour)), + }) + require.NoError(t, err) + assert.NotEmpty(t, resp.ApiKey.Value) +} + +func TestAPIv1_ListApiKeys(t *testing.T) { + srv, client := apiClient(t) + ctx := context.Background() + + before, err := client.ListApiKeys(ctx) + require.NoError(t, err) + + exp := time.Now().Add(time.Hour) + _, _, err = srv.State().CreateAPIKey(&exp) + require.NoError(t, err) + + after, err := client.ListApiKeys(ctx) + require.NoError(t, err) + require.Len(t, after.ApiKeys, len(before.ApiKeys)+1) + + for i := 1; i < len(after.ApiKeys); i++ { + assert.LessOrEqual(t, after.ApiKeys[i-1].ID.Value, after.ApiKeys[i].ID.Value) + } +} + +func TestAPIv1_ExpireApiKey(t *testing.T) { + srv, client := apiClient(t) + ctx := context.Background() + + exp := time.Now().Add(time.Hour) + _, key, err := srv.State().CreateAPIKey(&exp) + require.NoError(t, err) + + require.NoError(t, client.ExpireApiKey(ctx, &apiv1.ExpireApiKeyReq{ + ID: apiv1.NewOptUint64(key.ID), + })) + + got, err := srv.State().GetAPIKeyByID(key.ID) + require.NoError(t, err) + require.NotNil(t, got.Expiration) + assert.True(t, got.Expiration.Before(time.Now()), "key should be expired") + + // Both id and prefix is a 400. + requireProblem(t, client.ExpireApiKey(ctx, &apiv1.ExpireApiKeyReq{ + ID: apiv1.NewOptUint64(1), + Prefix: apiv1.NewOptString("abc"), + }), http.StatusBadRequest) + + // Neither id nor prefix is a 400. + requireProblem(t, client.ExpireApiKey(ctx, &apiv1.ExpireApiKeyReq{}), http.StatusBadRequest) +} + +func TestAPIv1_DeleteApiKey(t *testing.T) { + srv, client := apiClient(t) + ctx := context.Background() + + exp := time.Now().Add(time.Hour) + _, key, err := srv.State().CreateAPIKey(&exp) + require.NoError(t, err) + + require.NoError(t, client.DeleteApiKey(ctx, apiv1.DeleteApiKeyParams{Prefix: key.Prefix})) + + _, err = srv.State().GetAPIKeyByID(key.ID) + require.Error(t, err, "key should be gone") + + // Unknown prefix is a 404. + requireProblem(t, client.DeleteApiKey(ctx, apiv1.DeleteApiKeyParams{Prefix: "nonexistent"}), http.StatusNotFound) +}