diff --git a/PROJECTS/beginner/canary-token-generator/backend/internal/token/dto.go b/PROJECTS/beginner/canary-token-generator/backend/internal/token/dto.go index ed326d85..9165c3d2 100644 --- a/PROJECTS/beginner/canary-token-generator/backend/internal/token/dto.go +++ b/PROJECTS/beginner/canary-token-generator/backend/internal/token/dto.go @@ -6,6 +6,8 @@ package token import ( "encoding/json" "time" + + "github.com/CarterPerez-dev/cybersecurity-projects/canary-token-generator/backend/internal/event" ) type CreateRequest struct { @@ -53,3 +55,44 @@ func (t *Token) ToResponse(triggerURL, manageURL string) Response { Metadata: t.Metadata, } } + +type ManageTokenView struct { + ID string `json:"id"` + Type Type `json:"type"` + Memo string `json:"memo"` + Filename *string `json:"filename"` + AlertChannel AlertChannel `json:"alert_channel"` + CreatedAt time.Time `json:"created_at"` + TriggerCount int64 `json:"trigger_count"` + LastTriggered *time.Time `json:"last_triggered"` + Enabled bool `json:"enabled"` + TriggerURL string `json:"trigger_url"` +} + +type ManagePage struct { + NextCursor string `json:"next_cursor"` + HasMore bool `json:"has_more"` +} + +type ManageResponse struct { + Token ManageTokenView `json:"token"` + Events []event.Response `json:"events"` + EventsTotal int64 `json:"events_total"` + EventsSilencedActive int64 `json:"events_silenced_active"` + Page ManagePage `json:"page"` +} + +func (t *Token) ToManageView(triggerURL string) ManageTokenView { + return ManageTokenView{ + ID: t.ID, + Type: t.Type, + Memo: t.Memo, + Filename: t.Filename, + AlertChannel: t.AlertChannel, + CreatedAt: t.CreatedAt, + TriggerCount: t.TriggerCount, + LastTriggered: t.LastTriggered, + Enabled: t.Enabled, + TriggerURL: triggerURL, + } +} diff --git a/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler.go b/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler.go index a248fa79..ce6a3f2d 100644 --- a/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler.go +++ b/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler.go @@ -8,8 +8,10 @@ import ( "encoding/base64" "encoding/json" "errors" + "fmt" "log/slog" "net/http" + "strconv" "strings" "github.com/go-chi/chi/v5" @@ -19,7 +21,8 @@ import ( ) const ( - urlParamTokenID = "id" + urlParamTokenID = "id" + urlParamManageID = "manage_id" headerContentType = "Content-Type" headerLocation = "Location" @@ -30,17 +33,27 @@ const ( errorCodeInternalError = "INTERNAL_ERROR" errorCodeUnknownType = "UNKNOWN_TYPE" errorCodeGenerateFailed = "GENERATE_FAILED" + errorCodeNotFound = "NOT_FOUND" + errorCodeBadCursor = "BAD_CURSOR" respMessageValidation = "request validation failed" respMessageBadJSON = "invalid JSON body" respMessageInternalError = "internal server error" respMessageGenerateFailed = "artifact generation failed" respMessageUnknownType = "unknown token type" + respMessageNotFound = "not found" + respMessageBadCursor = "invalid cursor" kubeconfigPathPrefix = "/k/" createTokenBodyMaxBytes = 64 * 1024 fingerprintBodyMaxBytes = 64 * 1024 + + manageDefaultPageSize = 20 + manageMaxPageSize = 100 + + queryParamCursor = "cursor" + queryParamLimit = "limit" ) type EventRecorder interface { @@ -55,10 +68,25 @@ type FingerprintRecorder interface { ) error } +type EventQuery interface { + ListByToken( + ctx context.Context, + tokenID string, + opts event.ListOptions, + ) (event.ListResult, error) + CountByToken(ctx context.Context, tokenID string) (int64, error) +} + +type DedupCounter interface { + CountActiveDedup(ctx context.Context, tokenID string) (int64, error) +} + type Handler struct { svc *Service events EventRecorder fingerprintRecorder FingerprintRecorder + eventQuery EventQuery + dedupCounter DedupCounter logger *slog.Logger } @@ -66,6 +94,8 @@ func NewHandler( svc *Service, events EventRecorder, fingerprint FingerprintRecorder, + eventQuery EventQuery, + dedupCounter DedupCounter, logger *slog.Logger, ) *Handler { if logger == nil { @@ -75,6 +105,8 @@ func NewHandler( svc: svc, events: events, fingerprintRecorder: fingerprint, + eventQuery: eventQuery, + dedupCounter: dedupCounter, logger: logger, } } @@ -84,6 +116,11 @@ func (h *Handler) RegisterAPIRoutes(r chi.Router) { r.Post("/tokens", h.CreateToken) } +func (h *Handler) RegisterManageRoutes(r chi.Router) { + r.Get("/m/{"+urlParamManageID+"}", h.GetManage) + r.Delete("/m/{"+urlParamManageID+"}", h.DeleteManage) +} + func (h *Handler) RegisterTriggerRoutes(r chi.Router) { r.Get("/c/{"+urlParamTokenID+"}", h.HandleTrigger) r.Post("/c/{"+urlParamTokenID+"}/fingerprint", h.HandleFingerprint) @@ -168,6 +205,158 @@ func (h *Handler) HandleTrigger(w http.ResponseWriter, r *http.Request) { h.writeTriggerResponse(w, r, resp) } +func (h *Handler) GetManage(w http.ResponseWriter, r *http.Request) { + manageID := chi.URLParam(r, urlParamManageID) + if manageID == "" { + h.writeJSON(w, http.StatusNotFound, envelopeError( + errorCodeNotFound, respMessageNotFound, + )) + return + } + + cursor, err := parseCursor(r.URL.Query().Get(queryParamCursor)) + if err != nil { + h.writeJSON(w, http.StatusBadRequest, envelopeError( + errorCodeBadCursor, respMessageBadCursor, + )) + return + } + limit := parseLimit(r.URL.Query().Get(queryParamLimit)) + + tok, err := h.svc.GetByManageID(r.Context(), manageID) + if err != nil { + h.logger.ErrorContext(r.Context(), "manage: get by manage id", + "error", err, "manage_id", manageID) + h.writeJSON(w, http.StatusInternalServerError, envelopeError( + errorCodeInternalError, respMessageInternalError, + )) + return + } + if tok == nil { + h.writeJSON(w, http.StatusNotFound, envelopeError( + errorCodeNotFound, respMessageNotFound, + )) + return + } + + events, total, silenced := h.gatherManageData(r, tok.ID, cursor, limit) + + resp := ManageResponse{ + Token: tok.ToManageView(h.svc.TriggerURL(tok.ID)), + Events: events, + EventsTotal: total, + EventsSilencedActive: silenced, + Page: buildPage(events, limit), + } + h.writeJSON(w, http.StatusOK, envelopeData(resp)) +} + +func (h *Handler) DeleteManage(w http.ResponseWriter, r *http.Request) { + manageID := chi.URLParam(r, urlParamManageID) + if manageID == "" { + h.writeJSON(w, http.StatusNotFound, envelopeError( + errorCodeNotFound, respMessageNotFound, + )) + return + } + + if err := h.svc.DeleteByManageID(r.Context(), manageID); err != nil { + if errors.Is(err, ErrNotFound) { + h.writeJSON(w, http.StatusNotFound, envelopeError( + errorCodeNotFound, respMessageNotFound, + )) + return + } + h.logger.ErrorContext(r.Context(), "manage: delete", + "error", err, "manage_id", manageID) + h.writeJSON(w, http.StatusInternalServerError, envelopeError( + errorCodeInternalError, respMessageInternalError, + )) + return + } + w.WriteHeader(http.StatusNoContent) +} + +func (h *Handler) gatherManageData( + r *http.Request, + tokenID string, + cursor int64, + limit int, +) (events []event.Response, total, silenced int64) { + if h.eventQuery == nil { + return nil, 0, 0 + } + list, err := h.eventQuery.ListByToken( + r.Context(), tokenID, event.ListOptions{Cursor: cursor, Limit: limit}, + ) + if err != nil { + h.logger.ErrorContext(r.Context(), "manage: list events", + "error", err, "token_id", tokenID) + } + for i := range list.Events { + events = append(events, list.Events[i].ToResponse()) + } + + if total, err = h.eventQuery.CountByToken( + r.Context(), + tokenID, + ); err != nil { + h.logger.ErrorContext(r.Context(), "manage: count events", + "error", err, "token_id", tokenID) + total = 0 + } + + if h.dedupCounter != nil { + var dErr error + if silenced, dErr = h.dedupCounter.CountActiveDedup( + r.Context(), tokenID, + ); dErr != nil { + h.logger.WarnContext(r.Context(), "manage: count active dedup", + "error", dErr, "token_id", tokenID) + silenced = 0 + } + } + return events, total, silenced +} + +func buildPage(events []event.Response, limit int) ManagePage { + if len(events) < limit || len(events) == 0 { + return ManagePage{} + } + last := events[len(events)-1] + return ManagePage{ + NextCursor: strconv.FormatInt(last.ID, 10), + HasMore: true, + } +} + +func parseCursor(raw string) (int64, error) { + raw = strings.TrimSpace(raw) + if raw == "" { + return 0, nil + } + v, err := strconv.ParseInt(raw, 10, 64) + if err != nil || v < 0 { + return 0, fmt.Errorf("invalid cursor: %q", raw) + } + return v, nil +} + +func parseLimit(raw string) int { + raw = strings.TrimSpace(raw) + if raw == "" { + return manageDefaultPageSize + } + v, err := strconv.Atoi(raw) + if err != nil || v <= 0 { + return manageDefaultPageSize + } + if v > manageMaxPageSize { + return manageMaxPageSize + } + return v +} + func (h *Handler) HandleFingerprint(w http.ResponseWriter, r *http.Request) { id := chi.URLParam(r, urlParamTokenID) if id == "" || h.fingerprintRecorder == nil { diff --git a/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler_test.go b/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler_test.go index 3f6c3ab3..6f59b268 100644 --- a/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler_test.go +++ b/PROJECTS/beginner/canary-token-generator/backend/internal/token/handler_test.go @@ -78,13 +78,20 @@ func newWebbugHandler( ManageURL: "https://canary.example.com", }, ) - return token.NewHandler(svc, rec, nil, quietHandlerLogger()), repo, rec + return token.NewHandler( + svc, + rec, + nil, + nil, + nil, + quietHandlerLogger(), + ), repo, rec } func TestGetTypes_Returns7Types(t *testing.T) { svc := token.NewService(newFakeRepo(), token.MapRegistry{}, token.ServiceConfig{BaseURL: "https://x.test"}) - h := token.NewHandler(svc, nil, nil, quietHandlerLogger()) + h := token.NewHandler(svc, nil, nil, nil, nil, quietHandlerLogger()) r := chi.NewRouter() h.RegisterAPIRoutes(r) @@ -348,7 +355,7 @@ func exposeArtifactToJSON(a generators.Artifact) token.ArtifactJSON { repo := newFakeRepo() svc := token.NewService(repo, token.MapRegistry{token.TypeWebbug: gen}, token.ServiceConfig{BaseURL: "https://x"}) - h := token.NewHandler(svc, nil, nil, quietHandlerLogger()) + h := token.NewHandler(svc, nil, nil, nil, nil, quietHandlerLogger()) r := chi.NewRouter() h.RegisterAPIRoutes(r) @@ -369,3 +376,268 @@ func exposeArtifactToJSON(a generators.Artifact) token.ArtifactJSON { } return resp.Data.Artifact } + +type fakeEventQuery struct { + listResult event.ListResult + listErr error + countN int64 + countErr error + lastTokID string + lastOpts event.ListOptions +} + +func (f *fakeEventQuery) ListByToken( + _ context.Context, + tokenID string, + opts event.ListOptions, +) (event.ListResult, error) { + f.lastTokID = tokenID + f.lastOpts = opts + return f.listResult, f.listErr +} + +func (f *fakeEventQuery) CountByToken( + _ context.Context, + _ string, +) (int64, error) { + return f.countN, f.countErr +} + +type fakeDedupCounter struct { + n int64 + err error +} + +func (f *fakeDedupCounter) CountActiveDedup( + _ context.Context, + _ string, +) (int64, error) { + return f.n, f.err +} + +func newManageHandler( + t *testing.T, + eq *fakeEventQuery, + dc *fakeDedupCounter, +) (*token.Handler, *fakeRepo) { + t.Helper() + repo := newFakeRepo() + svc := token.NewService(repo, token.MapRegistry{}, + token.ServiceConfig{ + BaseURL: "https://canary.example.com", + ManageURL: "https://canary.example.com", + }) + return token.NewHandler(svc, nil, nil, eq, dc, quietHandlerLogger()), repo +} + +func seedToken(t *testing.T, repo *fakeRepo, manageID string) *token.Token { + t.Helper() + tok := &token.Token{ + ID: "tok" + manageID[:8], + ManageID: manageID, + Type: token.TypeWebbug, + Memo: "manage-test", + AlertChannel: token.ChannelWebhook, + WebhookURL: strPtr("https://x/h"), + CreatedIP: "1.1.1.1", + CreatedFP: "fp", + Metadata: json.RawMessage(`{}`), + Enabled: true, + TriggerCount: 5, + } + require.NoError(t, repo.Insert(context.Background(), tok)) + return tok +} + +func strPtr(s string) *string { return &s } + +func decodeManage(t *testing.T, body []byte) struct { + Success bool `json:"success"` + Data token.ManageResponse `json:"data"` +} { + t.Helper() + var resp struct { + Success bool `json:"success"` + Data token.ManageResponse `json:"data"` + } + require.NoError(t, json.Unmarshal(body, &resp)) + return resp +} + +func TestGetManage_HappyPath(t *testing.T) { + t.Parallel() + eq := &fakeEventQuery{ + countN: 17, + listResult: event.ListResult{ + Events: []event.Event{ + {ID: 42, TokenID: "x", SourceIP: "1.2.3.4"}, + {ID: 41, TokenID: "x", SourceIP: "5.6.7.8"}, + }, + HasMore: false, + }, + } + dc := &fakeDedupCounter{n: 3} + h, repo := newManageHandler(t, eq, dc) + tok := seedToken(t, repo, "11111111-1111-1111-1111-111111111111") + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + req := httptest.NewRequest(http.MethodGet, "/m/"+tok.ManageID, nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + + require.Equal(t, http.StatusOK, w.Code, "body=%s", w.Body.String()) + resp := decodeManage(t, w.Body.Bytes()) + require.True(t, resp.Success) + require.Equal(t, tok.ID, resp.Data.Token.ID) + require.Equal( + t, + "https://canary.example.com/c/"+tok.ID, + resp.Data.Token.TriggerURL, + ) + require.Equal(t, int64(5), resp.Data.Token.TriggerCount) + require.Len(t, resp.Data.Events, 2) + require.Equal(t, int64(17), resp.Data.EventsTotal) + require.Equal(t, int64(3), resp.Data.EventsSilencedActive) +} + +func TestGetManage_404OnUnknownManageID(t *testing.T) { + t.Parallel() + h, _ := newManageHandler(t, &fakeEventQuery{}, &fakeDedupCounter{}) + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP( + w, + httptest.NewRequest(http.MethodGet, "/m/does-not-exist", nil), + ) + require.Equal(t, http.StatusNotFound, w.Code) + require.Contains(t, w.Body.String(), "NOT_FOUND") +} + +func TestGetManage_400OnBadCursor(t *testing.T) { + t.Parallel() + h, repo := newManageHandler(t, &fakeEventQuery{}, &fakeDedupCounter{}) + tok := seedToken(t, repo, "22222222-2222-2222-2222-222222222222") + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, + "/m/"+tok.ManageID+"?cursor=notanumber", nil)) + require.Equal(t, http.StatusBadRequest, w.Code) + require.Contains(t, w.Body.String(), "BAD_CURSOR") +} + +func TestGetManage_400OnNegativeCursor(t *testing.T) { + t.Parallel() + h, repo := newManageHandler(t, &fakeEventQuery{}, &fakeDedupCounter{}) + tok := seedToken(t, repo, "33333333-3333-3333-3333-333333333333") + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, + "/m/"+tok.ManageID+"?cursor=-1", nil)) + require.Equal(t, http.StatusBadRequest, w.Code) +} + +func TestGetManage_PaginationCursorAndHasMore(t *testing.T) { + t.Parallel() + full := []event.Event{} + for i := 20; i > 0; i-- { + full = append(full, event.Event{ID: int64(i), SourceIP: "1.1.1.1"}) + } + eq := &fakeEventQuery{ + listResult: event.ListResult{ + Events: full, + HasMore: true, + NextCursor: 1, + }, + countN: 50, + } + h, repo := newManageHandler(t, eq, &fakeDedupCounter{}) + tok := seedToken(t, repo, "44444444-4444-4444-4444-444444444444") + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/m/"+tok.ManageID, nil)) + require.Equal(t, http.StatusOK, w.Code) + + resp := decodeManage(t, w.Body.Bytes()) + require.Len(t, resp.Data.Events, 20) + require.True(t, resp.Data.Page.HasMore) + require.Equal(t, "1", resp.Data.Page.NextCursor, + "cursor is the ID of the last event returned") +} + +func TestGetManage_LimitParamRespected(t *testing.T) { + t.Parallel() + eq := &fakeEventQuery{} + h, repo := newManageHandler(t, eq, &fakeDedupCounter{}) + tok := seedToken(t, repo, "55555555-5555-5555-5555-555555555555") + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, + "/m/"+tok.ManageID+"?limit=5", nil)) + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, 5, eq.lastOpts.Limit) +} + +func TestGetManage_LimitCappedAtMax(t *testing.T) { + t.Parallel() + eq := &fakeEventQuery{} + h, repo := newManageHandler(t, eq, &fakeDedupCounter{}) + tok := seedToken(t, repo, "66666666-6666-6666-6666-666666666666") + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP(w, httptest.NewRequest(http.MethodGet, + "/m/"+tok.ManageID+"?limit=999", nil)) + require.Equal(t, http.StatusOK, w.Code) + require.LessOrEqual(t, eq.lastOpts.Limit, 100, + "limit should be capped at manageMaxPageSize") +} + +func TestDeleteManage_HappyPath(t *testing.T) { + t.Parallel() + h, repo := newManageHandler(t, &fakeEventQuery{}, &fakeDedupCounter{}) + tok := seedToken(t, repo, "77777777-7777-7777-7777-777777777777") + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP( + w, + httptest.NewRequest(http.MethodDelete, "/m/"+tok.ManageID, nil), + ) + require.Equal(t, http.StatusNoContent, w.Code) +} + +func TestDeleteManage_404OnUnknownManageID(t *testing.T) { + t.Parallel() + h, _ := newManageHandler(t, &fakeEventQuery{}, &fakeDedupCounter{}) + + r := chi.NewRouter() + h.RegisterManageRoutes(r) + + w := httptest.NewRecorder() + r.ServeHTTP( + w, + httptest.NewRequest(http.MethodDelete, "/m/does-not-exist", nil), + ) + require.Equal(t, http.StatusNotFound, w.Code) +}