1use core::fmt;
45
46use alloc::{
47 string::{String, ToString},
48 vec::Vec,
49};
50
51use bounded_static::IntoBoundedStatic;
52use log::trace;
53use thiserror::Error;
54
55use crate::{
56 coroutine::*,
57 rfc5321::types::{domain::Domain, reply_code::ReplyCode},
58 send::*,
59 smtp_try,
60};
61
62pub struct SmtpHeloCommand<'a> {
64 pub domain: Domain<'a>,
66}
67
68impl<'a> From<SmtpHeloCommand<'a>> for Vec<u8> {
69 fn from(cmd: SmtpHeloCommand<'a>) -> Vec<u8> {
70 let mut buf = String::from("HELO ");
71 buf.push_str(&cmd.domain.to_string());
72 buf.push_str("\r\n");
73 buf.into_bytes()
74 }
75}
76
77#[derive(Clone, Debug, Error)]
79pub enum SmtpHeloError {
80 #[error("SMTP HELO failed: rejected {code} {message}")]
81 Rejected { code: u16, message: String },
82 #[error("SMTP HELO failed: {0}")]
83 Send(#[from] SendSmtpCommandError),
84}
85
86pub struct SmtpHelo {
88 state: State,
89}
90
91impl SmtpHelo {
92 pub fn new(domain: Domain<'_>) -> Self {
93 let cmd = SmtpHeloCommand {
94 domain: domain.into_static(),
95 };
96
97 Self {
98 state: State::Send(SendSmtpCommand::new(cmd)),
99 }
100 }
101}
102
103impl SmtpCoroutine for SmtpHelo {
104 type Yield = SmtpYield;
105 type Return = Result<(), SmtpHeloError>;
106
107 fn resume(&mut self, arg: Option<&[u8]>) -> SmtpCoroutineState<Self::Yield, Self::Return> {
108 loop {
109 trace!("helo: {}", self.state);
110
111 match &mut self.state {
112 State::Send(send) => {
113 let out = smtp_try!(send, arg);
114
115 if out.response.code == ReplyCode::OK {
116 return SmtpCoroutineState::Complete(Ok(()));
117 }
118
119 let code = out.response.code.code();
120 let message = out.response.text().to_string();
121 return SmtpCoroutineState::Complete(Err(SmtpHeloError::Rejected {
122 code,
123 message,
124 }));
125 }
126 }
127 }
128 }
129}
130
131enum State {
132 Send(SendSmtpCommand<SmtpHeloCommand<'static>>),
133}
134
135impl fmt::Display for State {
136 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
137 match self {
138 Self::Send(_) => f.write_str("send helo"),
139 }
140 }
141}
142
143#[cfg(test)]
144mod tests {
145 use alloc::borrow::Cow;
146
147 use super::*;
148
149 fn domain() -> Domain<'static> {
150 Domain(Cow::Borrowed("example.com"))
151 }
152
153 #[test]
154 fn success_returns_ok() {
155 let mut helo = SmtpHelo::new(domain());
156
157 let bytes = expect_wants_write(&mut helo, None);
158 assert_eq!(bytes, b"HELO example.com\r\n");
159
160 expect_wants_read(&mut helo);
161 expect_complete_ok(&mut helo, b"250 server.example.com\r\n");
162 }
163
164 #[test]
165 fn rejected_returns_rejected_error() {
166 let mut helo = SmtpHelo::new(domain());
167 let _ = expect_wants_write(&mut helo, None);
168 expect_wants_read(&mut helo);
169
170 let err = expect_complete_err(&mut helo, b"550 bad domain\r\n");
171 let SmtpHeloError::Rejected { code, message } = err else {
172 panic!("expected SmtpHeloError::Rejected, got {err:?}");
173 };
174 assert_eq!(code, 550);
175 assert_eq!(message, "bad domain");
176 }
177
178 #[test]
179 fn eof_returns_eof_error() {
180 let mut helo = SmtpHelo::new(domain());
181 let _ = expect_wants_write(&mut helo, None);
182 expect_wants_read(&mut helo);
183
184 let err = expect_complete_err(&mut helo, b"");
185 assert!(matches!(
186 err,
187 SmtpHeloError::Send(SendSmtpCommandError::Eof)
188 ));
189 }
190
191 fn expect_wants_write(cor: &mut SmtpHelo, arg: Option<&[u8]>) -> Vec<u8> {
194 match cor.resume(arg) {
195 SmtpCoroutineState::Yielded(SmtpYield::WantsWrite(bytes)) => bytes,
196 state => panic!("expected WantsWrite, got {state:?}"),
197 }
198 }
199
200 fn expect_wants_read(cor: &mut SmtpHelo) {
201 match cor.resume(None) {
202 SmtpCoroutineState::Yielded(SmtpYield::WantsRead) => {}
203 state => panic!("expected WantsRead, got {state:?}"),
204 }
205 }
206
207 fn expect_complete_ok(cor: &mut SmtpHelo, reply: &[u8]) {
208 match cor.resume(Some(reply)) {
209 SmtpCoroutineState::Complete(Ok(())) => {}
210 state => panic!("expected Complete(Ok), got {state:?}"),
211 }
212 }
213
214 fn expect_complete_err(cor: &mut SmtpHelo, reply: &[u8]) -> SmtpHeloError {
215 match cor.resume(Some(reply)) {
216 SmtpCoroutineState::Complete(Err(err)) => err,
217 state => panic!("expected Complete(Err), got {state:?}"),
218 }
219 }
220}