1use std::io;
14use std::net::SocketAddr;
15use std::sync::Arc;
16use std::time::{Duration, Instant};
17
18use crate::wire::config::{Config, ErrorConvention, Handshake, HelloStyle, PushPolicy};
19use crate::wire::{encode_frame, read_request_with_limit, Request, Response, Value, PUSH_ID};
20use tokio::io::{AsyncWriteExt, BufReader, BufWriter};
21use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
22use tokio::net::{TcpListener, TcpStream};
23use tokio::sync::{mpsc, watch, Semaphore};
24
25use crate::server::dispatch::{AuthError, Credentials, Dispatch};
26use crate::server::errors::{format_bracket_code, format_err, NOAUTH, WRONGPASS};
27use crate::server::metrics::{Metrics, MetricsSnapshot};
28use crate::server::session::{PushSender, Session, WriteJob};
29
30const PROTO_VERSION: i64 = 1;
33
34const PRE_AUTH_COMMANDS: &[&str] = &["PING", "HELLO", "AUTH", "QUIT"];
36
37const WRITER_QUEUE_DEPTH: usize = 64;
39
40#[derive(Debug, Clone)]
42pub struct ServerInfo {
43 pub name: String,
45 pub version: String,
47}
48
49#[derive(Debug, Clone)]
52pub struct ListenerConfig {
53 pub addr: SocketAddr,
56 pub idle_timeout: Duration,
59 pub slow_threshold: Duration,
62 pub auth_required: bool,
75}
76
77impl ListenerConfig {
78 pub fn new(addr: SocketAddr) -> Self {
81 Self {
82 addr,
83 idle_timeout: Duration::ZERO,
84 slow_threshold: Duration::from_millis(1000),
85 auth_required: true,
86 }
87 }
88
89 pub fn open(mut self) -> Self {
95 self.auth_required = false;
96 self
97 }
98}
99
100impl Default for ListenerConfig {
101 fn default() -> Self {
103 Self::new(SocketAddr::from(([127, 0, 0, 1], 0)))
104 }
105}
106
107struct ConnShared<D> {
109 dispatch: Arc<D>,
110 profile: Config,
111 info: ServerInfo,
112 idle_timeout: Duration,
113 slow_threshold: Duration,
114 auth_required: bool,
115 metrics: Arc<Metrics>,
116}
117
118#[derive(Debug)]
125pub struct ListenerHandle {
126 local_addr: SocketAddr,
127 shutdown: watch::Sender<bool>,
128 metrics: Arc<Metrics>,
129 done: Option<mpsc::Receiver<()>>,
130}
131
132impl ListenerHandle {
133 pub fn local_addr(&self) -> SocketAddr {
135 self.local_addr
136 }
137
138 pub fn snapshot(&self) -> MetricsSnapshot {
140 self.metrics.snapshot()
141 }
142
143 pub async fn stop(mut self) {
146 let _ = self.shutdown.send(true);
147 if let Some(mut done) = self.done.take() {
148 let _ = done.recv().await;
151 }
152 }
153}
154
155impl Drop for ListenerHandle {
156 fn drop(&mut self) {
157 let _ = self.shutdown.send(true);
159 }
160}
161
162pub async fn spawn_listener<D: Dispatch>(
165 dispatch: Arc<D>,
166 profile: Config,
167 info: ServerInfo,
168 config: ListenerConfig,
169) -> io::Result<ListenerHandle> {
170 let listener = TcpListener::bind(config.addr).await?;
171 let local_addr = listener.local_addr()?;
172 let metrics = Arc::new(Metrics::default());
173 let (shutdown_tx, shutdown_rx) = watch::channel(false);
174 let (done_tx, done_rx) = mpsc::channel::<()>(1);
175
176 let shared = Arc::new(ConnShared {
177 dispatch,
178 profile,
179 info,
180 idle_timeout: config.idle_timeout,
181 slow_threshold: config.slow_threshold,
182 auth_required: config.auth_required,
183 metrics: Arc::clone(&metrics),
184 });
185
186 tokio::spawn(accept_loop(listener, shared, shutdown_rx, done_tx));
187
188 Ok(ListenerHandle {
189 local_addr,
190 shutdown: shutdown_tx,
191 metrics,
192 done: Some(done_rx),
193 })
194}
195
196async fn accept_loop<D: Dispatch>(
200 listener: TcpListener,
201 shared: Arc<ConnShared<D>>,
202 shutdown: watch::Receiver<bool>,
203 done: mpsc::Sender<()>,
204) {
205 let mut accept_shutdown = shutdown.clone();
206 let mut next_conn_id: u64 = 1;
207 loop {
208 let accepted = tokio::select! {
209 _ = accept_shutdown.wait_for(|stop| *stop) => break,
210 accepted = listener.accept() => accepted,
211 };
212 let Ok((stream, _peer)) = accepted else {
213 continue;
214 };
215 let conn_id = next_conn_id;
216 next_conn_id = next_conn_id.wrapping_add(1);
217 let ctx = Arc::clone(&shared);
218 let conn_shutdown = shutdown.clone();
219 let done_guard = done.clone();
220 ctx.metrics.connection_opened();
221 tokio::spawn(async move {
222 handle_connection(stream, &ctx, conn_id, conn_shutdown).await;
223 ctx.metrics.connection_closed();
224 drop(done_guard);
225 });
226 }
227 }
230
231async fn handle_connection<D: Dispatch>(
235 stream: TcpStream,
236 ctx: &ConnShared<D>,
237 conn_id: u64,
238 mut shutdown: watch::Receiver<bool>,
239) {
240 let _ = stream.set_nodelay(true);
243 let (read_half, write_half) = stream.into_split();
244 let mut reader = BufReader::new(read_half);
245
246 let (tx, rx) = mpsc::channel::<WriteJob>(WRITER_QUEUE_DEPTH);
247 let write_task = tokio::spawn(writer_task(
248 BufWriter::new(write_half),
249 rx,
250 Arc::clone(&ctx.metrics),
251 ctx.slow_threshold,
252 ));
253
254 let push = match ctx.profile.push {
257 PushPolicy::Enabled => Some(PushSender::new(tx.clone())),
258 PushPolicy::Reserved => None,
259 };
260 let starts_authenticated =
267 matches!(ctx.profile.handshake, Handshake::None) || !ctx.auth_required;
268 let session = Arc::new(Session::new(conn_id, starts_authenticated, push));
269
270 let permits = ctx.profile.max_in_flight.clamp(1, u32::MAX as usize) as u32;
271 let in_flight = Arc::new(Semaphore::new(permits as usize));
272
273 let mut first_frame = true;
274 loop {
275 let read = tokio::select! {
276 _ = shutdown.wait_for(|stop| *stop) => break,
277 read = read_next(&mut reader, ctx.profile.max_frame_bytes, ctx.idle_timeout) => read,
278 };
279 let Ok((req, in_bytes)) = read else { break };
284
285 if req.id == PUSH_ID {
288 let response = Response::err(req.id, push_refusal_error(&ctx.profile));
289 if !send_inline(&tx, response, in_bytes).await {
290 break;
291 }
292 continue;
293 }
294
295 if first_frame {
298 first_frame = false;
299 if matches!(ctx.profile.handshake, Handshake::HelloMandatory) && req.command != "HELLO"
300 {
301 let response = Response::err(req.id, hello_required_error(&ctx.profile));
302 let _ = send_inline(&tx, response, in_bytes).await;
303 break;
304 }
305 }
306
307 match req.command.as_str() {
311 "HELLO" if !matches!(ctx.profile.hello_style, HelloStyle::NotUsed) => {
313 let response = handle_hello(ctx, &session, req.id, &req.args).await;
314 if !send_inline(&tx, response, in_bytes).await {
315 break;
316 }
317 continue;
318 }
319 "AUTH" if matches!(ctx.profile.handshake, Handshake::AuthCommand) => {
321 let response = handle_auth(ctx, &session, req.id, &req.args).await;
322 if !send_inline(&tx, response, in_bytes).await {
323 break;
324 }
325 continue;
326 }
327 "PING" if !session.is_authenticated() => {
330 let response = builtin_ping(req.id, &req.args);
331 if !send_inline(&tx, response, in_bytes).await {
332 break;
333 }
334 continue;
335 }
336 "QUIT" if matches!(ctx.profile.handshake, Handshake::AuthCommand) => {
338 let response = Response::ok(req.id, Value::Str("OK".to_owned()));
339 let _ = send_inline(&tx, response, in_bytes).await;
340 break;
341 }
342 _ => {}
343 }
344
345 if !session.is_authenticated() {
347 match ctx.profile.handshake {
348 Handshake::None => {}
349 Handshake::AuthCommand => {
350 if !PRE_AUTH_COMMANDS.contains(&req.command.as_str()) {
351 let response = Response::err(req.id, NOAUTH);
352 if !send_inline(&tx, response, in_bytes).await {
353 break;
354 }
355 continue;
356 }
357 }
360 Handshake::HelloMandatory => {
361 let response = Response::err(req.id, hello_required_error(&ctx.profile));
362 if !send_inline(&tx, response, in_bytes).await {
363 break;
364 }
365 continue;
366 }
367 }
368 }
369
370 let Ok(permit) = Arc::clone(&in_flight).acquire_owned().await else {
374 break;
375 };
376 let dispatch = Arc::clone(&ctx.dispatch);
377 let session = Arc::clone(&session);
378 let tx = tx.clone();
379 tokio::spawn(async move {
380 let started = Instant::now();
381 let Request { id, command, args } = req;
382 let response = match dispatch.dispatch(&session, &command, args).await {
383 Ok(value) => Response::ok(id, value),
384 Err(message) => Response::err(id, message),
387 };
388 let _ = tx
389 .send(WriteJob::Response {
390 response,
391 in_bytes,
392 duration: started.elapsed(),
393 })
394 .await;
395 drop(permit);
398 });
399 }
400
401 let _ = in_flight.acquire_many(permits).await;
407 let _ = tx.send(WriteJob::Shutdown).await;
408 drop(tx);
409 let _ = write_task.await;
410}
411
412async fn read_next(
416 reader: &mut BufReader<OwnedReadHalf>,
417 max_frame_bytes: usize,
418 idle_timeout: Duration,
419) -> io::Result<(Request, usize)> {
420 if idle_timeout.is_zero() {
421 read_request_with_limit(reader, max_frame_bytes).await
422 } else {
423 match tokio::time::timeout(
424 idle_timeout,
425 read_request_with_limit(reader, max_frame_bytes),
426 )
427 .await
428 {
429 Ok(read) => read,
430 Err(_) => Err(io::Error::new(io::ErrorKind::TimedOut, "idle timeout")),
431 }
432 }
433}
434
435async fn send_inline(tx: &mpsc::Sender<WriteJob>, response: Response, in_bytes: usize) -> bool {
439 tx.send(WriteJob::Response {
440 response,
441 in_bytes,
442 duration: Duration::ZERO,
443 })
444 .await
445 .is_ok()
446}
447
448async fn writer_task(
453 mut writer: BufWriter<OwnedWriteHalf>,
454 mut rx: mpsc::Receiver<WriteJob>,
455 metrics: Arc<Metrics>,
456 slow_threshold: Duration,
457) {
458 'outer: while let Some(job) = rx.recv().await {
459 if !write_job(&mut writer, job, &metrics, slow_threshold).await {
460 break;
461 }
462 while let Ok(job) = rx.try_recv() {
463 if !write_job(&mut writer, job, &metrics, slow_threshold).await {
464 break 'outer;
465 }
466 }
467 if writer.flush().await.is_err() {
468 break;
469 }
470 }
471 let _ = writer.flush().await;
473}
474
475async fn write_job(
479 writer: &mut BufWriter<OwnedWriteHalf>,
480 job: WriteJob,
481 metrics: &Metrics,
482 slow_threshold: Duration,
483) -> bool {
484 let (response, in_bytes, duration) = match job {
485 WriteJob::Shutdown => return false,
486 WriteJob::Push(response) => {
487 let Ok(frame) = encode_frame(&response) else {
488 return true;
489 };
490 if writer.write_all(&frame).await.is_err() {
491 return false;
492 }
493 metrics.record_push(frame.len());
494 return true;
495 }
496 WriteJob::Response {
497 response,
498 in_bytes,
499 duration,
500 } => (response, in_bytes, duration),
501 };
502 let is_error = response.result.is_err();
503 let Ok(frame) = encode_frame(&response) else {
506 return true;
509 };
510 if writer.write_all(&frame).await.is_err() {
511 return false;
512 }
513 metrics.record_command(in_bytes, frame.len(), duration, is_error, slow_threshold);
515 true
516}
517
518async fn handle_hello<D: Dispatch>(
527 ctx: &ConnShared<D>,
528 session: &Session,
529 req_id: u32,
530 args: &[Value],
531) -> Response {
532 match ctx.profile.hello_style {
533 HelloStyle::NotUsed => {
535 Response::err(req_id, format_err("HELLO is not part of this profile"))
536 }
537 HelloStyle::ArgLess => Response::ok(
540 req_id,
541 Value::Map(vec![
542 (
543 Value::Str("server".to_owned()),
544 Value::Str(ctx.info.name.clone()),
545 ),
546 (
547 Value::Str("version".to_owned()),
548 Value::Str(ctx.info.version.clone()),
549 ),
550 (Value::Str("proto".to_owned()), Value::Int(PROTO_VERSION)),
551 (
552 Value::Str("id".to_owned()),
553 Value::Int(session.connection_id() as i64),
554 ),
555 (
556 Value::Str("authenticated".to_owned()),
557 Value::Bool(session.is_authenticated()),
558 ),
559 ]),
560 ),
561 HelloStyle::MapPayload => {
563 let creds = match parse_hello_credentials(args) {
564 Ok(creds) => creds,
565 Err(message) => return Response::err(req_id, message),
566 };
567 match ctx.dispatch.authenticate(creds).await {
568 Ok(principal) => {
569 let capabilities = ctx.dispatch.capabilities(&principal);
570 session.set_principal(principal);
571 Response::ok(
572 req_id,
573 Value::Map(vec![
574 (
575 Value::Str("protocol_version".to_owned()),
576 Value::Int(PROTO_VERSION),
577 ),
578 (
579 Value::Str("capabilities".to_owned()),
580 Value::Array(capabilities.into_iter().map(Value::Str).collect()),
581 ),
582 ]),
583 )
584 }
585 Err(err) => Response::err(req_id, auth_error_string(&ctx.profile, err)),
588 }
589 }
590 }
591}
592
593async fn handle_auth<D: Dispatch>(
596 ctx: &ConnShared<D>,
597 session: &Session,
598 req_id: u32,
599 args: &[Value],
600) -> Response {
601 let creds = match args {
602 [key] => value_str(key).map(Credentials::ApiKey),
603 [user, pass] => value_str(user)
604 .zip(value_str(pass))
605 .map(|(user, pass)| Credentials::UserPass(user, pass)),
606 _ => None,
607 };
608 let Some(creds) = creds else {
609 return Response::err(req_id, format_err("invalid arguments for 'AUTH'"));
610 };
611 match ctx.dispatch.authenticate(creds).await {
612 Ok(principal) => {
613 session.set_principal(principal);
614 Response::ok(req_id, Value::Str("OK".to_owned()))
615 }
616 Err(err) => Response::err(req_id, auth_error_string(&ctx.profile, err)),
617 }
618}
619
620fn builtin_ping(req_id: u32, args: &[Value]) -> Response {
623 match args {
624 [] => Response::ok(req_id, Value::Str("PONG".to_owned())),
625 [Value::Str(payload)] => Response::ok(req_id, Value::Str(payload.clone())),
626 [Value::Bytes(payload)] => Response::ok(req_id, Value::Bytes(payload.clone())),
627 [_] => Response::err(
628 req_id,
629 format_err("PING argument must be a string or bytes"),
630 ),
631 args => Response::err(
632 req_id,
633 format_err(&format!(
634 "wrong number of arguments for 'PING' ({})",
635 args.len()
636 )),
637 ),
638 }
639}
640
641fn parse_hello_credentials(args: &[Value]) -> Result<Credentials, String> {
645 let map = match args.first() {
646 None => return Ok(Credentials::None),
647 Some(map @ Value::Map(_)) => map,
648 Some(_) => return Err(format_err("HELLO expects a Map argument")),
649 };
650 if let Some(token) = map.map_get("token").and_then(value_str) {
651 Ok(Credentials::Token(token))
652 } else if let Some(key) = map.map_get("api_key").and_then(value_str) {
653 Ok(Credentials::ApiKey(key))
654 } else {
655 Ok(Credentials::None)
656 }
657}
658
659fn value_str(value: &Value) -> Option<String> {
662 match value {
663 Value::Str(text) => Some(text.clone()),
664 Value::Bytes(bytes) => String::from_utf8(bytes.clone()).ok(),
665 _ => None,
666 }
667}
668
669fn auth_error_string(profile: &Config, err: AuthError) -> String {
674 match err {
675 AuthError::Message(message) => message,
676 AuthError::InvalidCredentials => match profile.error_codes {
677 ErrorConvention::BracketCode | ErrorConvention::Both => {
678 format_bracket_code("unauthorized", "invalid credentials")
679 }
680 _ => WRONGPASS.to_owned(),
681 },
682 }
683}
684
685fn hello_required_error(profile: &Config) -> String {
687 match profile.error_codes {
688 ErrorConvention::BracketCode | ErrorConvention::Both => {
689 format_bracket_code("unauthorized", "authentication required: send HELLO first")
690 }
691 _ => NOAUTH.to_owned(),
692 }
693}
694
695fn push_refusal_error(profile: &Config) -> String {
697 const MESSAGE: &str = "request id u32::MAX is reserved for server push frames";
698 match profile.error_codes {
699 ErrorConvention::BracketCode | ErrorConvention::Both => {
700 format_bracket_code("reserved_frame_id", MESSAGE)
701 }
702 _ => format_err(MESSAGE),
703 }
704}