Update src/main.rs

This commit is contained in:
gigirassy
2026-05-28 06:00:33 +02:00
parent 071bbdb842
commit 93f149fcc3
+226 -119
View File
@@ -1,138 +1,245 @@
use bytes::Bytes;
use http_body_util::Full; use std::{io::Cursor, net::SocketAddr, sync::Arc, time::Duration};
use hyper::body::Incoming;
use hyper::server::conn::http1; use anyhow::Result;
use hyper::service::service_fn; use axum::{
use hyper::{Method, Request, Response, StatusCode}; extract::{Query, State},
use hyper_util::rt::TokioIo; http::{HeaderMap, HeaderValue, StatusCode},
response::{IntoResponse, Response},
routing::get,
Router,
};
use image::{
codecs::jpeg::JpegEncoder,
imageops::FilterType,
GenericImageView,
};
use mimalloc::MiMalloc; use mimalloc::MiMalloc;
use std::net::SocketAddr; use reqwest::Client;
use serde::Deserialize;
use tokio::net::TcpListener; use tokio::net::TcpListener;
use tower_http::trace::TraceLayer;
use tracing::{error, info};
#[global_allocator] #[global_allocator]
static GLOBAL: MiMalloc = MiMalloc; static GLOBAL: MiMalloc = MiMalloc;
async fn handle_http(req: Request<Incoming>) -> Result<Response<Full<Bytes>>, hyper::Error> { #[derive(Clone)]
if req.method() != Method::GET { struct AppState {
return Ok(Response::builder() client: Client,
.status(StatusCode::METHOD_NOT_ALLOWED)
.body(Full::new(Bytes::from("Only GET supported")))
.unwrap());
}
let url = match req.uri().to_string().parse::<reqwest::Url>() {
Ok(u) => u,
Err(_) => {
return Ok(Response::builder()
.status(StatusCode::BAD_REQUEST)
.body(Full::new(Bytes::from("Invalid URL")))
.unwrap())
}
};
let client = reqwest::Client::new();
let resp = match client.get(url).send().await {
Ok(r) => r,
Err(_) => {
return Ok(Response::builder()
.status(StatusCode::BAD_GATEWAY)
.body(Full::new(Bytes::from("Upstream error")))
.unwrap())
}
};
let headers = resp.headers().clone();
let content_type = headers
.get(reqwest::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let bytes = match resp.bytes().await {
Ok(b) => b,
Err(_) => {
return Ok(Response::builder()
.status(StatusCode::BAD_GATEWAY)
.body(Full::new(Bytes::from("Read error")))
.unwrap())
}
};
// Only compress images
if content_type.starts_with("image/") {
if let Ok(img) = image::load_from_memory(&bytes) {
let mut out = Vec::new();
match content_type {
ct if ct.contains("jpeg") || ct.contains("jpg") => {
let _ = img.write_to(
&mut std::io::Cursor::new(&mut out),
image::ImageFormat::Jpeg,
);
}
ct if ct.contains("png") => {
let _ = img.write_to(
&mut std::io::Cursor::new(&mut out),
image::ImageFormat::Png,
);
}
_ => {
out = bytes.to_vec();
}
}
return Ok(Response::builder()
.header("content-type", content_type)
.body(Full::new(Bytes::from(out)))
.unwrap());
}
}
Ok(Response::builder()
.body(Full::new(bytes))
.unwrap())
} }
// CONNECT tunneling (HTTPS_PROXY support, but no inspection) #[derive(Deserialize)]
async fn tunnel(mut client: tokio::net::TcpStream, host: String) { struct ImageQuery {
match tokio::net::TcpStream::connect(host).await { url: String,
Ok(mut server) => { }
let (mut cr, mut cw) = client.split();
let (mut sr, mut sw) = server.split();
let client_to_server = tokio::io::copy(&mut cr, &mut sw); #[derive(Deserialize)]
let server_to_client = tokio::io::copy(&mut sr, &mut cw); struct ThumbQuery {
url: String,
let _ = tokio::join!(client_to_server, server_to_client); w: Option<u32>,
}
Err(_) => {}
}
} }
#[tokio::main] #[tokio::main]
async fn main() -> anyhow::Result<()> { async fn main() -> Result<()> {
let addr: SocketAddr = "0.0.0.0:8544".parse()?; tracing_subscriber::fmt()
.with_env_filter("info")
.init();
let client = Client::builder()
.pool_idle_timeout(Duration::from_secs(60))
.pool_max_idle_per_host(8)
.tcp_keepalive(Duration::from_secs(30))
.connect_timeout(Duration::from_secs(10))
.timeout(Duration::from_secs(20))
.user_agent("PinataImageBackend/1.0")
.build()?;
let state = Arc::new(AppState { client });
let app = Router::new()
.route("/health", get(health))
.route("/fetch", get(fetch))
.route("/thumb", get(thumb))
.layer(TraceLayer::new_for_http())
.with_state(state);
let addr: SocketAddr = "0.0.0.0:8081".parse()?;
info!("image backend listening on {}", addr);
let listener = TcpListener::bind(addr).await?; let listener = TcpListener::bind(addr).await?;
loop { axum::serve(listener, app).await?;
let (stream, _) = listener.accept().await?;
tokio::spawn(async move {
let io = TokioIo::new(stream);
let svc = service_fn(|req: Request<Incoming>| async move { Ok(())
if req.method() == Method::CONNECT { }
// HTTPS tunneling (no compression possible here)
return Ok::<_, hyper::Error>(
Response::builder()
.status(200)
.body(Full::new(Bytes::new()))
.unwrap(),
);
}
handle_http(req).await async fn health() -> &'static str {
}); "ok"
}
let _ = http1::Builder::new().serve_connection(io, svc).await; async fn fetch(
}); State(state): State<Arc<AppState>>,
Query(query): Query<ImageQuery>,
) -> Response {
match process_fetch(state, query).await {
Ok(r) => r,
Err(e) => {
error!("fetch error: {:?}", e);
(
StatusCode::BAD_GATEWAY,
"failed to process image",
)
.into_response()
}
} }
}
async fn thumb(
State(state): State<Arc<AppState>>,
Query(query): Query<ThumbQuery>,
) -> Response {
match process_thumb(state, query).await {
Ok(r) => r,
Err(e) => {
error!("thumb error: {:?}", e);
(
StatusCode::BAD_GATEWAY,
"failed to process thumbnail",
)
.into_response()
}
}
}
fn validate_url(url: &str) -> Result<()> {
let parsed = reqwest::Url::parse(url)?;
if parsed.scheme() != "https" {
anyhow::bail!("https required");
}
let host = parsed.host_str().unwrap_or("");
if !host.eq_ignore_ascii_case("i.pinimg.com") {
anyhow::bail!("invalid host");
}
Ok(())
}
async fn download_image(client: &Client, url: &str) -> Result<bytes::Bytes> {
let resp = client.get(url).send().await?;
if !resp.status().is_success() {
anyhow::bail!("upstream status {}", resp.status());
}
Ok(resp.bytes().await?)
}
async fn process_fetch(
state: Arc<AppState>,
query: ImageQuery,
) -> Result<Response> {
validate_url(&query.url)?;
let data = download_image(&state.client, &query.url).await?;
let img = image::load_from_memory(&data)?;
let rgb = img.to_rgb8();
let mut out = Vec::with_capacity(data.len() / 2);
{
let mut cursor = Cursor::new(&mut out);
let mut encoder =
JpegEncoder::new_with_quality(&mut cursor, 75);
encoder.encode_image(&rgb)?;
}
let mut headers = HeaderMap::new();
headers.insert(
"content-type",
HeaderValue::from_static("image/jpeg"),
);
headers.insert(
"cache-control",
HeaderValue::from_static(
"public, max-age=86400, immutable",
),
);
Ok((headers, out).into_response())
}
async fn process_thumb(
state: Arc<AppState>,
query: ThumbQuery,
) -> Result<Response> {
validate_url(&query.url)?;
let width = query.w.unwrap_or(260).clamp(64, 2048);
let data = download_image(&state.client, &query.url).await?;
let img = image::load_from_memory(&data)?;
let (src_w, src_h) = img.dimensions();
if width >= src_w {
return process_fetch(
state,
ImageQuery {
url: query.url,
},
)
.await;
}
let height =
((src_h as f32 * width as f32) / src_w as f32) as u32;
let resized = img.resize_exact(
width,
height.max(1),
FilterType::Triangle,
);
let rgb = resized.to_rgb8();
let mut out = Vec::new();
{
let mut cursor = Cursor::new(&mut out);
let mut encoder =
JpegEncoder::new_with_quality(&mut cursor, 72);
encoder.encode_image(&rgb)?;
}
let mut headers = HeaderMap::new();
headers.insert(
"content-type",
HeaderValue::from_static("image/jpeg"),
);
headers.insert(
"cache-control",
HeaderValue::from_static(
"public, max-age=86400, immutable",
),
);
Ok((headers, out).into_response())
} }