package handler import ( "bytes" "context" "encoding/json" "net/http" "net/http/httptest" "testing" "github.com/agentmail/gateway/internal/db" "github.com/agentmail/gateway/internal/middleware" "github.com/agentmail/gateway/internal/repo" ) // 附件竞态回滚必须**真的**把邮件和幂等键都退掉。 // // 触发方式:同一个 attachment_id 传两次 —— `checkAttachable` 逐个查时两条都还 // 没挂载(都通过),`attachAll` 的原子 UPDATE 在第二条上改到 0 行 ⇒ 409。 // 这正是那条回滚路径要处理的形状("另一个请求把同一个附件挂走了"), // 只是这里由同一个请求自己造出竞态,无需真并发。 // // 背景:relayed_mails.mail_id 对 mails 是外键且**无 CASCADE**。原先 // BindRelayMail 在 attachAll **之前**执行,于是回滚块里的 DeleteMailByID 撞外键 // 失败(邮件留在库里),而 ReleaseRelay 的 WHERE 是 `mail_id IS NULL`,对已绑定的 // 行也是 no-op —— 注释写着"三件事都要退",实际一件都没退掉。 // // 判据两条都要: // 1. 邮件没留下(否则会重复投递给收件方); // 2. 幂等键退回去了(否则上游那条消息永远转不出来,重试只会 duplicate)。 // // 只断言 1 会放过"邮件删了但键被永久占住";只断言 2 会放过残余邮件。 func TestMailSendRollbackRemovesMailAndRelayKey(t *testing.T) { setupPermissionHandlerDB(t) ctx := context.Background() if err := repo.CreateOrUpdateAgent(ctx, "pi", "secret", "test", nil); err != nil { t.Fatal(err) } sid, err := repo.CreateSession(ctx, nil, "pi", "x", "/tmp") if err != nil { t.Fatal(err) } att, err := repo.CreateAttachment(ctx, "pi", "f.txt", "text/plain", 3, "sum-x") if err != nil { t.Fatal(err) } payload, _ := json.Marshal(map[string]any{ "to": "pi", "subject": "s", "body": "b", "from_session_id": sid.String(), "attachment_ids": []string{att.ID.String(), att.ID.String()}, "relay_key": "probe:rb-1", "relay": "summary", }) req := httptest.NewRequest(http.MethodPost, "/api/v1/mail/send", bytes.NewReader(payload)) req.Header.Set("Content-Type", "application/json") reqCtx := context.WithValue(req.Context(), middleware.AgentNameKey, "pi") resp := httptest.NewRecorder() SendMail(resp, req.WithContext(reqCtx)) // 前提:触发条件成立,确实走到了回滚那条路。 // // 这一条不能省:若哪天 parseAttachmentIDs 开始去重(或 attachAll 变得容忍 // 重复 id),请求会变成 200 且**根本没建邮件**,于是下面两个计数都是 0 —— // 用例会"通过",但它一次都没验到回滚。断言触发本身,才不会静默退化成空跑。 if resp.Code != http.StatusConflict { t.Fatalf("重复 attachment_id 应触发 409 以走到回滚路径,实际 HTTP %d: %s", resp.Code, resp.Body.String()) } var mails, relays int _ = db.DB.QueryRowContext(ctx, `SELECT count(*) FROM mails`).Scan(&mails) _ = db.DB.QueryRowContext(ctx, `SELECT count(*) FROM relayed_mails WHERE relay_key='probe:rb-1'`).Scan(&relays) if mails != 0 { t.Errorf("回滚没删掉邮件(mails=%d)", mails) } if relays != 0 { t.Errorf("回滚没退掉幂等键(relays=%d)", relays) } }