Documentation
use std::{net::SocketAddr, sync::Arc};

use bytes::{Buf, Bytes};
use h3_quinn::quinn;
use http::Response;
use http_body_util::{BodyExt, Full};
use rustls::ServerConfig;

use crate::{CertLoader, LoadCert, Result, Route, proxy};

pub async fn srv<D: LoadCert>(
  addr: SocketAddr,
  route: Arc<Route>,
  cert_loader: Arc<CertLoader<D>>,
) -> Result<()> {
  // 1. TLS and QUIC configuration
  let mut tls_config = ServerConfig::builder()
    .with_no_client_auth()
    .with_cert_resolver(cert_loader);
  tls_config.alpn_protocols = vec![b"h3".to_vec()];

  let server_config = quinn::ServerConfig::with_crypto(Arc::new(
    quinn::crypto::rustls::QuicServerConfig::try_from(tls_config)?,
  ));

  // 2. Create QUIC endpoint
  let endpoint = quinn::Endpoint::server(server_config, addr)?;
  println!("h3 srv listening on {}", addr);

  // 3. Accept incoming connections
  while let Some(conn) = endpoint.accept().await {
    let route = route.clone();
    tokio::spawn(async move {
      match conn.await {
        Ok(new_conn) => {
          // 4. Handle H3 connection
          let h3_conn = h3::server::Connection::new(h3_quinn::Connection::new(new_conn));
          let mut h3_conn = match h3_conn.await {
            Ok(conn) => conn,
            Err(err) => {
              eprintln!("h3 connection error: {}", err);
              return;
            }
          };

          // 5. Accept incoming requests on the H3 connection
          loop {
            match h3_conn.accept().await {
              Ok(Some(resolver)) => {
                let route = route.clone();
                tokio::spawn(async move {
                  match resolver.resolve_request().await {
                    Ok((req, mut stream)) => {
                      let (parts, _) = req.into_parts();

                      // Get the request body from the stream
                      let mut body_vec = Vec::new();
                      while let Ok(Some(mut chunk)) = stream.recv_data().await {
                        body_vec.extend_from_slice(chunk.copy_to_bytes(chunk.remaining()).as_ref());
                      }
                      let body_bytes = Bytes::from(body_vec);

                      let req = http::Request::from_parts(parts, Full::new(body_bytes));

                      let resp = proxy(req, route).await;
                      let (parts, mut body) = resp.into_parts();
                      let resp = Response::from_parts(parts, ());
                      match stream.send_response(resp).await {
                        Ok(_) => {
                          while let Some(chunk) = body.frame().await {
                            match chunk {
                              Ok(frame) => {
                                if let Some(data) = frame.data_ref()
                                  && let Err(e) = stream.send_data(data.clone()).await
                                {
                                  eprintln!("h3 send data error: {}", e);
                                  break;
                                }
                              }
                              Err(e) => {
                                eprintln!("h3 body chunk error: {}", e);
                                break;
                              }
                            }
                          }
                        }
                        Err(e) => {
                          eprintln!("h3 send response error: {}", e);
                        }
                      }
                    }
                    Err(err) => {
                      eprintln!("h3 request resolve error: {}", err);
                    }
                  }
                });
              }
              Ok(None) => {
                // Connection closed
                break;
              }
              Err(err) => {
                eprintln!("h3 accept error: {}", err);
                break;
              }
            }
          }
        }
        Err(err) => {
          eprintln!("h3 accepting connection failed: {}", err);
        }
      }
    });
  }

  Ok(())
}