From 0b057df2a3a764db0bfc5992939f615274112387 Mon Sep 17 00:00:00 2001 From: claude Date: Sat, 15 Aug 2026 17:19:13 +0400 Subject: [PATCH] Give reminder cancellation its own store and IPC path (V-719) CancelReminder replaces the cancelled half of MarkReminder, which stays delivery-only. Cancellation has to win against the start of an external send, so it refuses when the occurrence has a pending, sent or unknown outbox row, and clears the delivery group inside the same transaction. BeginDeliveryAttempt takes the mirror lock for reminder sends, so no interleaving lets both operations report success. Cancelling one member of a collapsed catch-up bundle invalidates the cached phrase on every pending sibling; a later retry would otherwise keep saying "three reminders" after one was removed. Legacy rows carry the empty delivery group from migration 25, so they only count as this occurrence when they began at or after its next-fire boundary. Without that bound one old success would make a recurring series permanently uncancellable. ListPendingReminders returns cancellable rows in firing order, with no limit by default, because spoken resolution must not miss an old reminder that newer fired history pushed out of ListReminders' window. Cancellation is ordinary authenticated write authority: it prevents a future send and cannot create one. cmd/e2eprobe drives both from outside. --no-verify: master is the working branch this session by the owner's call. Co-Authored-By: Claude Opus 5 --- cmd/e2eprobe/main.go | 284 ++++++++++++++++++++++ internal/auth/auth_test.go | 4 +- internal/auth/policy.go | 17 +- internal/ipc/api.go | 4 +- internal/ipc/client.go | 15 ++ internal/ipc/coreapi.go | 2 + internal/ipc/ipc_test.go | 49 +++- internal/ipc/maperr_test.go | 1 + internal/ipc/server.go | 6 + internal/ipc/storeapi.go | 11 + internal/ipc/unimplemented.go | 6 + internal/ipc/wire.go | 31 ++- internal/store/delivery.go | 56 +++++ internal/store/delivery_test.go | 15 +- internal/store/reminders.go | 116 ++++++++- internal/store/reminders_cancel_test.go | 243 ++++++++++++++++++ internal/store/reminders_delivery_test.go | 24 +- 17 files changed, 846 insertions(+), 38 deletions(-) create mode 100644 cmd/e2eprobe/main.go create mode 100644 internal/store/reminders_cancel_test.go diff --git a/cmd/e2eprobe/main.go b/cmd/e2eprobe/main.go new file mode 100644 index 0000000..e2a80aa --- /dev/null +++ b/cmd/e2eprobe/main.go @@ -0,0 +1,284 @@ +// e2eprobe is a temporary typed IPC driver used by the 2026-08-15 isolated +// whole-Maven acceptance session. It is removed after the session; keeping the +// driver inside the module lets it import Maven's internal IPC contract rather +// than peeking into sqlite. +package main + +import ( + "context" + "encoding/json" + "errors" + "flag" + "fmt" + "math" + "os" + "strconv" + "strings" + "time" + + "github.com/kami/maven/internal/ipc" + "github.com/kami/maven/internal/router" + "github.com/kami/maven/internal/store" +) + +func main() { + if err := run(os.Args[1:]); err != nil { + fmt.Fprintln(os.Stderr, "e2eprobe:", err) + os.Exit(1) + } +} + +func run(args []string) error { + fs := flag.NewFlagSet("e2eprobe", flag.ContinueOnError) + sock := fs.String("sock", "", "mavend unix socket") + if err := fs.Parse(args); err != nil { + return err + } + argv := fs.Args() + if len(argv) == 0 { + return errors.New("usage: e2eprobe -sock PATH COMMAND [ARGS]") + } + ctx, cancel := context.WithTimeout(context.Background(), 90*time.Second) + defer cancel() + if argv[0] == "score-pair" { + out, err := scorePair(ctx, argv) + if err != nil { + return err + } + return encode(out) + } + if argv[0] == "parse-task-status" { + if len(argv) != 2 { + return errors.New("parse-task-status needs TEXT") + } + parsed, ok := router.ParseTaskStatus(argv[1]) + return encode(map[string]any{"accepted": ok, "parsed": parsed}) + } + if *sock == "" { + return errors.New("usage: e2eprobe -sock PATH COMMAND [ARGS]") + } + + cli, err := ipc.DialWait(*sock, 15*time.Second) + if err != nil { + return err + } + defer cli.Close() + + var out any + switch argv[0] { + case "ping": + out, err = cli.Ping(ctx) + case "chat": + if len(argv) < 3 { + return errors.New("chat needs CONVERSATION TEXT") + } + out, err = cli.Chat(ctx, argv[1], strings.Join(argv[2:], " ")) + case "create-reminder": + if len(argv) < 3 || len(argv) > 4 { + return errors.New("create-reminder needs RFC3339 TEXT [CRON]") + } + fire, parseErr := time.Parse(time.RFC3339, argv[1]) + if parseErr != nil { + return parseErr + } + cron := "" + if len(argv) == 4 { + cron = argv[3] + } + id, createErr := cli.CreateReminder(ctx, fire, `{"text":`+quote(argv[2])+`}`, cron) + out, err = map[string]any{"id": id}, createErr + case "cancel-reminder": + id, parseErr := oneID(argv) + if parseErr != nil { + return parseErr + } + err = cli.CancelReminder(ctx, id) + out = map[string]any{"cancelled": id} + case "mark-reminder": + if len(argv) != 3 { + return errors.New("mark-reminder needs ID STATUS") + } + id, parseErr := strconv.ParseInt(argv[1], 10, 64) + if parseErr != nil { + return parseErr + } + err = cli.MarkReminder(ctx, id, argv[2]) + out = map[string]any{"marked": id, "status": argv[2]} + case "reminders": + n, parseErr := optionalN(argv, 200) + if parseErr != nil { + return parseErr + } + out, err = cli.ListReminders(ctx, n) + case "pending-reminders": + n, parseErr := optionalN(argv, 0) + if parseErr != nil { + return parseErr + } + out, err = cli.ListPendingReminders(ctx, n) + case "create-task": + if len(argv) != 2 { + return errors.New("create-task needs TEXT") + } + out, err = cli.CaptureTask(ctx, ipc.CaptureTaskReq{ + Text: argv[1], Source: "tap:web", Status: store.TaskOpen, Ts: time.Now(), + }) + case "tasks": + status := "live" + if len(argv) == 2 { + status = argv[1] + } else if len(argv) != 1 { + return errors.New("tasks takes optional STATUS") + } + out, err = cli.ListTasks(ctx, status) + case "notes": + n, parseErr := optionalN(argv, 50) + if parseErr != nil { + return parseErr + } + out, err = cli.RecentNotes(ctx, n) + case "query-notes": + if len(argv) != 2 { + return errors.New("query-notes needs TEXT") + } + embedder, embedErr := router.NewONNXEmbedder( + "models/embedder/multilingual-e5-small/model_quantized.onnx", + "models/embedder/multilingual-e5-small/tokenizer.json", + "deps/onnxruntime-linux-x64-1.26.0/lib/libonnxruntime.so", + ) + if embedErr != nil { + return embedErr + } + defer embedder.Close() + vec, embedErr := router.EmbedQuery(ctx, embedder, argv[1]) + if embedErr != nil { + return embedErr + } + out, err = cli.QueryNotes(ctx, vec, 10) + case "score-pair": + out, err = scorePair(ctx, argv) + case "facts": + n, parseErr := optionalN(argv, 50) + if parseErr != nil { + return parseErr + } + out, err = cli.RecentFacts(ctx, n) + case "decisions": + n, parseErr := optionalN(argv, 50) + if parseErr != nil { + return parseErr + } + out, err = cli.TurnDecisions(ctx, n) + case "events": + n, parseErr := optionalN(argv, 50) + if parseErr != nil { + return parseErr + } + out, err = cli.RecentEvents(ctx, n) + case "eco-traces": + n, parseErr := optionalN(argv, 50) + if parseErr != nil { + return parseErr + } + out, err = cli.RecentEcosystemTraces(ctx, n) + case "delivery-attempts": + status := "" + if len(argv) == 2 { + status = argv[1] + } else if len(argv) != 1 { + return errors.New("delivery-attempts takes optional STATUS") + } + out, err = cli.DeliveryAttempts(ctx, status, 200) + case "nudges": + n, parseErr := optionalN(argv, 50) + if parseErr != nil { + return parseErr + } + out, err = cli.RecentNudges(ctx, n) + case "tools": + status := "" + if len(argv) == 2 { + status = argv[1] + } else if len(argv) != 1 { + return errors.New("tools takes optional STATUS") + } + out, err = cli.ListTools(ctx, status) + case "plan": + out, err = cli.DayPlan(ctx) + case "correct": + if len(argv) != 3 { + return errors.New("correct needs TRACE_ID SHOULD_BE") + } + id, parseErr := strconv.ParseInt(argv[1], 10, 64) + if parseErr != nil { + return parseErr + } + err = cli.CorrectTurn(ctx, id, argv[2]) + out = map[string]any{"corrected": id, "should_be": argv[2]} + default: + return fmt.Errorf("unknown command %q", argv[0]) + } + if err != nil { + return err + } + return encode(out) +} + +func encode(out any) error { + enc := json.NewEncoder(os.Stdout) + enc.SetIndent("", " ") + return enc.Encode(out) +} + +func scorePair(ctx context.Context, argv []string) (any, error) { + if len(argv) != 3 { + return nil, errors.New("score-pair needs QUERY PASSAGE") + } + embedder, err := router.NewONNXEmbedder( + "models/embedder/multilingual-e5-small/model_quantized.onnx", + "models/embedder/multilingual-e5-small/tokenizer.json", + "deps/onnxruntime-linux-x64-1.26.0/lib/libonnxruntime.so", + ) + if err != nil { + return nil, err + } + defer embedder.Close() + qvec, err := router.EmbedQuery(ctx, embedder, argv[1]) + if err != nil { + return nil, err + } + pvec, err := router.EmbedPassage(ctx, embedder, argv[2]) + if err != nil { + return nil, err + } + if len(qvec) != len(pvec) { + return nil, fmt.Errorf("embedding widths differ: %d != %d", len(qvec), len(pvec)) + } + var dot float64 + for i := range qvec { + dot += float64(qvec[i]) * float64(pvec[i]) + } + return map[string]any{"score": math.Round(dot*1e9) / 1e9}, nil +} + +func quote(s string) string { + b, _ := json.Marshal(s) + return string(b) +} + +func oneID(argv []string) (int64, error) { + if len(argv) != 2 { + return 0, errors.New("command needs ID") + } + return strconv.ParseInt(argv[1], 10, 64) +} + +func optionalN(argv []string, fallback int) (int, error) { + if len(argv) == 1 { + return fallback, nil + } + if len(argv) != 2 { + return 0, errors.New("command takes optional N") + } + return strconv.Atoi(argv[1]) +} diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go index 569339b..79a2996 100644 --- a/internal/auth/auth_test.go +++ b/internal/auth/auth_test.go @@ -68,7 +68,8 @@ func TestRequirement_Table(t *testing.T) { reads := []ipc.Method{ ipc.MethodLatestFact, ipc.MethodLatestFactBySource, ipc.MethodSince, ipc.MethodPresence, ipc.MethodRecentOutcomes, - ipc.MethodCreateReminder, ipc.MethodMarkReminder, + ipc.MethodCreateReminder, + ipc.MethodListPendingReminders, ipc.MethodRecordNudge, ipc.MethodResolveNudge, ipc.MethodChat, } @@ -424,6 +425,7 @@ func TestRequirement_SwapModel(t *testing.T) { func TestRequirement_ListMutation(t *testing.T) { for _, m := range []ipc.Method{ ipc.MethodIngestMail, ipc.MethodSetTaskStatus, + ipc.MethodMarkReminder, ipc.MethodCancelReminder, } { if got := Requirement(m); got != AuthWrite { t.Errorf("%s authority = %v; want AuthWrite", m, got) diff --git a/internal/auth/policy.go b/internal/auth/policy.go index 510a1f1..13ebd00 100644 --- a/internal/auth/policy.go +++ b/internal/auth/policy.go @@ -19,9 +19,10 @@ type Authority int8 const ( // AuthRead — read methods (LatestFact, LatestFactBySource, Since, Presence, - // RecentOutcomes) and state mutations a module legitimately makes - // (CreateReminder, MarkReminder, RecordNudge, ResolveNudge). The Enrollment - // already gated caller identity; any enrolled module may use these. + // RecentOutcomes) and additive state mutations a module legitimately makes + // (CreateReminder, RecordNudge, ResolveNudge). The Enrollment already gated + // caller identity; any enrolled module may use these. Terminal list changes + // such as MarkReminder and CancelReminder sit at AuthWrite below. AuthRead Authority = 0 // AuthWrite — WriteFact. Need enrollment + source-scope match. The @@ -128,6 +129,14 @@ func Requirement(m ipc.Method) Authority { // from under him. Same reasoning as WriteFact: a module gets to add to // its own corner, not to erase his. return AuthWrite + case ipc.MethodMarkReminder, ipc.MethodCancelReminder: + // Cancelling removes a standing reason for Maven to speak. It is a + // mutation, like resolving a task, but not a step-up act: both voice and + // the web must be able to make Maven quieter at the owner's request. + // MarkReminder now accepts only the delivery-side fired transition, but + // it is terminal too and therefore belongs on the same enrolled-writer + // rung rather than the generic reader rung. + return AuthWrite case ipc.MethodAssertStepUp: return AuthRead case ipc.MethodLatestFact, @@ -136,7 +145,7 @@ func Requirement(m ipc.Method) Authority { ipc.MethodPresence, ipc.MethodRecentOutcomes, ipc.MethodCreateReminder, - ipc.MethodMarkReminder, + ipc.MethodListPendingReminders, ipc.MethodRecordNudge, ipc.MethodResolveNudge, // Task capture (Vikunja #130). Listed explicitly rather than left to diff --git a/internal/ipc/api.go b/internal/ipc/api.go index cdf8074..15b4c6c 100644 --- a/internal/ipc/api.go +++ b/internal/ipc/api.go @@ -605,7 +605,8 @@ type idReq struct { ID int64 `json:"id"` } -// markReminderReq — pending→fired|cancelled. +// markReminderReq — pending→fired. Cancellation has its own transactional +// method because it must serialize against delivery attempts. type markReminderReq struct { ID int64 `json:"id"` Status string `json:"status"` @@ -1000,6 +1001,7 @@ var ( ErrNudgeOutcome = errors.New("ipc: nudge already resolved") ErrReminderNotFound = errors.New("ipc: reminder not found") ErrReminderState = errors.New("ipc: reminder not in a mutable state") + ErrReminderInFlight = errors.New("ipc: reminder delivery already started") ErrUnknownMethod = errors.New("ipc: unknown method") ErrBadParams = errors.New("ipc: bad params") // ErrForbidden — the caller's authority doesn't cover this call. The diff --git a/internal/ipc/client.go b/internal/ipc/client.go index 08c885c..460554d 100644 --- a/internal/ipc/client.go +++ b/internal/ipc/client.go @@ -74,6 +74,7 @@ var readOnlyMethods = map[Method]bool{ MethodSince: true, MethodPresence: true, MethodListReminders: true, + MethodListPendingReminders: true, MethodRecentOutcomes: true, MethodRecentFacts: true, MethodRecentActiveFacts: true, @@ -327,6 +328,8 @@ func hydrate(e *RpcError) error { return fmt.Errorf("%w: %s", ErrReminderNotFound, e.Message) case codeReminderState: return fmt.Errorf("%w: %s", ErrReminderState, e.Message) + case codeReminderInFlight: + return fmt.Errorf("%w: %s", ErrReminderInFlight, e.Message) case codeToolNotFound: return fmt.Errorf("%w: %s", ErrToolNotFound, e.Message) case codeNoSuchTrace: @@ -398,6 +401,10 @@ func (c *Client) MarkReminder(ctx context.Context, id int64, status string) erro return c.call(ctx, MethodMarkReminder, markReminderReq{ID: id, Status: status}, nil) } +func (c *Client) CancelReminder(ctx context.Context, id int64) error { + return c.call(ctx, MethodCancelReminder, idReq{ID: id}, nil) +} + func (c *Client) ListReminders(ctx context.Context, n int) ([]Reminder, error) { var out []Reminder if err := c.call(ctx, MethodListReminders, nReq{N: n}, &out); err != nil { @@ -406,6 +413,14 @@ func (c *Client) ListReminders(ctx context.Context, n int) ([]Reminder, error) { return out, nil } +func (c *Client) ListPendingReminders(ctx context.Context, n int) ([]Reminder, error) { + var out []Reminder + if err := c.call(ctx, MethodListPendingReminders, nReq{N: n}, &out); err != nil { + return nil, err + } + return out, nil +} + func (c *Client) RecordNudge(ctx context.Context, rule, channel, message string, ts time.Time) (int64, error) { var r idResp if err := c.call(ctx, MethodRecordNudge, recordNudgeReq{Rule: rule, Channel: channel, Message: message, Ts: ts}, &r); err != nil { diff --git a/internal/ipc/coreapi.go b/internal/ipc/coreapi.go index cf3d5dc..10c66d6 100644 --- a/internal/ipc/coreapi.go +++ b/internal/ipc/coreapi.go @@ -42,7 +42,9 @@ type FactAPI interface { type ReminderAPI interface { CreateReminder(ctx context.Context, fire time.Time, payload, cron string) (int64, error) MarkReminder(ctx context.Context, id int64, status string) error + CancelReminder(ctx context.Context, id int64) error ListReminders(ctx context.Context, n int) ([]Reminder, error) + ListPendingReminders(ctx context.Context, n int) ([]Reminder, error) } // NudgeAPI — proactive sends Maven proposed, their outcomes, and the outbox diff --git a/internal/ipc/ipc_test.go b/internal/ipc/ipc_test.go index 87b3c34..80b8e84 100644 --- a/internal/ipc/ipc_test.go +++ b/internal/ipc/ipc_test.go @@ -245,6 +245,9 @@ func TestStoreAPI_Direct(t *testing.T) { if err := api.MarkReminder(ctx, 99999, "weird"); !errors.Is(err, ErrReminderState) { t.Fatalf("MarkReminder weird: got %v, want ErrReminderState", err) } + if err := api.CancelReminder(ctx, 99999); !errors.Is(err, ErrReminderNotFound) { + t.Fatalf("CancelReminder missing: got %v, want ErrReminderNotFound", err) + } // resolve nonexistent nudge ⇒ ErrNudgeNotFound if err := api.ResolveNudge(ctx, 99999, "acted", time.Now()); !errors.Is(err, ErrNudgeNotFound) { t.Fatalf("ResolveNudge none: got %v, want ErrNudgeNotFound", err) @@ -256,7 +259,7 @@ func TestStoreAPI_Direct(t *testing.T) { // the test that catches the boundary bugs: param shape mismatch, sentinel // code drift, dto mapping, framing interleaving. func TestClient_E2E(t *testing.T) { - _, _, cli, _ := newServerWithStore(t) + _, _, cli, st := newServerWithStore(t) ctx := context.Background() now := time.Now().UTC().Truncate(time.Millisecond) @@ -311,11 +314,31 @@ func TestClient_E2E(t *testing.T) { t.Fatalf("Presence cold-start = %+v, want away/0", pres) } - // reminder lifecycle: create → mark fired → re-mark ⇒ ErrReminderState. + // reminder lifecycle: both pending rows cross the wire in firing order; + // one cancels and disappears from that read, while the other still follows + // the existing fired transition. + cancelID, err := cli.CreateReminder(ctx, now.Add(30*time.Minute), `{"text":"cancel me"}`, "") + if err != nil { + t.Fatalf("CreateReminder cancellation candidate: %v", err) + } rid, err := cli.CreateReminder(ctx, now.Add(time.Hour), `{"text":"wake me 7"}`, "") if err != nil { t.Fatalf("CreateReminder: %v", err) } + pending, err := cli.ListPendingReminders(ctx, 0) + if err != nil || len(pending) != 2 || pending[0].ID != cancelID || pending[1].ID != rid { + t.Fatalf("ListPendingReminders = %+v, %v", pending, err) + } + if err := cli.CancelReminder(ctx, cancelID); err != nil { + t.Fatalf("CancelReminder: %v", err) + } + if err := cli.CancelReminder(ctx, cancelID); !errors.Is(err, ErrReminderState) { + t.Fatalf("CancelReminder twice: got %v, want ErrReminderState", err) + } + pending, err = cli.ListPendingReminders(ctx, 0) + if err != nil || len(pending) != 1 || pending[0].ID != rid { + t.Fatalf("pending after cancellation = %+v, %v", pending, err) + } if err := cli.MarkReminder(ctx, rid, "fired"); err != nil { t.Fatalf("MarkReminder fired: %v", err) } @@ -323,6 +346,28 @@ func TestClient_E2E(t *testing.T) { t.Fatalf("MarkReminder twice: got %v, want ErrReminderState", err) } + flightID, err := cli.CreateReminder(ctx, now.Add(2*time.Hour), `{"text":"already leaving"}`, "") + if err != nil { + t.Fatal(err) + } + storePending, err := st.ListPendingReminders(ctx, 0) + if err != nil || len(storePending) != 1 || storePending[0].ID != flightID { + t.Fatalf("store pending = %+v, %v", storePending, err) + } + const deliveryGroup = "reminder:ipc-in-flight" + if err := st.CacheReminderPhrase(ctx, storePending, deliveryGroup, "leaving", "leaving", "neutral"); err != nil { + t.Fatal(err) + } + if _, err := st.BeginDeliveryAttempt(ctx, "reminder", "", flightID, deliveryGroup, "telegram", "hash", now); err != nil { + t.Fatal(err) + } + if err := cli.MarkReminder(ctx, flightID, "cancelled"); !errors.Is(err, ErrReminderState) { + t.Fatalf("legacy MarkReminder cancellation bypass: got %v, want ErrReminderState", err) + } + if err := cli.CancelReminder(ctx, flightID); !errors.Is(err, ErrReminderInFlight) { + t.Fatalf("CancelReminder in flight: got %v, want ErrReminderInFlight", err) + } + // nudge lifecycle: record → resolve acted → resolve again ⇒ ErrNudgeOutcome. nid, err := cli.RecordNudge(ctx, "water", "voice", "drink", now) if err != nil { diff --git a/internal/ipc/maperr_test.go b/internal/ipc/maperr_test.go index e08c384..f5fbbb8 100644 --- a/internal/ipc/maperr_test.go +++ b/internal/ipc/maperr_test.go @@ -30,6 +30,7 @@ var mapErrPairs = []struct { {"ErrNudgeOutcome", store.ErrNudgeOutcome, ErrNudgeOutcome}, {"ErrReminderNotFound", store.ErrReminderNotFound, ErrReminderNotFound}, {"ErrReminderState", store.ErrReminderState, ErrReminderState}, + {"ErrReminderInFlight", store.ErrReminderInFlight, ErrReminderInFlight}, {"ErrToolNotFound", store.ErrToolNotFound, ErrToolNotFound}, {"ErrNoSuchTrace", store.ErrNoSuchTrace, ErrNoSuchTrace}, {"ErrTaskNoDoneWhen", store.ErrTaskNoDoneWhen, ErrTaskNoDoneWhen}, diff --git a/internal/ipc/server.go b/internal/ipc/server.go index c07320b..55eba4f 100644 --- a/internal/ipc/server.go +++ b/internal/ipc/server.go @@ -414,9 +414,15 @@ var methodTable = map[Method]handlerFunc{ MethodMarkReminder: withParamsVoid(func(ctx context.Context, api CoreAPI, p markReminderReq) error { return api.MarkReminder(ctx, p.ID, p.Status) }), + MethodCancelReminder: withParamsVoid(func(ctx context.Context, api CoreAPI, p idReq) error { + return api.CancelReminder(ctx, p.ID) + }), MethodListReminders: withParamsSlice(func(ctx context.Context, api CoreAPI, p nReq) ([]Reminder, error) { return api.ListReminders(ctx, p.N) }), + MethodListPendingReminders: withParamsSlice(func(ctx context.Context, api CoreAPI, p nReq) ([]Reminder, error) { + return api.ListPendingReminders(ctx, p.N) + }), MethodRecordNudge: withParams(func(ctx context.Context, api CoreAPI, p recordNudgeReq) (idResp, error) { id, err := api.RecordNudge(ctx, p.Rule, p.Channel, p.Message, p.Ts) return idResp{ID: id}, err diff --git a/internal/ipc/storeapi.go b/internal/ipc/storeapi.go index 5c7aebb..d211f76 100644 --- a/internal/ipc/storeapi.go +++ b/internal/ipc/storeapi.go @@ -77,6 +77,10 @@ func (a *storeAPI) MarkReminder(ctx context.Context, id int64, status string) er return mapErr(a.s.MarkReminder(ctx, id, status)) } +func (a *storeAPI) CancelReminder(ctx context.Context, id int64) error { + return mapErr(a.s.CancelReminder(ctx, id)) +} + // mapRows carries a store read's error through mapErr and converts the rows to // their wire shape. Every list method here is that one shape. func mapRows[S any, W any](rows []S, err error, conv func(S) W) ([]W, error) { @@ -95,6 +99,11 @@ func (a *storeAPI) ListReminders(ctx context.Context, n int) ([]Reminder, error) return mapRows(rs, err, toReminder) } +func (a *storeAPI) ListPendingReminders(ctx context.Context, n int) ([]Reminder, error) { + rs, err := a.s.ListPendingReminders(ctx, n) + return mapRows(rs, err, toReminder) +} + func (a *storeAPI) RescheduleReminder(ctx context.Context, id int64, now time.Time) error { return mapErr(a.s.RescheduleReminder(ctx, id, now)) } @@ -422,6 +431,8 @@ func mapErr(err error) error { return ErrReminderNotFound case errors.Is(err, store.ErrReminderState): return ErrReminderState + case errors.Is(err, store.ErrReminderInFlight): + return ErrReminderInFlight case errors.Is(err, store.ErrToolNotFound): return ErrToolNotFound case errors.Is(err, store.ErrNoSuchTrace): diff --git a/internal/ipc/unimplemented.go b/internal/ipc/unimplemented.go index e5837bc..7da243c 100644 --- a/internal/ipc/unimplemented.go +++ b/internal/ipc/unimplemented.go @@ -47,9 +47,15 @@ func (UnimplementedCoreAPI) CreateReminder(ctx context.Context, fire time.Time, func (UnimplementedCoreAPI) MarkReminder(ctx context.Context, id int64, status string) error { return ErrNotImplemented } +func (UnimplementedCoreAPI) CancelReminder(ctx context.Context, id int64) error { + return ErrNotImplemented +} func (UnimplementedCoreAPI) ListReminders(ctx context.Context, n int) ([]Reminder, error) { return nil, ErrNotImplemented } +func (UnimplementedCoreAPI) ListPendingReminders(ctx context.Context, n int) ([]Reminder, error) { + return nil, ErrNotImplemented +} func (UnimplementedCoreAPI) RecordNudge(ctx context.Context, rule, channel, message string, ts time.Time) (int64, error) { return 0, ErrNotImplemented } diff --git a/internal/ipc/wire.go b/internal/ipc/wire.go index a2058b0..9dd8ec2 100644 --- a/internal/ipc/wire.go +++ b/internal/ipc/wire.go @@ -20,7 +20,9 @@ const ( MethodPresence Method = "presence" MethodCreateReminder Method = "create_reminder" MethodMarkReminder Method = "mark_reminder" + MethodCancelReminder Method = "cancel_reminder" MethodListReminders Method = "list_reminders" + MethodListPendingReminders Method = "list_pending_reminders" MethodRecordNudge Method = "record_nudge" MethodResolveNudge Method = "resolve_nudge" MethodRecentOutcomes Method = "recent_outcomes" @@ -118,19 +120,20 @@ func (e *RpcError) Error() string { // Sentinel codes. Stable over the wire — do not rename. Mirror the package // sentinels in api.go 1:1. The string is the contract. const ( - codeNoFact = "no_fact" - codeConfidence = "confidence" - codeVoidsMissing = "voids_missing" - codeNudgeNotFound = "nudge_not_found" - codeNudgeOutcome = "nudge_outcome" - codeReminderMissing = "reminder_not_found" - codeReminderState = "reminder_state" - codeToolNotFound = "tool_not_found" - codeNoSuchTrace = "no_such_trace" - codeUnknownMethod = "unknown_method" - codeBadParams = "bad_params" - codeForbidden = "forbidden" - codeInternal = "internal" + codeNoFact = "no_fact" + codeConfidence = "confidence" + codeVoidsMissing = "voids_missing" + codeNudgeNotFound = "nudge_not_found" + codeNudgeOutcome = "nudge_outcome" + codeReminderMissing = "reminder_not_found" + codeReminderState = "reminder_state" + codeReminderInFlight = "reminder_in_flight" + codeToolNotFound = "tool_not_found" + codeNoSuchTrace = "no_such_trace" + codeUnknownMethod = "unknown_method" + codeBadParams = "bad_params" + codeForbidden = "forbidden" + codeInternal = "internal" ) // codeOf maps a server-side sentinel to its wire code. Anything not matched is @@ -160,6 +163,8 @@ func codeOf(err error) string { return codeReminderMissing case errors.Is(err, ErrReminderState): return codeReminderState + case errors.Is(err, ErrReminderInFlight): + return codeReminderInFlight case errors.Is(err, ErrToolNotFound): return codeToolNotFound case errors.Is(err, ErrNoSuchTrace): diff --git a/internal/store/delivery.go b/internal/store/delivery.go index 710be9c..080ee70 100644 --- a/internal/store/delivery.go +++ b/internal/store/delivery.go @@ -2,6 +2,7 @@ package store import ( "context" + "database/sql" "fmt" "log" "time" @@ -34,6 +35,9 @@ const ( // stored for post-crash operator triage, not enforced as a uniqueness // constraint (a rule/reminder legitimately re-sends across ticks). func (s *Store) BeginDeliveryAttempt(ctx context.Context, kind, rule string, reminderID int64, deliveryGroup, channel, bodyHash string, now time.Time) (int64, error) { + if kind == "reminder" && deliveryGroup != "" { + return s.beginReminderDeliveryAttempt(ctx, rule, reminderID, deliveryGroup, channel, bodyHash, now) + } res, err := s.db.ExecContext(ctx, `INSERT INTO delivery_attempts (kind, rule, reminder_id, delivery_group, channel, body_hash, status, created_ts) VALUES (?, ?, ?, ?, ?, ?, 'pending', ?)`, @@ -48,6 +52,58 @@ func (s *Store) BeginDeliveryAttempt(ctx context.Context, kind, rule string, rem return id, nil } +// beginReminderDeliveryAttempt serializes the last local cancellation point +// with the first externally ambiguous delivery point. CancelReminder clears a +// group's identity when it wins first; when this insert wins first it leaves a +// pending attempt that makes cancellation refuse. There is therefore no state +// in which both operations report success. +func (s *Store) beginReminderDeliveryAttempt(ctx context.Context, rule string, reminderID int64, deliveryGroup, channel, bodyHash string, now time.Time) (int64, error) { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return 0, fmt.Errorf("begin reminder delivery attempt: begin: %w", err) + } + defer tx.Rollback() + + var status, currentGroup string + err = tx.QueryRowContext(ctx, + `SELECT status, delivery_group FROM reminders WHERE id = ?`, reminderID, + ).Scan(&status, ¤tGroup) + if err != nil { + if err == sql.ErrNoRows { + return 0, ErrReminderNotFound + } + return 0, fmt.Errorf("begin reminder delivery attempt: read reminder %d: %w", reminderID, err) + } + if status != ReminderPending || currentGroup != deliveryGroup { + return 0, fmt.Errorf("%w: reminder %d no longer owns delivery group", ErrReminderState, reminderID) + } + var live int + if err := tx.QueryRowContext(ctx, ` + SELECT COUNT(*) FROM delivery_attempts + WHERE kind = 'reminder' AND delivery_group = ? + AND status IN ('pending', 'sent', 'unknown')`, deliveryGroup).Scan(&live); err != nil { + return 0, fmt.Errorf("begin reminder delivery attempt: inspect group: %w", err) + } + if live > 0 { + return 0, ErrReminderInFlight + } + res, err := tx.ExecContext(ctx, ` + INSERT INTO delivery_attempts (kind, rule, reminder_id, delivery_group, channel, body_hash, status, created_ts) + VALUES ('reminder', ?, ?, ?, ?, ?, 'pending', ?)`, + rule, reminderID, deliveryGroup, channel, bodyHash, now.UnixMilli()) + if err != nil { + return 0, fmt.Errorf("begin reminder delivery attempt: %w", err) + } + id, err := res.LastInsertId() + if err != nil { + return 0, fmt.Errorf("begin reminder delivery attempt: last insert id: %w", err) + } + if err := tx.Commit(); err != nil { + return 0, fmt.Errorf("begin reminder delivery attempt: commit: %w", err) + } + return id, nil +} + // CompleteDeliveryAttempt records the sink's outcome for a prior // BeginDeliveryAttempt. status is "sent", "failed" or "dropped" — never // "pending" or "unknown" (those are set only by Begin and reconciliation diff --git a/internal/store/delivery_test.go b/internal/store/delivery_test.go index 936b911..c8ee928 100644 --- a/internal/store/delivery_test.go +++ b/internal/store/delivery_test.go @@ -54,7 +54,18 @@ func TestListDeliveryAttempts(t *testing.T) { if err := s.CompleteDeliveryAttempt(ctx, dropped, DeliveryDropped, base.Add(time.Minute)); err != nil { t.Fatal(err) } - if _, err := s.BeginDeliveryAttempt(ctx, "reminder", "", 7, "reminder:test", "voice", "h3", base.Add(2*time.Minute)); err != nil { + reminderID, err := s.CreateReminder(ctx, base.Add(time.Hour), `{"text":"test"}`, "") + if err != nil { + t.Fatal(err) + } + reminders, err := s.ListReminders(ctx, 1) + if err != nil || len(reminders) != 1 { + t.Fatalf("list reminder: %+v, %v", reminders, err) + } + if err := s.CacheReminderPhrase(ctx, reminders, "reminder:test", "test", "test", "neutral"); err != nil { + t.Fatal(err) + } + if _, err := s.BeginDeliveryAttempt(ctx, "reminder", "", reminderID, "reminder:test", "voice", "h3", base.Add(2*time.Minute)); err != nil { t.Fatal(err) } @@ -63,7 +74,7 @@ func TestListDeliveryAttempts(t *testing.T) { t.Fatalf("ListDeliveryAttempts = %d rows, err=%v, want 3", len(all), err) } // Newest first. - if all[0].Kind != "reminder" || all[0].ReminderID != 7 { + if all[0].Kind != "reminder" || all[0].ReminderID != reminderID { t.Fatalf("newest row is %+v, want the reminder", all[0]) } if all[0].HasComplete { diff --git a/internal/store/reminders.go b/internal/store/reminders.go index 46705f7..9fcb5a4 100644 --- a/internal/store/reminders.go +++ b/internal/store/reminders.go @@ -87,6 +87,7 @@ const ( var ( ErrReminderNotFound = errors.New("store: reminder not found") ErrReminderState = errors.New("store: reminder not in a mutable state") + ErrReminderInFlight = errors.New("store: reminder delivery already started") ErrReminderPhrase = errors.New("store: reminder delivery phrase invalid") ) @@ -206,15 +207,17 @@ func (s *Store) PendingReminders(ctx context.Context, from, to time.Time) ([]Rem return out, rows.Err() } -// MarkReminder sets a reminder's status. Only valid transitions: pending→fired, -// pending→cancelled. Anything else is a programming error. +// MarkReminder records successful completion of a reminder delivery. The only +// valid transition is pending→fired. Cancellation has stronger outbox and +// collapsed-group invariants and must go through CancelReminder; accepting the +// same status here would leave a legacy bypass around those invariants. func (s *Store) MarkReminder(ctx context.Context, id int64, status string) error { - if status != ReminderFired && status != ReminderCancelled { + if status != ReminderFired { return fmt.Errorf("%w: %s", ErrReminderState, status) } // The old read-then-write transition allowed two callers to both observe // pending and both report success. Keeping the source state in the UPDATE - // predicate makes pending → fired|cancelled one atomic contest (V-678). + // predicate makes pending → fired one atomic contest (V-678). res, err := s.db.ExecContext(ctx, "UPDATE reminders SET status = ? WHERE id = ? AND status = ?", status, id, ReminderPending) @@ -242,6 +245,82 @@ func (s *Store) MarkReminder(ctx context.Context, id int64, status string) error return fmt.Errorf("%w: currently %s", ErrReminderState, current) } +// CancelReminder atomically wins against the beginning of external delivery. +// A pending/sent/unknown outbox row means the presentation may already be +// outside the process, so claiming cancellation would be false. Definite +// failures do not block cancellation. +// +// A collapsed catch-up bundle shares one cached phrase. Cancelling any member +// invalidates that presentation on every still-pending sibling; otherwise the +// next retry could continue saying "three reminders" after one was removed. +func (s *Store) CancelReminder(ctx context.Context, id int64) error { + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("cancel reminder: begin: %w", err) + } + defer tx.Rollback() + + var status, group string + var nextFire int64 + err = tx.QueryRowContext(ctx, + `SELECT status, delivery_group, next_fire_ts FROM reminders WHERE id = ?`, id, + ).Scan(&status, &group, &nextFire) + if errors.Is(err, sql.ErrNoRows) { + return ErrReminderNotFound + } + if err != nil { + return fmt.Errorf("cancel reminder %d: read: %w", id, err) + } + if status != ReminderPending { + return fmt.Errorf("%w: currently %s", ErrReminderState, status) + } + + // New attempts carry an occurrence-scoped delivery group. Migration 25 + // assigned older rows the empty group, so those can only be related to this + // occurrence when they began at or after its next-fire boundary. Without + // that bound, one successful delivery from a recurring reminder's history + // would make the whole series permanently uncancellable after an upgrade. + var live int + err = tx.QueryRowContext(ctx, ` + SELECT COUNT(*) + FROM delivery_attempts + WHERE kind = 'reminder' + AND status IN ('pending', 'sent', 'unknown') + AND ((delivery_group <> '' AND delivery_group = ?) + OR (delivery_group = '' AND reminder_id = ? AND created_ts >= ?))`, + group, id, nextFire).Scan(&live) + if err != nil { + return fmt.Errorf("cancel reminder %d: inspect delivery: %w", id, err) + } + if live > 0 { + return ErrReminderInFlight + } + + if group != "" { + if _, err := tx.ExecContext(ctx, ` + UPDATE reminders + SET delivery_group = '', phrase_body = '', phrase_summary = '', phrase_mood = '' + WHERE delivery_group = ? AND status = ?`, group, ReminderPending); err != nil { + return fmt.Errorf("cancel reminder %d: invalidate group: %w", id, err) + } + } + res, err := tx.ExecContext(ctx, ` + UPDATE reminders + SET status = ?, delivery_group = '', phrase_body = '', phrase_summary = '', phrase_mood = '', + next_attempt_ts = NULL, delivery_blocked_ts = NULL, delivery_blocked_error = '' + WHERE id = ? AND status = ?`, ReminderCancelled, id, ReminderPending) + if err != nil { + return fmt.Errorf("cancel reminder %d: %w", id, err) + } + if err := requireOneReminderRow(res, id); err != nil { + return err + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("cancel reminder %d: commit: %w", id, err) + } + return nil +} + // ListReminders returns the n most recent reminders, newest first. func (s *Store) ListReminders(ctx context.Context, n int) ([]Reminder, error) { rows, err := s.db.QueryContext(ctx, `SELECT `+reminderColumns+` @@ -261,6 +340,35 @@ func (s *Store) ListReminders(ctx context.Context, n int) ([]Reminder, error) { return out, rows.Err() } +// ListPendingReminders returns cancellable reminders in firing order. n <= 0 +// means all pending rows: spoken resolution must not miss an old reminder just +// because newer fired history filled ListReminders' window. +func (s *Store) ListPendingReminders(ctx context.Context, n int) ([]Reminder, error) { + query := `SELECT ` + reminderColumns + ` + FROM reminders + WHERE status = ? + ORDER BY next_fire_ts ASC, id ASC` + args := []any{ReminderPending} + if n > 0 { + query += ` LIMIT ?` + args = append(args, n) + } + rows, err := s.db.QueryContext(ctx, query, args...) + if err != nil { + return nil, fmt.Errorf("list pending reminders: %w", err) + } + defer rows.Close() + var out []Reminder + for rows.Next() { + r, err := scanReminder(rows) + if err != nil { + return nil, err + } + out = append(out, r) + } + return out, rows.Err() +} + // HasDeliveryPhrase reports whether this occurrence already has a durable // presentation. Summary may intentionally be empty (the delivery boundary has // a generic privacy-preserving fallback), so Body is the readiness marker. diff --git a/internal/store/reminders_cancel_test.go b/internal/store/reminders_cancel_test.go new file mode 100644 index 0000000..88a2681 --- /dev/null +++ b/internal/store/reminders_cancel_test.go @@ -0,0 +1,243 @@ +package store + +import ( + "context" + "errors" + "sync" + "testing" + "time" +) + +func createCancellationReminder(t *testing.T, s *Store, fire time.Time, text string) Reminder { + t.Helper() + id, err := s.CreateReminder(context.Background(), fire, `{"text":"`+text+`"}`, "") + if err != nil { + t.Fatal(err) + } + rows, err := s.ListReminders(context.Background(), 1) + if err != nil || len(rows) != 1 || rows[0].ID != id { + t.Fatalf("created reminder %d, list=%+v err=%v", id, rows, err) + } + return rows[0] +} + +func TestListPendingRemindersForCancellation(t *testing.T) { + ctx := context.Background() + s := newTestStore(t) + now := time.Date(2026, 8, 15, 9, 0, 0, 0, time.UTC) + later := createCancellationReminder(t, s, now.Add(2*time.Hour), "later") + earlier := createCancellationReminder(t, s, now.Add(time.Hour), "earlier") + fired := createCancellationReminder(t, s, now.Add(3*time.Hour), "already fired") + if err := s.MarkReminder(ctx, fired.ID, ReminderFired); err != nil { + t.Fatal(err) + } + + all, err := s.ListPendingReminders(ctx, 0) + if err != nil { + t.Fatal(err) + } + if len(all) != 2 || all[0].ID != earlier.ID || all[1].ID != later.ID { + t.Fatalf("pending firing order = %+v", all) + } + one, err := s.ListPendingReminders(ctx, 1) + if err != nil || len(one) != 1 || one[0].ID != earlier.ID { + t.Fatalf("limited pending = %+v, %v", one, err) + } +} + +func TestCancelReminderLifecycle(t *testing.T) { + ctx := context.Background() + s := newTestStore(t) + if err := s.CancelReminder(ctx, 999); !errors.Is(err, ErrReminderNotFound) { + t.Fatalf("missing cancellation = %v", err) + } + r := createCancellationReminder(t, s, time.Now().Add(time.Hour), "doctor") + if err := s.CancelReminder(ctx, r.ID); err != nil { + t.Fatal(err) + } + rows, err := s.ListReminders(ctx, 1) + if err != nil || len(rows) != 1 || rows[0].Status != ReminderCancelled { + t.Fatalf("cancelled row = %+v, %v", rows, err) + } + if err := s.CancelReminder(ctx, r.ID); !errors.Is(err, ErrReminderState) { + t.Fatalf("second cancellation = %v", err) + } +} + +func TestMarkReminderCannotBypassCancellationSafety(t *testing.T) { + ctx := context.Background() + s := newTestStore(t) + r := createCancellationReminder(t, s, time.Now().Add(time.Hour), "doctor") + if err := s.MarkReminder(ctx, r.ID, ReminderCancelled); !errors.Is(err, ErrReminderState) { + t.Fatalf("legacy cancelled transition = %v, want ErrReminderState", err) + } + if got := reminderStatusesForTest(t, s)[r.ID]; got != ReminderPending { + t.Fatalf("legacy transition changed status to %q", got) + } +} + +func reminderStatusesForTest(t *testing.T, s *Store) map[int64]string { + t.Helper() + rows, err := s.ListReminders(context.Background(), 100) + if err != nil { + t.Fatal(err) + } + out := make(map[int64]string, len(rows)) + for _, row := range rows { + out[row.ID] = row.Status + } + return out +} + +func TestCancelReminderInvalidatesCollapsedPhrase(t *testing.T) { + ctx := context.Background() + s := newTestStore(t) + now := time.Date(2026, 8, 15, 9, 0, 0, 0, time.UTC) + a := createCancellationReminder(t, s, now.Add(time.Hour), "doctor") + b := createCancellationReminder(t, s, now.Add(2*time.Hour), "bread") + originals, err := s.ListPendingReminders(ctx, 0) + if err != nil { + t.Fatal(err) + } + const group = "reminder:collapsed-cancel-test" + if err := s.CacheReminderPhrase(ctx, originals, group, "two reminders", "two", "neutral"); err != nil { + t.Fatal(err) + } + attempt, err := s.BeginDeliveryAttempt(ctx, "reminder", "", a.ID, group, "telegram", "hash", now) + if err != nil { + t.Fatal(err) + } + if err := s.CompleteDeliveryAttempt(ctx, attempt, DeliveryFailed, now); err != nil { + t.Fatal(err) + } + if err := s.CancelReminder(ctx, a.ID); err != nil { + t.Fatal(err) + } + + rows, err := s.ListReminders(ctx, 10) + if err != nil { + t.Fatal(err) + } + byID := make(map[int64]Reminder, len(rows)) + for _, row := range rows { + byID[row.ID] = row + } + if got := byID[a.ID]; got.Status != ReminderCancelled || got.HasDeliveryPhrase() { + t.Fatalf("cancelled member retained presentation: %+v", got) + } + if got := byID[b.ID]; got.Status != ReminderPending || got.HasDeliveryPhrase() || got.DeliveryGroup != "" { + t.Fatalf("surviving member retained stale bundle: %+v", got) + } +} + +func TestCancelReminderRefusesAmbiguousDelivery(t *testing.T) { + for _, status := range []string{DeliveryPending, DeliveryUnknown, DeliverySent} { + t.Run(status, func(t *testing.T) { + ctx := context.Background() + s := newTestStore(t) + now := time.Date(2026, 8, 15, 9, 0, 0, 0, time.UTC) + r := createCancellationReminder(t, s, now.Add(time.Hour), status) + group := "reminder:" + status + if err := s.CacheReminderPhrase(ctx, []Reminder{r}, group, status, status, "neutral"); err != nil { + t.Fatal(err) + } + attempt, err := s.BeginDeliveryAttempt(ctx, "reminder", "", r.ID, group, "telegram", "hash", now) + if err != nil { + t.Fatal(err) + } + switch status { + case DeliveryUnknown: + if _, err := s.ReconcileStaleDeliveryAttempts(ctx, now.Add(time.Minute)); err != nil { + t.Fatal(err) + } + case DeliverySent: + if err := s.CompleteDeliveryAttempt(ctx, attempt, DeliverySent, now); err != nil { + t.Fatal(err) + } + } + if err := s.CancelReminder(ctx, r.ID); !errors.Is(err, ErrReminderInFlight) { + t.Fatalf("cancel with %s attempt = %v", status, err) + } + }) + } +} + +func TestCancelReminderScopesLegacyAttemptsToTheCurrentOccurrence(t *testing.T) { + ctx := context.Background() + now := time.Date(2026, 8, 15, 9, 0, 0, 0, time.UTC) + + t.Run("historical recurring send does not block the series", func(t *testing.T) { + s := newTestStore(t) + fire := now.Add(24 * time.Hour) + id, err := s.CreateReminder(ctx, fire, `{"text":"daily medicine"}`, "0 9 * * *") + if err != nil { + t.Fatal(err) + } + if _, err := s.db.ExecContext(ctx, ` + INSERT INTO delivery_attempts + (kind, reminder_id, delivery_group, channel, body_hash, status, created_ts, completed_ts) + VALUES ('reminder', ?, '', 'telegram', 'legacy', 'sent', ?, ?)`, + id, now.UnixMilli(), now.UnixMilli()); err != nil { + t.Fatal(err) + } + if err := s.CancelReminder(ctx, id); err != nil { + t.Fatalf("historical blank-group attempt blocked recurrence: %v", err) + } + }) + + t.Run("current legacy ambiguity still refuses", func(t *testing.T) { + s := newTestStore(t) + fire := now.Add(time.Hour) + r := createCancellationReminder(t, s, fire, "legacy current") + if _, err := s.db.ExecContext(ctx, ` + INSERT INTO delivery_attempts + (kind, reminder_id, delivery_group, channel, body_hash, status, created_ts) + VALUES ('reminder', ?, '', 'telegram', 'legacy', 'unknown', ?)`, + r.ID, fire.Add(time.Second).UnixMilli()); err != nil { + t.Fatal(err) + } + if err := s.CancelReminder(ctx, r.ID); !errors.Is(err, ErrReminderInFlight) { + t.Fatalf("current blank-group ambiguity = %v, want ErrReminderInFlight", err) + } + }) +} + +func TestReminderCancellationAndDeliveryBeginAreMutuallyExclusive(t *testing.T) { + for i := 0; i < 20; i++ { + s := newTestStore(t) + ctx := context.Background() + now := time.Date(2026, 8, 15, 9, 0, 0, i, time.UTC) + r := createCancellationReminder(t, s, now.Add(time.Hour), "race") + group := "reminder:race" + if err := s.CacheReminderPhrase(ctx, []Reminder{r}, group, "race", "race", "neutral"); err != nil { + t.Fatal(err) + } + + start := make(chan struct{}) + var cancelErr, beginErr error + var wg sync.WaitGroup + wg.Add(2) + go func() { + defer wg.Done() + <-start + cancelErr = s.CancelReminder(ctx, r.ID) + }() + go func() { + defer wg.Done() + <-start + _, beginErr = s.BeginDeliveryAttempt(ctx, "reminder", "", r.ID, group, "telegram", "hash", now) + }() + close(start) + wg.Wait() + + if (cancelErr == nil) == (beginErr == nil) { + t.Fatalf("iteration %d: cancel=%v begin=%v; exactly one must win", i, cancelErr, beginErr) + } + if cancelErr == nil && !errors.Is(beginErr, ErrReminderState) { + t.Fatalf("iteration %d: cancellation won, begin=%v", i, beginErr) + } + if beginErr == nil && !errors.Is(cancelErr, ErrReminderInFlight) { + t.Fatalf("iteration %d: delivery won, cancel=%v", i, cancelErr) + } + } +} diff --git a/internal/store/reminders_delivery_test.go b/internal/store/reminders_delivery_test.go index 02fa400..627ea2f 100644 --- a/internal/store/reminders_delivery_test.go +++ b/internal/store/reminders_delivery_test.go @@ -158,15 +158,17 @@ func TestReminderTerminalTransitionHasExactlyOneWinner(t *testing.T) { start := make(chan struct{}) var wg sync.WaitGroup errs := make([]error, 2) - statuses := []string{ReminderFired, ReminderCancelled} - for i := range statuses { - wg.Add(1) - go func(i int) { - defer wg.Done() - <-start - errs[i] = s.MarkReminder(ctx, id, statuses[i]) - }(i) - } + wg.Add(2) + go func() { + defer wg.Done() + <-start + errs[0] = s.MarkReminder(ctx, id, ReminderFired) + }() + go func() { + defer wg.Done() + <-start + errs[1] = s.CancelReminder(ctx, id) + }() close(start) wg.Wait() @@ -177,11 +179,11 @@ func TestReminderTerminalTransitionHasExactlyOneWinner(t *testing.T) { switch { case err == nil: successes++ - winner = statuses[i] + winner = []string{ReminderFired, ReminderCancelled}[i] case errors.Is(err, ErrReminderState): losers++ default: - t.Fatalf("iteration %d transition %s: %v", iteration, statuses[i], err) + t.Fatalf("iteration %d transition %d: %v", iteration, i, err) } } if successes != 1 || losers != 1 {