diff --git a/server/internal/handler/mail.go b/server/internal/handler/mail.go index ef1a1d7..a8bb116 100644 --- a/server/internal/handler/mail.go +++ b/server/internal/handler/mail.go @@ -481,10 +481,6 @@ func SendMail(w http.ResponseWriter, r *http.Request) { Error(w, http.StatusInternalServerError, "Failed to create mail") return } - if relay != "" { - // 关联失败不影响功能,只是少一条审计记录 - _ = repo.BindRelayMail(r.Context(), agentName, relayKey, mailID) - } if proposal != nil { // 记不上提议不该让发信失败:邮件本身已经入库,提议是旁支信息 _ = repo.SetMailRenameProposal(r.Context(), mailID, proposal.Alias, proposal.Reason) @@ -497,6 +493,13 @@ func SendMail(w http.ResponseWriter, r *http.Request) { // // 三件事都要退:邮件本身、本次往返预算、relay 幂等键。 // 错误均忽略:响应已由 attachAll 写出,回滚失败只能记日志。 + // + // 次序要紧:**必须在这里还没 BindRelayMail 时才可能真正回滚**。 + // relayed_mails.mail_id 对 mails 是外键且**无 CASCADE**,一旦绑定, + // 上面那句 DELETE 会直接撞外键失败,邮件留在库里;而 ReleaseRelay 的 + // WHERE 是 `mail_id IS NULL`,对已绑定的行也是 no-op —— 于是"三件事 + // 都要退"一件都退不掉,发件方收到 4xx 重试,收件方还会看到那封残余邮件。 + // 所以 BindRelayMail 放在这个回滚块**之后**(见下)。 _ = repo.DeleteMailByID(r.Context(), mailID) if !relayFree { repo.RefundSessionBudget(r.Context(), sessionID) @@ -507,6 +510,15 @@ func SendMail(w http.ResponseWriter, r *http.Request) { return } + if relay != "" { + // 关联失败不影响功能,只是少一条审计记录。 + // + // 放在 attachAll **之后**:上面那条回滚路径要能真的删掉邮件,就必须 + // 在删的时候还没有任何行引用它。绑定是纯审计关联,notify 的载荷里 + // 没有 relay 字段(models.Mail 零 relay 字段),所以推迟到这里没有副作用。 + _ = repo.BindRelayMail(r.Context(), agentName, relayKey, mailID) + } + notifyRecipients(r.Context(), to, ccList, sessionID, mailID, agentName, req.Subject, parentIDString(parentMailID)) // 回传会话别名与本任务剩余往返,让发件方知道后续用什么地址续谈、还能发几封 diff --git a/server/internal/handler/mail_rollback_attach_test.go b/server/internal/handler/mail_rollback_attach_test.go new file mode 100644 index 0000000..9de65a8 --- /dev/null +++ b/server/internal/handler/mail_rollback_attach_test.go @@ -0,0 +1,79 @@ +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) + } +}