internal/inbound/inbound.go (view raw)
1package inbound
2
3import (
4 "bytes"
5 "context"
6 "crypto/tls"
7 "database/sql"
8 "errors"
9 "fmt"
10 "io"
11 "log/slog"
12 "net"
13 "net/mail"
14 "os"
15 "strings"
16 "sync"
17 "time"
18
19 "postern/internal/db"
20 "postern/internal/model"
21 "postern/internal/policy"
22
23 "github.com/emersion/go-imap/v2/imapserver"
24 "github.com/emersion/go-sasl"
25 smtpserver "github.com/emersion/go-smtp"
26 "github.com/google/uuid"
27)
28
29type MessageDeliverer interface {
30 SubmitMessage(ctx context.Context, conn net.Conn, user model.User, r io.Reader) error
31}
32
33type Config struct {
34 DB *db.DB
35 Deliverer MessageDeliverer
36 Listen string
37 TLS *tls.Config
38 PolicyEngine *policy.Engine
39}
40type Server struct {
41 db *db.DB
42 listen string
43 deliverer MessageDeliverer
44 tlsConfig *tls.Config
45 policyEngine *policy.Engine
46
47 mu sync.Mutex
48 smtpSrv *smtpserver.Server
49 status string
50}
51
52func NewServer(config *Config) (*Server, error) {
53 return &Server{
54 db: config.DB,
55 listen: config.Listen,
56 deliverer: config.Deliverer,
57 tlsConfig: config.TLS,
58 policyEngine: config.PolicyEngine,
59 }, nil
60}
61
62func (i *Server) Start() error {
63 be := &backend{
64 db: i.db,
65 deliverer: i.deliverer,
66 policyEngine: i.policyEngine,
67 }
68 s := smtpserver.NewServer(be)
69
70 s.Addr = i.listen
71 s.AllowInsecureAuth = false
72 s.TLSConfig = i.tlsConfig
73 s.MaxMessageBytes = 20_000_000
74
75 i.mu.Lock()
76 i.smtpSrv = s
77 i.status = "running"
78 i.mu.Unlock()
79
80 err := s.ListenAndServe()
81 if err != nil {
82 return err
83 }
84
85 return nil
86}
87
88func (i *Server) Stop() error {
89 i.mu.Lock()
90 defer i.mu.Unlock()
91 i.status = "stopped"
92 if i.smtpSrv != nil {
93 return i.smtpSrv.Close()
94 }
95 return nil
96}
97
98func (i *Server) Status() string {
99 i.mu.Lock()
100 defer i.mu.Unlock()
101 return i.status
102}
103
104type smtpSession struct {
105 db *db.DB
106 conn *smtpserver.Conn
107 deliverer MessageDeliverer
108 authenticatedUser *model.SMTPAuth
109 toAddress string
110 fromAddress string
111 helo string
112 destinationUser model.User
113 sessionID string
114 policyEngine *policy.Engine
115 logger *slog.Logger
116}
117
118// The backend implements SMTP server methods.
119type backend struct {
120 db *db.DB
121 deliverer MessageDeliverer
122 policyEngine *policy.Engine
123}
124
125// NewSession is called after client greeting (EHLO, HELO).
126func (bkd *backend) NewSession(c *smtpserver.Conn) (smtpserver.Session, error) {
127 sessionID := uuid.NewString()
128 l := slog.With("component", "inbound", "remote_ip", c.Conn().RemoteAddr().String(), "session_id", sessionID)
129 l.Info("new session")
130 return &smtpSession{
131 db: bkd.db,
132 conn: c,
133 deliverer: bkd.deliverer,
134 authenticatedUser: &model.SMTPAuth{},
135 sessionID: sessionID,
136 logger: l,
137 policyEngine: bkd.policyEngine,
138 }, nil
139}
140
141// AuthMechanisms returns a slice of available auth mechanisms
142func (s *smtpSession) AuthMechanisms() []string {
143 return []string{sasl.Plain}
144}
145
146// Auth is the handler for supported authenticators.
147func (s *smtpSession) Auth(_ string) (sasl.Server, error) {
148 return sasl.NewPlainServer(func(identity, username, password string) error {
149 u, err := s.db.GetSMTPAuthUser(context.Background(), username)
150 if err != nil {
151 if errors.Is(err, sql.ErrNoRows) {
152 s.logger.Info("user not found")
153 } else {
154 s.logger.Error("getting user", "error", err)
155 }
156 return imapserver.ErrAuthFailed
157 }
158
159 err = u.VerifyPassword([]byte(password))
160 if err != nil {
161 slog.Info("invalid password", "username", username)
162 return imapserver.ErrAuthFailed
163 }
164
165 s.authenticatedUser = &u
166 s.logger.Info("authenticated", "user_id", u.ID, "username", username)
167 return nil
168 }), nil
169}
170
171func (s *smtpSession) Reset() {
172 s.toAddress = ""
173 s.fromAddress = ""
174}
175
176func (s *smtpSession) Logout() error {
177 s.authenticatedUser = nil
178 return nil
179}
180
181func (s *smtpSession) Mail(from string, _ *smtpserver.MailOptions) error {
182 if s.authenticatedUser == nil {
183 return smtpserver.ErrAuthFailed
184 }
185
186 s.logger.Info("MAIL FROM", "from", from)
187 s.fromAddress = from
188
189 return nil
190
191}
192
193func (s *smtpSession) Rcpt(to string, _ *smtpserver.RcptOptions) error {
194 if s.authenticatedUser == nil {
195 return smtpserver.ErrAuthFailed
196 }
197 ctx := context.Background()
198 s.logger.InfoContext(ctx, "RCPT TO", "to", to)
199
200 addr, err := mail.ParseAddress(to)
201 if err != nil {
202 s.logger.InfoContext(ctx, "RCPT TO", "error", err)
203 return &smtpserver.SMTPError{
204 Code: 501,
205 EnhancedCode: smtpserver.EnhancedCode{5, 1, 3},
206 Message: "Bad recipient address syntax",
207 }
208 }
209
210 atIndex := strings.LastIndex(addr.Address, "@")
211 if atIndex == -1 {
212 s.logger.InfoContext(ctx, "RCPT TO", "error", "malformed email: missing @ domain separator")
213 return &smtpserver.SMTPError{
214 Code: 501,
215 EnhancedCode: smtpserver.EnhancedCode{5, 1, 3},
216 Message: "Bad recipient address syntax",
217 }
218 }
219
220 localPart := addr.Address[:atIndex]
221 domain := addr.Address[atIndex+1:]
222 baseName, tag, _ := strings.Cut(localPart, "+")
223
224 s.logger.InfoContext(ctx, "parsed address", "baseName", baseName, "tag", tag, "domain", domain)
225
226 u, err := s.db.GetUserForAddress(ctx, fmt.Sprintf("%s@%s", baseName, domain))
227 if err != nil {
228 if errors.Is(err, sql.ErrNoRows) {
229 s.logger.InfoContext(ctx, "RCPT TO address not found")
230 return &smtpserver.SMTPError{
231 Code: 550,
232 EnhancedCode: smtpserver.EnhancedCode{5, 1, 1},
233 Message: "Mailbox unavailable",
234 }
235 }
236 s.logger.ErrorContext(ctx, "RCPT TO get address", "error", err)
237 return &smtpserver.SMTPError{
238 Code: 451,
239 EnhancedCode: smtpserver.EnhancedCode{4, 3, 0},
240 Message: "Temporary local problem, please try again later",
241 }
242 }
243
244 s.destinationUser = u
245 s.toAddress = to
246
247 s.logger.InfoContext(ctx, "RCPT TO", "destination_user", u.Name, "destination_user_id", u.ID)
248
249 return nil
250}
251
252func (s *smtpSession) Data(r io.Reader) error {
253 if s.authenticatedUser == nil {
254 return smtpserver.ErrAuthFailed
255 }
256
257 ctx := context.Background()
258 s.logger.InfoContext(ctx, "DATA")
259
260 var buffer bytes.Buffer
261
262 // Prepend headers
263 buffer.WriteString(fmt.Sprintf("Return-Path: %s\r\n", s.fromAddress))
264 myHostname, err := os.Hostname()
265 if err != nil {
266 myHostname = "postern.local"
267 }
268 with := "ESMTPSA" // we only allow encrypted (S) and authenticated (A) connections for inbound
269 receivedTimestampLayout := "Mon, 02 Jan 2006 15:04:05 -0700 (UTC)"
270 receivedTimestamp := time.Now().UTC().Format(receivedTimestampLayout)
271 buffer.WriteString(fmt.Sprintf("Received: from %s (%s) by %s (Postern p25.dev) with %s id %s for <%s>; %s\r\n",
272 s.conn.Hostname(), s.conn.Conn().RemoteAddr().String(), myHostname, with, s.sessionID, s.toAddress, receivedTimestamp))
273
274 _, err = io.Copy(&buffer, r)
275 if err != nil {
276 s.logger.ErrorContext(ctx, "Copy", "error", err)
277 return errInternalServerError
278 }
279
280 message := buffer.Bytes()
281 err = s.deliverer.SubmitMessage(ctx, s.conn.Conn(), s.destinationUser, bytes.NewReader(message))
282 if err != nil {
283 if errors.Is(err, model.ErrMissingFromHeader) {
284 return &smtpserver.SMTPError{
285 Code: 550,
286 EnhancedCode: smtpserver.EnhancedCode{5, 7, 1},
287 Message: "From header is required but missing",
288 }
289 }
290 s.logger.ErrorContext(ctx, "SubmitMessage", "error", err)
291 return errInternalServerError
292 }
293 return nil
294}