package smtp_server import ( "bytes" "crypto/tls" "fmt" "io" "log" "net" "strconv" "strings" "time" "mail_go/config" "mail_go/internal/connhub" "mail_go/internal/db" "mail_go/internal/imap_server" "mail_go/internal/mailutil" "mail_go/internal/outbound" "mail_go/internal/storage" "mail_go/internal/store" "mail_go/internal/tlsutil" "github.com/emersion/go-message/mail" "github.com/emersion/go-sasl" "github.com/emersion/go-smtp" ) type smtpMode int const ( smtpModeInbound smtpMode = iota smtpModeSubmission smtpModeImplicitTLS ) // SMTPServer wraps go-smtp servers and provides local mail delivery. type SMTPServer struct { stores *store.Stores storage *storage.AttachmentStorage outbound *outbound.Manager cfg config.SMTPConfig banCfg config.BanConfig tlsLoader *tlsutil.Loader hub *connhub.Hub pusher imap_server.Pusher // 本地投递成功推送(IMAP 新邮件),可空 } // NewSMTPServer creates a new SMTP server instance. tlsLoader may be nil // when TLS is not configured. func NewSMTPServer(cfg config.SMTPConfig, stores *store.Stores, attStorage *storage.AttachmentStorage, ob *outbound.Manager, tlsLoader *tlsutil.Loader, banCfg config.BanConfig, hub *connhub.Hub, pusher imap_server.Pusher) *SMTPServer { return &SMTPServer{stores: stores, storage: attStorage, outbound: ob, cfg: cfg, banCfg: banCfg, tlsLoader: tlsLoader, hub: hub, pusher: pusher} } func (s *SMTPServer) tlsConfig() (*tls.Config, error) { if s.tlsLoader == nil { return nil, fmt.Errorf("SMTP TLS certificate or key not configured") } // GetCertificate 每次握手按需重载证书,证书更新后无需重启服务 return &tls.Config{GetCertificate: s.tlsLoader.GetCertificate}, nil } func (s *SMTPServer) newServer(addr string, mode smtpMode, tlsConfig *tls.Config) *smtp.Server { be := &smtpBackend{server: s, mode: mode} srv := smtp.NewServer(be) srv.Addr = addr srv.Domain = s.cfg.Domain srv.MaxMessageBytes = s.cfg.MaxMessage srv.AllowInsecureAuth = tlsConfig == nil srv.ReadTimeout = 60 * time.Second srv.WriteTimeout = 60 * time.Second srv.TLSConfig = tlsConfig return srv } // Start starts the inbound SMTP server. func (s *SMTPServer) Start() error { tlsConfig, err := s.tlsConfig() if err != nil { log.Printf("SMTP STARTTLS 未启用: %v", err) } log.Printf("SMTP server listening on %s", s.cfg.Addr) return s.newServer(s.cfg.Addr, smtpModeInbound, tlsConfig).ListenAndServe() } // StartTLS starts the implicit TLS SMTP submission server. func (s *SMTPServer) StartTLS() error { tlsConfig, err := s.tlsConfig() if err != nil { return err } log.Printf("SMTPS server listening on %s", s.cfg.TLSAddr) return s.newServer(s.cfg.TLSAddr, smtpModeImplicitTLS, tlsConfig).ListenAndServeTLS() } // StartSubmission starts the SMTP submission server with STARTTLS support. func (s *SMTPServer) StartSubmission() error { tlsConfig, err := s.tlsConfig() if err != nil { return err } log.Printf("SMTP submission server listening on %s", s.cfg.SubmissionAddr) return s.newServer(s.cfg.SubmissionAddr, smtpModeSubmission, tlsConfig).ListenAndServe() } // smtpBackend implements the smtp.Backend interface. type smtpBackend struct { server *SMTPServer mode smtpMode } // NewSession creates a new SMTP session for the incoming connection. func (be *smtpBackend) NewSession(c *smtp.Conn) (smtp.Session, error) { clientIP := store.ClientIPFromAddr(c.Conn().RemoteAddr()) conn := be.server.hub.Register("smtp", clientIP, be.server.sessionPort(be.mode), be.server.tlsActive(c)) if conn != nil { // 强制断开:关闭底层连接后 go-smtp 读到 EOF,正常走 Logout 收尾 raw := c.Conn() conn.SetDisconnect(func() { _ = raw.Close() }) } return &smtpSession{ backend: be, mode: be.mode, rcpts: make([]string, 0), clientIP: clientIP, startedAt: time.Now(), port: be.server.sessionPort(be.mode), conn: conn, }, nil } // tlsActive 判断当前连接是否处于 TLS 加密状态(implicit TLS 或 STARTTLS)。 func (s *SMTPServer) tlsActive(c *smtp.Conn) bool { _, ok := c.TLSConnectionState() return ok } // sessionPort 返回该会话监听的端口号(区分明文/TLS/提交端口),解析失败返回 0。 func (s *SMTPServer) sessionPort(mode smtpMode) int { addr := s.cfg.Addr switch mode { case smtpModeSubmission: addr = s.cfg.SubmissionAddr case smtpModeImplicitTLS: addr = s.cfg.TLSAddr } _, portStr, err := net.SplitHostPort(addr) if err != nil { return 0 } port, err := strconv.Atoi(portStr) if err != nil { return 0 } return port } // smtpSession implements the smtp.Session interface for handling a single connection. type smtpSession struct { backend *smtpBackend mode smtpMode from string rcpts []string localRcpts []string externalRcpts []string authenticated bool userID uint email string user *db.User clientIP string // 会话日志累积状态 startedAt time.Time port int authTried bool authOK bool authUsername string failReason string // 首个失败原因 msgCount int // 成功处理的邮件数(本地投递 + 外发队列) detailParts []string // 连接追踪 conn *connhub.Conn } // AuthMechanisms returns supported SMTP AUTH mechanisms. func (s *smtpSession) AuthMechanisms() []string { return []string{sasl.Plain} } // recordFail 记录会话中第一个失败原因(日志用途,不改变协议行为)。 func (s *smtpSession) recordFail(reason string) { if s.failReason == "" { s.failReason = reason } } // recordDetail 追加一条操作摘要。 func (s *smtpSession) recordDetail(part string) { s.detailParts = append(s.detailParts, part) } // Auth authenticates the user with SASL PLAIN credentials. func (s *smtpSession) Auth(mech string) (sasl.Server, error) { if mech != sasl.Plain { s.recordFail("不支持的认证机制") return nil, smtp.ErrAuthUnknownMechanism } return sasl.NewPlainServer(func(identity, username, password string) error { s.authTried = true s.authUsername = username s.conn.Touch() // 已封禁 IP 一律拒绝认证(防协议层暴力破解) if banned, _ := s.backend.server.stores.Bans.IsBanned(s.clientIP); banned { s.recordFail("IP已被封禁") return smtp.ErrAuthFailed } user, err := s.backend.server.stores.Users.AuthenticateLogin(username, password) if err != nil { // 认证失败计数,达到阈值按档位封禁(与 Web 登录共用 ban_entries)。 // 用户名不存在(枚举型爆破)跳过宽限首次触发即封。 s.backend.server.stores.RecordAuthFailure( s.clientIP, s.backend.server.banCfg.MaxFailAttempts, s.backend.server.banCfg.BanDurationMin, "邮件协议认证失败次数过多", s.backend.server.stores.Users.LoginExists(username), ) s.recordFail("用户名或密码错误") return smtp.ErrAuthFailed } // 登录成功清零失败计数(与 Web 登录一致):防止合法用户 IP // 因失败计数只增不减被反复误封。 s.backend.server.stores.Bans.ResetFail(s.clientIP) domainName := user.Domain.Name if domainName == "" { domain, err := s.backend.server.stores.Domains.GetByID(user.DomainID) if err == nil { domainName = domain.Name } } if domainName == "" { s.recordFail("用户名或密码错误") return smtp.ErrAuthFailed } s.authenticated = true s.authOK = true s.userID = user.ID s.user = user s.email = user.Username + "@" + domainName if s.conn != nil { s.conn.SetUser(s.email) } return nil }), nil } // Mail records the sender address (MAIL FROM command). func (s *smtpSession) Mail(from string, opts *smtp.MailOptions) error { if s.mode != smtpModeInbound && !s.authenticated { s.recordFail("未认证用户尝试发信") return smtp.ErrAuthRequired } // Authenticated users may only send as themselves, preventing spoofing. if s.authenticated && !strings.EqualFold(strings.TrimSpace(from), s.email) { s.recordFail("发件人地址与登录用户不一致") return fmt.Errorf("sender address must match authenticated user") } s.from = from s.rcpts = s.rcpts[:0] s.localRcpts = s.localRcpts[:0] s.externalRcpts = s.externalRcpts[:0] return nil } // Rcpt validates and records a recipient address (RCPT TO command). // Local recipients are delivered to the mailbox; external recipients are // allowed only for authenticated users and go to the outbound queue, // which prevents open relay. func (s *smtpSession) Rcpt(to string, opts *smtp.RcptOptions) error { to = strings.TrimSpace(to) if to == "" { s.recordFail("无效的收件人地址") return fmt.Errorf("invalid recipient address: %s", to) } if _, err := s.localUserByEmail(to); err == nil { s.rcpts = append(s.rcpts, to) s.localRcpts = append(s.localRcpts, to) return nil } // External recipient: only authenticated local users may relay. if !s.authenticated { s.recordFail("中继访问被拒绝") return fmt.Errorf("relay access denied: %s", to) } // Sender verification must have been enforced in Mail() already. ob := s.backend.server.outbound if ob == nil || !ob.Enabled() { s.recordFail("外部投递未启用") return fmt.Errorf("external delivery is disabled: %s", to) } s.rcpts = append(s.rcpts, to) s.externalRcpts = append(s.externalRcpts, to) return nil } func (s *smtpSession) localUserByEmail(email string) (*db.User, error) { return s.backend.server.stores.Users.GetByEmail(strings.TrimSpace(email)) } // Data handles the message body and stores it for local recipients. // External recipients (authenticated sessions only) are queued for // outbound delivery. func (s *smtpSession) Data(r io.Reader) error { s.conn.Touch() if len(s.rcpts) == 0 { s.recordFail("未指定收件人") return fmt.Errorf("no accepted recipients") } data, err := io.ReadAll(r) if err != nil { s.recordFail("读取邮件数据失败") return fmt.Errorf("failed to read message data: %w", err) } parsed, err := parseSMTPMessage(data) if err != nil { s.recordFail("邮件格式解析失败") return err } // 附件配额:入站路径此前只记账不检查,外部匿名发件人可无限落盘直至 // 磁盘写满。口径与 Web 上传一致(只计附件)。任一本地收件人配额不足 // 则整封以 452 临时失败拒收(发件方按退避重试),避免部分投递; // 附件总量在 DATA 阶段才可知,无法在 RCPT 阶段按收件人拒绝。 attTotal := attachmentSize(parsed) rcptUser, charged, qerr := s.reserveLocalQuota(attTotal) if qerr != nil { return qerr } // Local recipients: deliver to INBOX. localDelivered := 0 storedByUser := make(map[uint]int64, len(charged)) for _, rcpt := range s.localRcpts { if _, reserved := rcptUser[rcpt]; attTotal > 0 && !reserved { // 预扣阶段未覆盖的收件人(当时不存在):跳过,避免绕过配额 log.Printf("SMTP: recipient %s missing at quota reservation, skipping", rcpt) continue } user, err := s.localUserByEmail(rcpt) if err != nil { log.Printf("SMTP: recipient not found %s, skipping", rcpt) continue } msg, stored, err := s.saveMessage(user.ID, "INBOX", parsed, data, false) if err != nil { log.Printf("SMTP: failed to create message for %s: %v", rcpt, err) continue } storedByUser[user.ID] += stored log.Printf("SMTP: message delivered to %s", rcpt) localDelivered++ // 本地投递成功 → IMAP 新邮件推送(IDLE 客户端实时收到通知) if pusher := s.backend.server.pusher; pusher != nil && msg != nil { pusher.PushNewMessage(user.Username+"@"+user.Domain.Name, msg) } } // 配额对账:退还预扣量中未真正落库的部分(收件人被删、落库失败、 // 部分附件保存失败),保证配额与实际占用一致。 s.reconcileQuota(charged, storedByUser) s.msgCount += localDelivered // External recipients: queue for outbound delivery. externalQueued := 0 if len(s.externalRcpts) > 0 { ob := s.backend.server.outbound if ob == nil { s.recordFail("外部投递服务不可用") return fmt.Errorf("outbound delivery is unavailable") } maxRcpt := ob.MaxRecipients() if maxRcpt > 0 && len(s.externalRcpts) > maxRcpt { s.recordFail("外部收件人数量超出限制") return fmt.Errorf("too many external recipients: %d (max %d)", len(s.externalRcpts), maxRcpt) } for _, rcpt := range s.externalRcpts { if _, err := ob.Enqueue(s.user, s.email, rcpt, data); err != nil { s.recordFail("外发队列投递失败") return fmt.Errorf("failed to queue external recipient %s: %v", rcpt, err) } log.Printf("SMTP: external message queued for %s", rcpt) externalQueued++ } } s.msgCount += externalQueued if s.authenticated && s.userID != 0 && s.mode != smtpModeInbound { // Sent 副本尽力而为:本地/外发投递已完成,不因配额回滚整封; // 附件量先原子预扣,保存后按实际落库量对账退还。 if attTotal > 0 { if ok, err := s.backend.server.stores.Users.TryReserveQuota(s.userID, attTotal); err != nil { log.Printf("SMTP: quota check for sent copy of %s failed: %v", s.email, err) } else if !ok { log.Printf("SMTP: sender %s over quota, skipping sent copy", s.email) } else if _, stored, err := s.saveMessage(s.userID, "Sent", parsed, data, true); err != nil { _ = s.backend.server.stores.Users.UpdateUsedBytes(s.userID, -attTotal) log.Printf("SMTP: failed to save sent copy for %s: %v", s.email, err) } else if refund := attTotal - stored; refund > 0 { _ = s.backend.server.stores.Users.UpdateUsedBytes(s.userID, -refund) } } else if _, _, err := s.saveMessage(s.userID, "Sent", parsed, data, true); err != nil { log.Printf("SMTP: failed to save sent copy for %s: %v", s.email, err) } } s.recordDetail(fmt.Sprintf("MAIL FROM:<%s> RCPT×%d 本地投递%d 外发%d", s.from, len(s.rcpts), localDelivered, externalQueued)) return nil } type parsedSMTPMessage struct { messageID string fromAddr string toAddr string ccAddr string subject string textBody string htmlBody string date time.Time attachments []*parsedAttachment } // parsedAttachment holds an extracted MIME attachment part. type parsedAttachment struct { fileName string contentType string data []byte } func parseSMTPMessage(data []byte) (*parsedSMTPMessage, error) { mr, err := mail.CreateReader(bytes.NewReader(data)) if err != nil { return nil, fmt.Errorf("failed to parse MIME message: %w", err) } header := mr.Header msg := &parsedSMTPMessage{} msg.fromAddr = mailutil.FormatAddressList(&header, "From") msg.toAddr = mailutil.FormatAddressList(&header, "To") msg.ccAddr = mailutil.FormatAddressList(&header, "Cc") msg.subject, _ = header.Subject() msg.messageID, _ = header.MessageID() msg.date, _ = header.Date() if msg.date.IsZero() { msg.date = time.Now() } for { p, err := mr.NextPart() if err == io.EOF { break } if err != nil { log.Printf("SMTP: error reading MIME part: %v", err) break } switch h := p.Header.(type) { case *mail.InlineHeader: contentType, params, _ := h.ContentType() buf, readErr := io.ReadAll(p.Body) if readErr != nil { log.Printf("SMTP: error reading inline part: %v", readErr) continue } charset := "" if cs, ok := params["charset"]; ok { charset = cs } decoded := mailutil.DecodeCharset(buf, charset) if strings.HasPrefix(contentType, "text/plain") { msg.textBody = decoded } else if strings.HasPrefix(contentType, "text/html") { msg.htmlBody = decoded } case *mail.AttachmentHeader: filename, _ := h.Filename() if filename == "" { filename = "unnamed_attachment" } contentType, _, _ := h.ContentType() buf, readErr := io.ReadAll(p.Body) if readErr != nil { log.Printf("SMTP: error reading attachment part: %v", readErr) continue } msg.attachments = append(msg.attachments, &parsedAttachment{ fileName: filename, contentType: contentType, data: buf, }) } } if msg.textBody == "" && msg.htmlBody == "" { msg.textBody = string(data) } return msg, nil } // attachmentSize 汇总邮件附件总字节数(配额记账口径与 Web 上传一致)。 func attachmentSize(parsed *parsedSMTPMessage) int64 { total := int64(0) for _, att := range parsed.attachments { total += int64(len(att.data)) } return total } // reserveLocalQuota 为本地收件人预扣附件配额(TryReserveQuota 原子预扣, // 防止并发投递绕过配额)。任一收件人配额不足则回退全部已预扣量,并返回 // 452 临时失败(普通 error 会被 go-smtp 映射为 554 永久失败,语义错误)。 // 同一收件人重复 RCPT TO 时投递多份,预扣量按用户聚合。 // 返回 rcpt→userID 映射与各用户预扣量,供投递与对账使用。 func (s *smtpSession) reserveLocalQuota(attTotal int64) (map[string]uint, map[uint]int64, *smtp.SMTPError) { rcptUser := make(map[string]uint, len(s.localRcpts)) if attTotal <= 0 { return rcptUser, nil, nil } quotaNeed := make(map[uint]int64, len(s.localRcpts)) for _, rcpt := range s.localRcpts { if _, ok := rcptUser[rcpt]; ok { continue } user, err := s.localUserByEmail(rcpt) if err != nil { continue // 投递阶段同样会跳过 } rcptUser[rcpt] = user.ID quotaNeed[user.ID] += attTotal } charged := make(map[uint]int64, len(quotaNeed)) for uid, amount := range quotaNeed { ok, err := s.backend.server.stores.Users.TryReserveQuota(uid, amount) if err != nil { s.refundQuota(charged) s.recordFail("配额检查失败") return nil, nil, &smtp.SMTPError{ Code: 452, EnhancedCode: smtp.EnhancedCode{4, 3, 0}, Message: fmt.Sprintf("Temporary quota check failure: %v", err), } } if !ok { s.refundQuota(charged) log.Printf("SMTP: recipient over quota (needs %d bytes), rejecting message", amount) s.recordFail("收件人邮箱配额不足") return nil, nil, &smtp.SMTPError{ Code: 452, EnhancedCode: smtp.EnhancedCode{4, 2, 2}, Message: "Insufficient system storage", } } charged[uid] = amount } return rcptUser, charged, nil } // refundQuota 回退预扣阶段已扣减的配额(整封拒收时保持原子性)。 func (s *smtpSession) refundQuota(charged map[uint]int64) { for uid, amount := range charged { _ = s.backend.server.stores.Users.UpdateUsedBytes(uid, -amount) } } // reconcileQuota 对账:退还预扣量中未真正落库的部分(收件人被删、落库 // 失败、部分附件保存失败),使配额与实际占用一致。 func (s *smtpSession) reconcileQuota(charged, storedByUser map[uint]int64) { for uid, amount := range charged { if refund := amount - storedByUser[uid]; refund > 0 { _ = s.backend.server.stores.Users.UpdateUsedBytes(uid, -refund) } } } // saveMessage 存储一封邮件及其附件。配额由调用方预扣并对账:本函数返回 // 实际落库的附件字节数(stored),调用方据此退还未落库的预扣量; // Mails.Create 失败时附件尚未处理,调用方应退还全部预扣量。 func (s *smtpSession) saveMessage(userID uint, folder string, parsed *parsedSMTPMessage, data []byte, read bool) (*db.Message, int64, error) { msg := &db.Message{ UserID: userID, MessageID: parsed.messageID, Folder: folder, FromAddr: parsed.fromAddr, ToAddr: parsed.toAddr, CcAddr: parsed.ccAddr, Subject: parsed.subject, TextBody: parsed.textBody, HtmlBody: parsed.htmlBody, RawData: string(data), IsRead: read, IsFlagged: false, Date: parsed.date, } if err := s.backend.server.stores.Mails.Create(msg); err != nil { return nil, 0, err } // Persist attachments to disk and link them to the message so that the // Web mail UI can list/download them. Quota is reserved by the caller // (TryReserveQuota) and reconciled against the returned stored size. var stored int64 for _, att := range parsed.attachments { relPath, err := s.backend.server.storage.Save(att.fileName, att.data) if err != nil { log.Printf("SMTP: failed to save attachment %s: %v", att.fileName, err) continue } rec := &db.Attachment{ MessageID: msg.ID, FileName: att.fileName, FilePath: relPath, ContentType: att.contentType, FileSize: int64(len(att.data)), } if err := s.backend.server.stores.Attachments.Create(rec); err != nil { log.Printf("SMTP: failed to create attachment record: %v", err) continue } stored += int64(len(att.data)) } return msg, stored, nil } // Reset clears the session state for the next message on the same connection. func (s *smtpSession) Reset() { s.from = "" s.rcpts = s.rcpts[:0] s.localRcpts = s.localRcpts[:0] s.externalRcpts = s.externalRcpts[:0] } // Logout is called when the SMTP connection is closed. func (s *smtpSession) Logout() error { s.writeProtocolLog() if s.conn != nil { s.conn.Close() } return nil } // writeProtocolLog 汇总本会话状态写入协议调用日志(供后台分析攻击/滥用)。 func (s *smtpSession) writeProtocolLog() { success := s.failReason == "" detail := strings.Join(s.detailParts, "; ") username := s.authUsername if username == "" && s.email != "" { username = s.email } if detail == "" { if s.authTried { if s.authOK { detail = "AUTH 成功" } else { detail = "AUTH 失败" } } else if success { detail = "连接建立,无邮件操作" } } if success && s.authTried && !s.authOK { success = false } entry := &db.ProtocolLog{ Protocol: db.ProtocolSMTP, Port: s.port, ClientIP: s.clientIP, Username: username, Success: success, FailReason: s.failReason, Detail: detail, MsgCount: s.msgCount, DurationMs: time.Since(s.startedAt).Milliseconds(), CreatedAt: time.Now(), } if err := s.backend.server.stores.ProtocolLogs.Create(entry); err != nil { log.Printf("SMTP: 写入协议日志失败: %v", err) } }