3996 lines
128 KiB
Rust
3996 lines
128 KiB
Rust
use std::{
|
||
collections::HashMap,
|
||
env,
|
||
fs::OpenOptions,
|
||
io,
|
||
io::SeekFrom,
|
||
io::Write,
|
||
net::{IpAddr, SocketAddr},
|
||
path::{Path, PathBuf},
|
||
sync::{Arc, Mutex},
|
||
time::{Duration, Instant, SystemTime, UNIX_EPOCH},
|
||
};
|
||
|
||
use async_trait::async_trait;
|
||
use bytes::Bytes;
|
||
use httpdate::{fmt_http_date, parse_http_date};
|
||
use pingora::{
|
||
ConnectTimedout, ConnectionClosed, Error, ErrorSource, HTTPStatus, ReadError, ReadTimedout,
|
||
Result as PingoraResult, WriteError, WriteTimedout,
|
||
listeners::tls::TlsSettings,
|
||
modules::http::{HttpModule, HttpModuleBuilder, HttpModules, Module},
|
||
prelude::HttpPeer,
|
||
protocols::http::compression::ResponseCompressionCtx,
|
||
server::{Server, configuration::Opt},
|
||
};
|
||
use pingora_http::{Method, RequestHeader, ResponseHeader};
|
||
use pingora_proxy::{FailToProxy, ProxyHttp, Session, http_proxy_service};
|
||
use shared_logging::{OtelConfig, init_tracing};
|
||
use tokio::io::{AsyncReadExt, AsyncSeekExt};
|
||
use tracing::{error, info, warn};
|
||
use uuid::Uuid;
|
||
|
||
const DEFAULT_LISTEN_ADDR: &str = "127.0.0.1:18081";
|
||
const DEFAULT_API_UPSTREAM: &str = "127.0.0.1:8082";
|
||
const DEFAULT_SPACETIME_UPSTREAM: &str = "127.0.0.1:3101";
|
||
const DEFAULT_WEB_ROOT: &str = "/srv/genarrative/web";
|
||
const DEFAULT_ACME_ROOT: &str = "/var/www/html";
|
||
const DEFAULT_MAINTENANCE_FILE: &str = "/var/lib/genarrative/maintenance/enabled";
|
||
const DEFAULT_MAINTENANCE_PAGE_FILE: &str = "/var/lib/genarrative/maintenance/page.html";
|
||
const DEFAULT_MAX_API_BODY_BYTES: u64 = 64 * 1024 * 1024;
|
||
const DEFAULT_GZIP_LEVEL: u32 = 5;
|
||
const DEFAULT_GZIP_MIN_LENGTH_BYTES: u64 = 1024;
|
||
const DEFAULT_UPSTREAM_CONNECT_TIMEOUT_MS: u64 = 3_000;
|
||
const DEFAULT_UPSTREAM_DEFAULT_READ_TIMEOUT_SECONDS: u64 = 60;
|
||
const DEFAULT_UPSTREAM_API_READ_TIMEOUT_SECONDS: u64 = 3_600;
|
||
const DEFAULT_UPSTREAM_LONG_READ_TIMEOUT_SECONDS: u64 = 3_600;
|
||
const DEFAULT_UPSTREAM_WRITE_TIMEOUT_SECONDS: u64 = 3_600;
|
||
const DEFAULT_HTML_CACHE_CONTROL: &str = "no-cache";
|
||
const DEFAULT_ASSET_CACHE_CONTROL: &str = "public, max-age=31536000, immutable";
|
||
const DEFAULT_STATIC_CACHE_CONTROL: &str = "no-cache";
|
||
const DEFAULT_INSTANCE_COUNT: u32 = 1;
|
||
const DEFAULT_LOG_FILTER: &str = "info,pingora=info,pingora_gateway=info";
|
||
const SHADOW_PROBE_PATH: &str = "/__genarrative_pingora/healthz";
|
||
const SHADOW_PROBE_HEADER: &str = "x-genarrative-pingora-probe";
|
||
const PAYLOAD_TOO_LARGE_CONTEXT: &str = "genarrative_payload_too_large";
|
||
const PROTECTION_STATE_TTL: Duration = Duration::from_secs(600);
|
||
const PROTECTION_CLEANUP_INTERVAL: Duration = Duration::from_secs(60);
|
||
const MAIN_SPA_PATHS: &[&str] = &["/", "/creation", "/editor/canvas", "/profile", "/project"];
|
||
|
||
#[derive(Clone, Debug)]
|
||
struct GatewayConfig {
|
||
listen_addr: String,
|
||
tls_listen_addr: Option<String>,
|
||
tls_cert_file: Option<PathBuf>,
|
||
tls_key_file: Option<PathBuf>,
|
||
http_redirect_listen_addr: Option<String>,
|
||
http_redirect_target_scheme: String,
|
||
api_upstream: SocketAddr,
|
||
spacetime_upstream: SocketAddr,
|
||
gitea_hosts: Vec<String>,
|
||
gitea_upstream: Option<SocketAddr>,
|
||
web_root: PathBuf,
|
||
acme_root: PathBuf,
|
||
maintenance_file: PathBuf,
|
||
maintenance_page_file: PathBuf,
|
||
forwarded_proto: String,
|
||
max_api_body_bytes: u64,
|
||
gzip_enabled: bool,
|
||
gzip_level: u32,
|
||
gzip_min_length_bytes: u64,
|
||
compression: CompressionConfig,
|
||
static_cache: StaticCacheConfig,
|
||
upstream_timeouts: UpstreamTimeoutConfig,
|
||
shadow_probe_token: Option<String>,
|
||
trust_x_forwarded_for: bool,
|
||
instance_count: u32,
|
||
shared_protection_confirmed: bool,
|
||
protection: ProtectionConfig,
|
||
log_filter: String,
|
||
access_log_file: Option<PathBuf>,
|
||
otel_enabled: bool,
|
||
}
|
||
|
||
impl GatewayConfig {
|
||
fn from_env() -> io::Result<Self> {
|
||
let config = Self {
|
||
listen_addr: read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_LISTEN",
|
||
DEFAULT_LISTEN_ADDR,
|
||
),
|
||
tls_listen_addr: read_optional_env("GENARRATIVE_PINGORA_GATEWAY_TLS_LISTEN"),
|
||
tls_cert_file: read_optional_env("GENARRATIVE_PINGORA_GATEWAY_TLS_CERT_FILE")
|
||
.map(PathBuf::from),
|
||
tls_key_file: read_optional_env("GENARRATIVE_PINGORA_GATEWAY_TLS_KEY_FILE")
|
||
.map(PathBuf::from),
|
||
http_redirect_listen_addr: read_optional_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTTP_REDIRECT_LISTEN",
|
||
),
|
||
http_redirect_target_scheme: read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTTP_REDIRECT_TARGET_SCHEME",
|
||
"https",
|
||
),
|
||
api_upstream: read_socket_addr_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_API_UPSTREAM",
|
||
DEFAULT_API_UPSTREAM,
|
||
)?,
|
||
spacetime_upstream: read_socket_addr_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_SPACETIME_UPSTREAM",
|
||
DEFAULT_SPACETIME_UPSTREAM,
|
||
)?,
|
||
gitea_hosts: read_host_list_env("GENARRATIVE_PINGORA_GATEWAY_GITEA_HOSTS")?,
|
||
gitea_upstream: read_optional_socket_addr_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_GITEA_UPSTREAM",
|
||
)?,
|
||
web_root: PathBuf::from(read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_WEB_ROOT",
|
||
DEFAULT_WEB_ROOT,
|
||
)),
|
||
acme_root: PathBuf::from(read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_ACME_ROOT",
|
||
DEFAULT_ACME_ROOT,
|
||
)),
|
||
maintenance_file: PathBuf::from(read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_MAINTENANCE_FILE",
|
||
DEFAULT_MAINTENANCE_FILE,
|
||
)),
|
||
maintenance_page_file: PathBuf::from(read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_MAINTENANCE_PAGE_FILE",
|
||
DEFAULT_MAINTENANCE_PAGE_FILE,
|
||
)),
|
||
forwarded_proto: read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_FORWARDED_PROTO",
|
||
"http",
|
||
),
|
||
max_api_body_bytes: read_u64_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_MAX_API_BODY_BYTES",
|
||
DEFAULT_MAX_API_BODY_BYTES,
|
||
)?,
|
||
gzip_enabled: read_bool_env("GENARRATIVE_PINGORA_GATEWAY_GZIP_ENABLED", true),
|
||
gzip_level: read_u32_env("GENARRATIVE_PINGORA_GATEWAY_GZIP_LEVEL", DEFAULT_GZIP_LEVEL)?,
|
||
gzip_min_length_bytes: read_u64_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_GZIP_MIN_LENGTH_BYTES",
|
||
DEFAULT_GZIP_MIN_LENGTH_BYTES,
|
||
)?,
|
||
compression: CompressionConfig::from_env()?,
|
||
static_cache: StaticCacheConfig::from_env(),
|
||
upstream_timeouts: UpstreamTimeoutConfig::from_env()?,
|
||
shadow_probe_token: read_optional_env("GENARRATIVE_PINGORA_GATEWAY_PROBE_TOKEN"),
|
||
trust_x_forwarded_for: read_bool_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_TRUST_X_FORWARDED_FOR",
|
||
false,
|
||
),
|
||
instance_count: read_u32_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_INSTANCE_COUNT",
|
||
DEFAULT_INSTANCE_COUNT,
|
||
)?,
|
||
shared_protection_confirmed: read_bool_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_SHARED_PROTECTION_CONFIRMED",
|
||
false,
|
||
),
|
||
protection: ProtectionConfig::from_env()?,
|
||
log_filter: read_env_or_default("GENARRATIVE_PINGORA_GATEWAY_LOG", DEFAULT_LOG_FILTER),
|
||
access_log_file: read_optional_env("GENARRATIVE_PINGORA_GATEWAY_ACCESS_LOG_FILE")
|
||
.map(PathBuf::from),
|
||
otel_enabled: read_bool_env("GENARRATIVE_PINGORA_GATEWAY_OTEL_ENABLED", false),
|
||
};
|
||
config.validate()?;
|
||
Ok(config)
|
||
}
|
||
|
||
fn validate(&self) -> io::Result<()> {
|
||
let listen_addr =
|
||
parse_socket_addr_config("GENARRATIVE_PINGORA_GATEWAY_LISTEN", &self.listen_addr)?;
|
||
|
||
if listen_addr == self.api_upstream || listen_addr == self.spacetime_upstream {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_LISTEN 不能与上游地址相同",
|
||
));
|
||
}
|
||
self.validate_gitea_config(listen_addr, None, None)?;
|
||
|
||
let tls_listen_addr = if let Some(tls_listen_addr) = &self.tls_listen_addr {
|
||
let tls_listen_addr = parse_socket_addr_config(
|
||
"GENARRATIVE_PINGORA_GATEWAY_TLS_LISTEN",
|
||
tls_listen_addr,
|
||
)?;
|
||
validate_distinct_listen_addr(
|
||
"GENARRATIVE_PINGORA_GATEWAY_TLS_LISTEN",
|
||
tls_listen_addr,
|
||
listen_addr,
|
||
self.api_upstream,
|
||
self.spacetime_upstream,
|
||
self.gitea_upstream,
|
||
)?;
|
||
self.validate_gitea_config(listen_addr, Some(tls_listen_addr), None)?;
|
||
validate_tls_file_config(&self.tls_cert_file, &self.tls_key_file)?;
|
||
Some(tls_listen_addr)
|
||
} else if self.tls_cert_file.is_some() || self.tls_key_file.is_some() {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"配置 TLS_CERT_FILE / TLS_KEY_FILE 时必须同时设置 GENARRATIVE_PINGORA_GATEWAY_TLS_LISTEN",
|
||
));
|
||
} else {
|
||
None
|
||
};
|
||
|
||
if let Some(http_redirect_listen_addr) = &self.http_redirect_listen_addr {
|
||
let Some(tls_listen_addr) = tls_listen_addr else {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"配置 GENARRATIVE_PINGORA_GATEWAY_HTTP_REDIRECT_LISTEN 时必须同时设置 TLS_LISTEN / TLS_CERT_FILE / TLS_KEY_FILE",
|
||
));
|
||
};
|
||
let http_redirect_listen_addr = parse_socket_addr_config(
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTTP_REDIRECT_LISTEN",
|
||
http_redirect_listen_addr,
|
||
)?;
|
||
validate_distinct_listen_addr(
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTTP_REDIRECT_LISTEN",
|
||
http_redirect_listen_addr,
|
||
listen_addr,
|
||
self.api_upstream,
|
||
self.spacetime_upstream,
|
||
self.gitea_upstream,
|
||
)?;
|
||
if http_redirect_listen_addr == tls_listen_addr {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTTP_REDIRECT_LISTEN 不能与 TLS_LISTEN 相同",
|
||
));
|
||
}
|
||
self.validate_gitea_config(
|
||
listen_addr,
|
||
Some(tls_listen_addr),
|
||
Some(http_redirect_listen_addr),
|
||
)?;
|
||
validate_redirect_target_scheme(&self.http_redirect_target_scheme)?;
|
||
}
|
||
|
||
if self.max_api_body_bytes == 0 {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_MAX_API_BODY_BYTES 必须大于 0",
|
||
));
|
||
}
|
||
|
||
if self.gzip_level > 9 {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_GZIP_LEVEL 必须在 0..=9 之间",
|
||
));
|
||
}
|
||
if self.gzip_min_length_bytes == 0 {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_GZIP_MIN_LENGTH_BYTES 必须大于 0",
|
||
));
|
||
}
|
||
self.compression.validate()?;
|
||
self.static_cache.validate()?;
|
||
self.upstream_timeouts.validate()?;
|
||
|
||
if let Some(token) = self.shadow_probe_token.as_deref() {
|
||
validate_shadow_probe_token(token)?;
|
||
}
|
||
|
||
if self.trust_x_forwarded_for
|
||
&& !read_bool_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_TRUSTED_FRONT_PROXY_CONFIRMED",
|
||
false,
|
||
)
|
||
{
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"启用 GENARRATIVE_PINGORA_GATEWAY_TRUST_X_FORWARDED_FOR=true 前必须同时设置 GENARRATIVE_PINGORA_GATEWAY_TRUSTED_FRONT_PROXY_CONFIRMED=true,确认前置代理会清洗 X-Forwarded-For",
|
||
));
|
||
}
|
||
|
||
if let Some(access_log_file) = &self.access_log_file {
|
||
validate_access_log_file(access_log_file)?;
|
||
}
|
||
|
||
self.protection.validate()?;
|
||
validate_instance_protection_boundary(
|
||
self.instance_count,
|
||
self.protection.enabled,
|
||
self.shared_protection_confirmed,
|
||
)?;
|
||
Ok(())
|
||
}
|
||
|
||
fn validate_gitea_config(
|
||
&self,
|
||
listen_addr: SocketAddr,
|
||
tls_listen_addr: Option<SocketAddr>,
|
||
http_redirect_listen_addr: Option<SocketAddr>,
|
||
) -> io::Result<()> {
|
||
for host in &self.gitea_hosts {
|
||
validate_configured_gateway_host("GENARRATIVE_PINGORA_GATEWAY_GITEA_HOSTS", host)?;
|
||
}
|
||
|
||
match (self.gitea_hosts.is_empty(), self.gitea_upstream) {
|
||
(true, Some(_)) => {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"配置 GENARRATIVE_PINGORA_GATEWAY_GITEA_UPSTREAM 时必须同时设置 GENARRATIVE_PINGORA_GATEWAY_GITEA_HOSTS",
|
||
));
|
||
}
|
||
(false, None) => {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"配置 GENARRATIVE_PINGORA_GATEWAY_GITEA_HOSTS 时必须同时设置 GENARRATIVE_PINGORA_GATEWAY_GITEA_UPSTREAM",
|
||
));
|
||
}
|
||
_ => {}
|
||
}
|
||
|
||
let Some(gitea_upstream) = self.gitea_upstream else {
|
||
return Ok(());
|
||
};
|
||
if gitea_upstream == self.api_upstream || gitea_upstream == self.spacetime_upstream {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_GITEA_UPSTREAM 不能与 API 或 SpacetimeDB 上游地址相同",
|
||
));
|
||
}
|
||
if gitea_upstream == listen_addr
|
||
|| Some(gitea_upstream) == tls_listen_addr
|
||
|| Some(gitea_upstream) == http_redirect_listen_addr
|
||
{
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_GITEA_UPSTREAM 不能与 Pingora 监听地址相同",
|
||
));
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
fn matches_gitea_host(&self, host: Option<&str>) -> bool {
|
||
let Some(host) = host.and_then(normalize_gateway_host) else {
|
||
return false;
|
||
};
|
||
self.gitea_hosts.iter().any(|candidate| candidate == &host)
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Debug)]
|
||
struct CompressionConfig {
|
||
gzip: bool,
|
||
raw_algorithms: String,
|
||
}
|
||
|
||
impl CompressionConfig {
|
||
fn from_env() -> io::Result<Self> {
|
||
let raw_algorithms =
|
||
read_env_or_default("GENARRATIVE_PINGORA_GATEWAY_COMPRESSION_ALGORITHMS", "gzip");
|
||
Self::parse(raw_algorithms)
|
||
}
|
||
|
||
fn parse(raw_algorithms: String) -> io::Result<Self> {
|
||
let mut gzip = false;
|
||
|
||
for algorithm in raw_algorithms
|
||
.split(',')
|
||
.map(str::trim)
|
||
.filter(|value| !value.is_empty())
|
||
{
|
||
match algorithm.to_ascii_lowercase().as_str() {
|
||
"gzip" => gzip = true,
|
||
_ => {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!(
|
||
"GENARRATIVE_PINGORA_GATEWAY_COMPRESSION_ALGORITHMS 当前只支持 gzip,收到 {algorithm:?}"
|
||
),
|
||
));
|
||
}
|
||
}
|
||
}
|
||
|
||
Ok(Self {
|
||
gzip,
|
||
raw_algorithms,
|
||
})
|
||
}
|
||
|
||
fn validate(&self) -> io::Result<()> {
|
||
if !self.gzip && !self.raw_algorithms.trim().is_empty() {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_COMPRESSION_ALGORITHMS 未启用任何受支持算法",
|
||
));
|
||
}
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Debug)]
|
||
struct UpstreamTimeoutConfig {
|
||
connect_timeout_ms: u64,
|
||
default_read_timeout_seconds: u64,
|
||
api_read_timeout_seconds: u64,
|
||
long_read_timeout_seconds: u64,
|
||
write_timeout_seconds: u64,
|
||
}
|
||
|
||
impl UpstreamTimeoutConfig {
|
||
fn from_env() -> io::Result<Self> {
|
||
Ok(Self {
|
||
connect_timeout_ms: read_u64_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_CONNECT_TIMEOUT_MS",
|
||
DEFAULT_UPSTREAM_CONNECT_TIMEOUT_MS,
|
||
)?,
|
||
default_read_timeout_seconds: read_u64_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_DEFAULT_READ_TIMEOUT_SECONDS",
|
||
DEFAULT_UPSTREAM_DEFAULT_READ_TIMEOUT_SECONDS,
|
||
)?,
|
||
api_read_timeout_seconds: read_u64_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_API_READ_TIMEOUT_SECONDS",
|
||
DEFAULT_UPSTREAM_API_READ_TIMEOUT_SECONDS,
|
||
)?,
|
||
long_read_timeout_seconds: read_u64_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_LONG_READ_TIMEOUT_SECONDS",
|
||
DEFAULT_UPSTREAM_LONG_READ_TIMEOUT_SECONDS,
|
||
)?,
|
||
write_timeout_seconds: read_u64_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_WRITE_TIMEOUT_SECONDS",
|
||
DEFAULT_UPSTREAM_WRITE_TIMEOUT_SECONDS,
|
||
)?,
|
||
})
|
||
}
|
||
|
||
fn validate(&self) -> io::Result<()> {
|
||
for (name, value) in [
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_CONNECT_TIMEOUT_MS",
|
||
self.connect_timeout_ms,
|
||
),
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_DEFAULT_READ_TIMEOUT_SECONDS",
|
||
self.default_read_timeout_seconds,
|
||
),
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_API_READ_TIMEOUT_SECONDS",
|
||
self.api_read_timeout_seconds,
|
||
),
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_LONG_READ_TIMEOUT_SECONDS",
|
||
self.long_read_timeout_seconds,
|
||
),
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_UPSTREAM_WRITE_TIMEOUT_SECONDS",
|
||
self.write_timeout_seconds,
|
||
),
|
||
] {
|
||
if value == 0 {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{name} 必须大于 0"),
|
||
));
|
||
}
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
fn read_timeout_for_route(&self, _route: &RouteDecision, path: &str) -> Duration {
|
||
let seconds = if is_long_upstream_route(path) {
|
||
self.long_read_timeout_seconds
|
||
} else if is_generic_api_proxy_path(path) {
|
||
self.api_read_timeout_seconds
|
||
} else {
|
||
self.default_read_timeout_seconds
|
||
};
|
||
Duration::from_secs(seconds)
|
||
}
|
||
|
||
fn connect_timeout(&self) -> Duration {
|
||
Duration::from_millis(self.connect_timeout_ms)
|
||
}
|
||
|
||
fn write_timeout(&self) -> Duration {
|
||
Duration::from_secs(self.write_timeout_seconds)
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Debug)]
|
||
struct StaticCacheConfig {
|
||
html_cache_control: String,
|
||
asset_cache_control: String,
|
||
static_cache_control: String,
|
||
}
|
||
|
||
impl StaticCacheConfig {
|
||
fn from_env() -> Self {
|
||
Self {
|
||
html_cache_control: read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTML_CACHE_CONTROL",
|
||
DEFAULT_HTML_CACHE_CONTROL,
|
||
),
|
||
asset_cache_control: read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_ASSET_CACHE_CONTROL",
|
||
DEFAULT_ASSET_CACHE_CONTROL,
|
||
),
|
||
static_cache_control: read_env_or_default(
|
||
"GENARRATIVE_PINGORA_GATEWAY_STATIC_CACHE_CONTROL",
|
||
DEFAULT_STATIC_CACHE_CONTROL,
|
||
),
|
||
}
|
||
}
|
||
|
||
fn validate(&self) -> io::Result<()> {
|
||
for (name, value) in [
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTML_CACHE_CONTROL",
|
||
self.html_cache_control.as_str(),
|
||
),
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_ASSET_CACHE_CONTROL",
|
||
self.asset_cache_control.as_str(),
|
||
),
|
||
(
|
||
"GENARRATIVE_PINGORA_GATEWAY_STATIC_CACHE_CONTROL",
|
||
self.static_cache_control.as_str(),
|
||
),
|
||
] {
|
||
if value.contains('\n') || value.contains('\r') || value.contains('\0') {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{name} 不能包含换行或 NUL 字符"),
|
||
));
|
||
}
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
#[derive(Clone)]
|
||
struct GenarrativeGateway {
|
||
config: Arc<GatewayConfig>,
|
||
protection: Arc<GatewayProtection>,
|
||
mode: GatewayMode,
|
||
}
|
||
|
||
#[derive(Debug)]
|
||
struct RequestContext {
|
||
request_id: String,
|
||
route: RouteDecision,
|
||
request_body_bytes_seen: u64,
|
||
protection_class: Option<ProtectionClass>,
|
||
protection_client: String,
|
||
protection_key: Option<ProtectionKey>,
|
||
started_at: Instant,
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum GatewayMode {
|
||
Proxy,
|
||
HttpRedirect,
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum CompressionAlgorithm {
|
||
Gzip,
|
||
}
|
||
|
||
impl CompressionAlgorithm {
|
||
fn header_value(self) -> &'static str {
|
||
match self {
|
||
Self::Gzip => "gzip",
|
||
}
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Copy)]
|
||
struct CompressionRequestConfig {
|
||
enabled: bool,
|
||
gzip: bool,
|
||
min_length_bytes: u64,
|
||
gzip_level: u32,
|
||
}
|
||
|
||
struct GatewayResponseCompressionBuilder {
|
||
config: CompressionRequestConfig,
|
||
}
|
||
|
||
struct GatewayResponseCompression {
|
||
config: CompressionRequestConfig,
|
||
inner: ResponseCompressionCtx,
|
||
}
|
||
|
||
impl HttpModuleBuilder for GatewayResponseCompressionBuilder {
|
||
fn init(&self) -> Module {
|
||
Box::new(GatewayResponseCompression {
|
||
config: self.config,
|
||
inner: ResponseCompressionCtx::new(self.config.gzip_level, false, false),
|
||
})
|
||
}
|
||
|
||
fn order(&self) -> i16 {
|
||
i16::MIN / 2
|
||
}
|
||
}
|
||
|
||
#[async_trait]
|
||
impl HttpModule for GatewayResponseCompression {
|
||
fn as_any(&self) -> &dyn std::any::Any {
|
||
self
|
||
}
|
||
|
||
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
|
||
self
|
||
}
|
||
|
||
async fn request_header_filter(&mut self, req: &mut RequestHeader) -> PingoraResult<()> {
|
||
normalize_accept_encoding_for_gateway_compression(req, self.config)?;
|
||
self.inner.request_filter(req);
|
||
Ok(())
|
||
}
|
||
|
||
async fn response_header_filter(
|
||
&mut self,
|
||
resp: &mut ResponseHeader,
|
||
_end_of_stream: bool,
|
||
) -> PingoraResult<()> {
|
||
if !should_allow_gateway_compression_for_response(resp, self.config.min_length_bytes) {
|
||
self.inner.adjust_level(0);
|
||
}
|
||
self.inner.response_header_filter(resp, _end_of_stream);
|
||
Ok(())
|
||
}
|
||
|
||
fn response_body_filter(
|
||
&mut self,
|
||
body: &mut Option<Bytes>,
|
||
end_of_stream: bool,
|
||
) -> PingoraResult<()> {
|
||
if !self.inner.is_enabled() {
|
||
return Ok(());
|
||
}
|
||
if let Some(compressed) = self
|
||
.inner
|
||
.response_body_filter(body.as_ref(), end_of_stream)
|
||
{
|
||
*body = Some(compressed);
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
fn response_done_filter(&mut self) -> PingoraResult<Option<Bytes>> {
|
||
if !self.inner.is_enabled() {
|
||
return Ok(None);
|
||
}
|
||
Ok(self.inner.response_body_filter(None, true))
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum ProxyTarget {
|
||
Api,
|
||
Spacetime,
|
||
Gitea,
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
|
||
enum ProtectionClass {
|
||
AdminApi,
|
||
Api,
|
||
Spacetime,
|
||
}
|
||
|
||
impl ProtectionClass {
|
||
fn as_str(self) -> &'static str {
|
||
match self {
|
||
ProtectionClass::AdminApi => "admin_api",
|
||
ProtectionClass::Api => "api",
|
||
ProtectionClass::Spacetime => "spacetime",
|
||
}
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum RejectReason {
|
||
Concurrent,
|
||
Rate,
|
||
}
|
||
|
||
impl RejectReason {
|
||
fn code(self) -> &'static str {
|
||
match self {
|
||
RejectReason::Concurrent => "GATEWAY_CONCURRENCY_LIMITED",
|
||
RejectReason::Rate => "GATEWAY_RATE_LIMITED",
|
||
}
|
||
}
|
||
|
||
fn message(self) -> &'static str {
|
||
match self {
|
||
RejectReason::Concurrent => "服务繁忙,请稍后重试",
|
||
RejectReason::Rate => "请求过于频繁,请稍后重试",
|
||
}
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug)]
|
||
struct ProtectionClassConfig {
|
||
max_concurrent: u32,
|
||
rate_per_second: u32,
|
||
burst: u32,
|
||
}
|
||
|
||
#[derive(Clone, Debug)]
|
||
struct ProtectionConfig {
|
||
enabled: bool,
|
||
admin_api: ProtectionClassConfig,
|
||
api: ProtectionClassConfig,
|
||
spacetime: ProtectionClassConfig,
|
||
}
|
||
|
||
impl ProtectionConfig {
|
||
fn from_env() -> io::Result<Self> {
|
||
Ok(Self {
|
||
enabled: read_bool_env("GENARRATIVE_PINGORA_GATEWAY_PROTECTION_ENABLED", true),
|
||
admin_api: ProtectionClassConfig::from_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_ADMIN_API",
|
||
64,
|
||
30,
|
||
16,
|
||
)?,
|
||
api: ProtectionClassConfig::from_env("GENARRATIVE_PINGORA_GATEWAY_API", 64, 300, 64)?,
|
||
spacetime: ProtectionClassConfig::from_env(
|
||
"GENARRATIVE_PINGORA_GATEWAY_SPACETIME",
|
||
256,
|
||
1000,
|
||
256,
|
||
)?,
|
||
})
|
||
}
|
||
|
||
fn class_config(&self, class: ProtectionClass) -> ProtectionClassConfig {
|
||
match class {
|
||
ProtectionClass::AdminApi => self.admin_api,
|
||
ProtectionClass::Api => self.api,
|
||
ProtectionClass::Spacetime => self.spacetime,
|
||
}
|
||
}
|
||
|
||
fn validate(&self) -> io::Result<()> {
|
||
for (name, config) in [
|
||
("admin_api", self.admin_api),
|
||
("api", self.api),
|
||
("spacetime", self.spacetime),
|
||
] {
|
||
config.validate(name)?;
|
||
}
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
impl ProtectionClassConfig {
|
||
fn from_env(
|
||
prefix: &str,
|
||
max_concurrent: u32,
|
||
rate_per_second: u32,
|
||
burst: u32,
|
||
) -> io::Result<Self> {
|
||
Ok(Self {
|
||
max_concurrent: read_u32_env(&format!("{prefix}_MAX_CONCURRENT"), max_concurrent)?,
|
||
rate_per_second: read_u32_env(&format!("{prefix}_RATE_PER_SECOND"), rate_per_second)?,
|
||
burst: read_u32_env(&format!("{prefix}_BURST"), burst)?,
|
||
})
|
||
}
|
||
|
||
fn is_disabled(&self) -> bool {
|
||
self.max_concurrent == 0 && self.rate_per_second == 0
|
||
}
|
||
|
||
fn bucket_capacity(&self) -> u32 {
|
||
self.rate_per_second.saturating_add(self.burst).max(1)
|
||
}
|
||
|
||
fn validate(&self, name: &str) -> io::Result<()> {
|
||
if self.burst > 0 && self.rate_per_second == 0 {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!(
|
||
"{name} 配置了 burst 但 RATE_PER_SECOND=0;若要关闭 RPS,请同时把 BURST 设为 0"
|
||
),
|
||
));
|
||
}
|
||
Ok(())
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum StaticRoot {
|
||
Web,
|
||
Acme,
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum StaticMode {
|
||
Exact,
|
||
SpaFallback,
|
||
}
|
||
|
||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||
enum LocalResponse {
|
||
RedirectPermanent { location: &'static str },
|
||
HttpToHttpsRedirect,
|
||
ShadowProbe,
|
||
Static { root: StaticRoot, mode: StaticMode },
|
||
NotFound,
|
||
}
|
||
|
||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||
enum RouteDecision {
|
||
Proxy {
|
||
target: ProxyTarget,
|
||
body_limit: Option<u64>,
|
||
},
|
||
Local(LocalResponse),
|
||
}
|
||
|
||
impl RouteDecision {
|
||
fn applies_maintenance_gate(&self) -> bool {
|
||
matches!(
|
||
self,
|
||
RouteDecision::Proxy {
|
||
target: ProxyTarget::Api | ProxyTarget::Spacetime,
|
||
..
|
||
} | RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Web,
|
||
..
|
||
})
|
||
)
|
||
}
|
||
|
||
fn is_api_like(&self) -> bool {
|
||
matches!(
|
||
self,
|
||
RouteDecision::Proxy {
|
||
target: ProxyTarget::Api,
|
||
..
|
||
}
|
||
)
|
||
}
|
||
|
||
fn proxy_target(&self) -> Option<ProxyTarget> {
|
||
match self {
|
||
RouteDecision::Proxy { target, .. } => Some(*target),
|
||
RouteDecision::Local(_) => None,
|
||
}
|
||
}
|
||
|
||
fn body_limit(&self) -> Option<u64> {
|
||
match self {
|
||
RouteDecision::Proxy { body_limit, .. } => *body_limit,
|
||
RouteDecision::Local(_) => None,
|
||
}
|
||
}
|
||
}
|
||
|
||
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
|
||
struct ProtectionKey {
|
||
class: ProtectionClass,
|
||
client: String,
|
||
}
|
||
|
||
#[derive(Debug)]
|
||
struct ClientProtectionState {
|
||
in_flight: u32,
|
||
tokens: f64,
|
||
last_refill: Instant,
|
||
last_seen: Instant,
|
||
}
|
||
|
||
#[derive(Debug)]
|
||
struct ProtectionState {
|
||
clients: HashMap<ProtectionKey, ClientProtectionState>,
|
||
last_cleanup: Instant,
|
||
}
|
||
|
||
struct GatewayProtection {
|
||
config: ProtectionConfig,
|
||
state: Mutex<ProtectionState>,
|
||
}
|
||
|
||
impl GatewayProtection {
|
||
fn new(config: ProtectionConfig) -> Self {
|
||
Self {
|
||
config,
|
||
state: Mutex::new(ProtectionState {
|
||
clients: HashMap::new(),
|
||
last_cleanup: Instant::now(),
|
||
}),
|
||
}
|
||
}
|
||
|
||
fn try_acquire(
|
||
self: &Arc<Self>,
|
||
class: ProtectionClass,
|
||
client: String,
|
||
) -> Result<Option<ProtectionKey>, RejectReason> {
|
||
if !self.config.enabled {
|
||
return Ok(None);
|
||
}
|
||
|
||
let class_config = self.config.class_config(class);
|
||
if class_config.is_disabled() {
|
||
return Ok(None);
|
||
}
|
||
|
||
let now = Instant::now();
|
||
let key = ProtectionKey { class, client };
|
||
let mut state = self
|
||
.state
|
||
.lock()
|
||
.unwrap_or_else(|poisoned| poisoned.into_inner());
|
||
cleanup_protection_state(&mut state, now);
|
||
|
||
let entry = state
|
||
.clients
|
||
.entry(key.clone())
|
||
.or_insert_with(|| ClientProtectionState {
|
||
in_flight: 0,
|
||
tokens: class_config.bucket_capacity() as f64,
|
||
last_refill: now,
|
||
last_seen: now,
|
||
});
|
||
refill_tokens(entry, class_config, now);
|
||
entry.last_seen = now;
|
||
|
||
if class_config.max_concurrent > 0 && entry.in_flight >= class_config.max_concurrent {
|
||
return Err(RejectReason::Concurrent);
|
||
}
|
||
|
||
if class_config.rate_per_second > 0 && entry.tokens < 1.0 {
|
||
return Err(RejectReason::Rate);
|
||
}
|
||
|
||
if class_config.rate_per_second > 0 {
|
||
entry.tokens -= 1.0;
|
||
}
|
||
entry.in_flight = entry.in_flight.saturating_add(1);
|
||
|
||
Ok(Some(key))
|
||
}
|
||
|
||
fn release(&self, key: &ProtectionKey) {
|
||
let mut state = self
|
||
.state
|
||
.lock()
|
||
.unwrap_or_else(|poisoned| poisoned.into_inner());
|
||
if let Some(entry) = state.clients.get_mut(key) {
|
||
entry.in_flight = entry.in_flight.saturating_sub(1);
|
||
entry.last_seen = Instant::now();
|
||
}
|
||
}
|
||
}
|
||
|
||
fn refill_tokens(
|
||
entry: &mut ClientProtectionState,
|
||
class_config: ProtectionClassConfig,
|
||
now: Instant,
|
||
) {
|
||
if class_config.rate_per_second == 0 {
|
||
entry.last_refill = now;
|
||
return;
|
||
}
|
||
|
||
let elapsed = now.saturating_duration_since(entry.last_refill);
|
||
let refill = elapsed.as_secs_f64() * class_config.rate_per_second as f64;
|
||
if refill > 0.0 {
|
||
entry.tokens = (entry.tokens + refill).min(class_config.bucket_capacity() as f64);
|
||
entry.last_refill = now;
|
||
}
|
||
}
|
||
|
||
fn cleanup_protection_state(state: &mut ProtectionState, now: Instant) {
|
||
if now.saturating_duration_since(state.last_cleanup) < PROTECTION_CLEANUP_INTERVAL {
|
||
return;
|
||
}
|
||
|
||
state.clients.retain(|_, entry| {
|
||
entry.in_flight > 0
|
||
|| now.saturating_duration_since(entry.last_seen) <= PROTECTION_STATE_TTL
|
||
});
|
||
state.last_cleanup = now;
|
||
}
|
||
|
||
fn main() -> std::result::Result<(), Box<dyn std::error::Error>> {
|
||
let config = Arc::new(GatewayConfig::from_env()?);
|
||
init_tracing(
|
||
&config.log_filter,
|
||
OtelConfig {
|
||
enabled: config.otel_enabled,
|
||
},
|
||
)?;
|
||
|
||
info!(
|
||
listen = %config.listen_addr,
|
||
tls_listen = ?config.tls_listen_addr,
|
||
http_redirect_listen = ?config.http_redirect_listen_addr,
|
||
api_upstream = %config.api_upstream,
|
||
spacetime_upstream = %config.spacetime_upstream,
|
||
gitea_hosts = ?config.gitea_hosts,
|
||
gitea_upstream = ?config.gitea_upstream,
|
||
web_root = %config.web_root.display(),
|
||
"Pingora shadow gateway 启动"
|
||
);
|
||
|
||
let opt = Opt::parse_args();
|
||
let mut server = Server::new(Some(opt))?;
|
||
server.bootstrap();
|
||
|
||
let protection = Arc::new(GatewayProtection::new(config.protection.clone()));
|
||
let gateway = GenarrativeGateway {
|
||
config: config.clone(),
|
||
protection: protection.clone(),
|
||
mode: GatewayMode::Proxy,
|
||
};
|
||
let listen_addr = gateway.config.listen_addr.clone();
|
||
let mut proxy = http_proxy_service(&server.configuration, gateway);
|
||
proxy.add_tcp(&listen_addr);
|
||
if let Some(tls_listen_addr) = config.tls_listen_addr.as_deref() {
|
||
let cert_path = config
|
||
.tls_cert_file
|
||
.as_deref()
|
||
.expect("TLS cert file validated")
|
||
.to_string_lossy()
|
||
.into_owned();
|
||
let key_path = config
|
||
.tls_key_file
|
||
.as_deref()
|
||
.expect("TLS key file validated")
|
||
.to_string_lossy()
|
||
.into_owned();
|
||
let mut tls_settings = TlsSettings::intermediate(&cert_path, &key_path)?;
|
||
tls_settings.enable_h2();
|
||
proxy.add_tls_with_settings(tls_listen_addr, None, tls_settings);
|
||
}
|
||
server.add_service(proxy);
|
||
|
||
if let Some(http_redirect_listen_addr) = config.http_redirect_listen_addr.clone() {
|
||
let redirect_gateway = GenarrativeGateway {
|
||
config,
|
||
protection,
|
||
mode: GatewayMode::HttpRedirect,
|
||
};
|
||
let mut redirect_service = http_proxy_service(&server.configuration, redirect_gateway);
|
||
redirect_service.add_tcp(&http_redirect_listen_addr);
|
||
server.add_service(redirect_service);
|
||
}
|
||
|
||
server.run_forever();
|
||
}
|
||
|
||
#[async_trait]
|
||
impl ProxyHttp for GenarrativeGateway {
|
||
type CTX = RequestContext;
|
||
|
||
fn init_downstream_modules(&self, modules: &mut HttpModules) {
|
||
let gzip_level = if self.config.gzip_enabled {
|
||
self.config.gzip_level
|
||
} else {
|
||
0
|
||
};
|
||
modules.add_module(Box::new(GatewayResponseCompressionBuilder {
|
||
config: CompressionRequestConfig {
|
||
enabled: gzip_level > 0,
|
||
gzip: self.config.compression.gzip,
|
||
min_length_bytes: self.config.gzip_min_length_bytes,
|
||
gzip_level,
|
||
},
|
||
}));
|
||
}
|
||
|
||
fn new_ctx(&self) -> Self::CTX {
|
||
RequestContext {
|
||
request_id: Uuid::new_v4().to_string(),
|
||
route: RouteDecision::Local(LocalResponse::NotFound),
|
||
request_body_bytes_seen: 0,
|
||
protection_class: None,
|
||
protection_client: String::new(),
|
||
protection_key: None,
|
||
started_at: Instant::now(),
|
||
}
|
||
}
|
||
|
||
async fn request_filter(
|
||
&self,
|
||
session: &mut Session,
|
||
ctx: &mut Self::CTX,
|
||
) -> PingoraResult<bool>
|
||
where
|
||
Self::CTX: Send + Sync,
|
||
{
|
||
let path = session.req_header().uri.path();
|
||
let host = host_or_authority(session);
|
||
ctx.route = self.classify_request(host.as_deref(), path);
|
||
apply_configured_body_limit(&mut ctx.route, self.config.max_api_body_bytes);
|
||
ctx.request_id = resolve_request_id(session);
|
||
|
||
if self.is_maintenance_enabled() {
|
||
let internal_bypass =
|
||
allows_internal_maintenance_bypass(request_source_ip(session).as_ref());
|
||
let should_apply_maintenance = !internal_bypass
|
||
&& !is_maintenance_page_asset(path)
|
||
&& (is_admin_request_path(path) || ctx.route.applies_maintenance_gate());
|
||
if should_apply_maintenance {
|
||
respond_maintenance(
|
||
session,
|
||
ctx.route.is_api_like(),
|
||
&self.config.maintenance_page_file,
|
||
&self.config.web_root,
|
||
)
|
||
.await?;
|
||
return Ok(true);
|
||
}
|
||
}
|
||
|
||
if let RouteDecision::Proxy { body_limit, .. } = ctx.route
|
||
&& let Some(limit) = body_limit
|
||
&& let Some(content_length) = content_length(session)
|
||
&& content_length > limit
|
||
{
|
||
respond_json(session, 413, payload_too_large_body()).await?;
|
||
return Ok(true);
|
||
}
|
||
|
||
if let Some(protection_class) = protection_class_for_route(&ctx.route, path) {
|
||
let protection_client =
|
||
protection_client_id(session, self.config.trust_x_forwarded_for);
|
||
match self
|
||
.protection
|
||
.try_acquire(protection_class, protection_client.clone())
|
||
{
|
||
Ok(key) => {
|
||
ctx.protection_class = Some(protection_class);
|
||
ctx.protection_client = protection_client;
|
||
ctx.protection_key = key;
|
||
}
|
||
Err(reason) => {
|
||
ctx.protection_class = Some(protection_class);
|
||
ctx.protection_client = protection_client;
|
||
respond_too_many_requests(session, reason).await?;
|
||
return Ok(true);
|
||
}
|
||
}
|
||
}
|
||
|
||
match &ctx.route {
|
||
RouteDecision::Proxy { .. } => Ok(false),
|
||
RouteDecision::Local(LocalResponse::RedirectPermanent { location }) => {
|
||
respond_redirect(session, location).await?;
|
||
Ok(true)
|
||
}
|
||
RouteDecision::Local(LocalResponse::HttpToHttpsRedirect) => {
|
||
respond_http_to_https_redirect(
|
||
session,
|
||
self.config.http_redirect_target_scheme.as_str(),
|
||
)
|
||
.await?;
|
||
Ok(true)
|
||
}
|
||
RouteDecision::Local(LocalResponse::ShadowProbe) => {
|
||
if !self.is_shadow_probe_authorized(session) {
|
||
respond_not_found(session).await?;
|
||
return Ok(true);
|
||
}
|
||
|
||
respond_shadow_probe(session, self.is_maintenance_enabled()).await?;
|
||
Ok(true)
|
||
}
|
||
RouteDecision::Local(LocalResponse::Static { root, mode }) => {
|
||
let root_path = match root {
|
||
StaticRoot::Web => &self.config.web_root,
|
||
StaticRoot::Acme => &self.config.acme_root,
|
||
};
|
||
serve_static(session, root_path, *root, *mode, &self.config.static_cache).await?;
|
||
Ok(true)
|
||
}
|
||
RouteDecision::Local(LocalResponse::NotFound) => {
|
||
respond_not_found(session).await?;
|
||
Ok(true)
|
||
}
|
||
}
|
||
}
|
||
|
||
async fn request_body_filter(
|
||
&self,
|
||
session: &mut Session,
|
||
body: &mut Option<Bytes>,
|
||
_end_of_stream: bool,
|
||
ctx: &mut Self::CTX,
|
||
) -> PingoraResult<()>
|
||
where
|
||
Self::CTX: Send + Sync,
|
||
{
|
||
let Some(limit) = ctx.route.body_limit() else {
|
||
return Ok(());
|
||
};
|
||
|
||
if let Some(body) = body.as_ref() {
|
||
ctx.request_body_bytes_seen = ctx
|
||
.request_body_bytes_seen
|
||
.saturating_add(body.len() as u64);
|
||
}
|
||
|
||
if ctx.request_body_bytes_seen > limit {
|
||
respond_json(session, 413, payload_too_large_body()).await?;
|
||
return Error::e_explain(HTTPStatus(413), PAYLOAD_TOO_LARGE_CONTEXT);
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
async fn upstream_peer(
|
||
&self,
|
||
_session: &mut Session,
|
||
ctx: &mut Self::CTX,
|
||
) -> PingoraResult<Box<HttpPeer>> {
|
||
let upstream = match ctx.route.proxy_target().unwrap_or(ProxyTarget::Api) {
|
||
ProxyTarget::Api => self.config.api_upstream,
|
||
ProxyTarget::Spacetime => self.config.spacetime_upstream,
|
||
ProxyTarget::Gitea => self
|
||
.config
|
||
.gitea_upstream
|
||
.expect("Gitea route requires configured upstream"),
|
||
};
|
||
|
||
let mut peer = HttpPeer::new(upstream, false, String::new());
|
||
peer.options.connection_timeout = Some(self.config.upstream_timeouts.connect_timeout());
|
||
peer.options.read_timeout = Some(
|
||
self.config
|
||
.upstream_timeouts
|
||
.read_timeout_for_route(&ctx.route, _session.req_header().uri.path()),
|
||
);
|
||
peer.options.write_timeout = Some(self.config.upstream_timeouts.write_timeout());
|
||
|
||
Ok(Box::new(peer))
|
||
}
|
||
|
||
async fn upstream_request_filter(
|
||
&self,
|
||
session: &mut Session,
|
||
upstream_request: &mut pingora_http::RequestHeader,
|
||
ctx: &mut Self::CTX,
|
||
) -> PingoraResult<()>
|
||
where
|
||
Self::CTX: Send + Sync,
|
||
{
|
||
upstream_request.insert_header("X-Request-Id", ctx.request_id.as_str())?;
|
||
upstream_request
|
||
.insert_header("X-Forwarded-Proto", self.config.forwarded_proto.as_str())?;
|
||
|
||
if let Some(host) = host_or_authority(session) {
|
||
upstream_request.insert_header("Host", host.clone())?;
|
||
upstream_request.insert_header("X-Forwarded-Host", host)?;
|
||
}
|
||
|
||
if let Some(client_ip) = client_ip(session) {
|
||
upstream_request.insert_header("X-Real-IP", client_ip.as_str())?;
|
||
let forwarded_for = append_forwarded_for(session, &client_ip);
|
||
upstream_request.insert_header("X-Forwarded-For", forwarded_for.as_str())?;
|
||
}
|
||
|
||
// 中文注释:SpacetimeDB 订阅走 WebSocket Upgrade;普通 API 连接头保持干净,贴近当前 Nginx 模板。
|
||
if ctx.route.proxy_target() == Some(ProxyTarget::Spacetime) && is_upgrade_request(session) {
|
||
upstream_request.insert_header("Connection", "Upgrade")?;
|
||
} else {
|
||
upstream_request.remove_header("connection");
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
async fn response_filter(
|
||
&self,
|
||
_session: &mut Session,
|
||
upstream_response: &mut ResponseHeader,
|
||
ctx: &mut Self::CTX,
|
||
) -> PingoraResult<()>
|
||
where
|
||
Self::CTX: Send + Sync,
|
||
{
|
||
upstream_response.insert_header("X-Genarrative-Gateway", "pingora-shadow")?;
|
||
if should_disable_accel_buffering(&ctx.route) {
|
||
upstream_response.insert_header("X-Accel-Buffering", "no")?;
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
async fn logging(
|
||
&self,
|
||
session: &mut Session,
|
||
error: Option<&pingora::Error>,
|
||
ctx: &mut Self::CTX,
|
||
) where
|
||
Self::CTX: Send + Sync,
|
||
{
|
||
let status = session
|
||
.response_written()
|
||
.map_or(0, |response| response.status.as_u16());
|
||
let elapsed_ms = ctx.started_at.elapsed().as_millis();
|
||
let method = session.req_header().method.as_str();
|
||
let path = session.req_header().uri.path();
|
||
let uri = session.req_header().uri.to_string();
|
||
let host = header_value(session, "host").unwrap_or_default();
|
||
let client_ip = client_ip(session).unwrap_or_default();
|
||
let content_length = content_length(session);
|
||
let protection_class = ctx.protection_class.map(|class| class.as_str());
|
||
let protection_client = ctx.protection_client.as_str();
|
||
let proxy_target = ctx
|
||
.route
|
||
.proxy_target()
|
||
.map(|target| format!("{target:?}"))
|
||
.unwrap_or_else(|| "Local".to_string());
|
||
let upstream = match ctx.route.proxy_target() {
|
||
Some(ProxyTarget::Api) => self.config.api_upstream.to_string(),
|
||
Some(ProxyTarget::Spacetime) => self.config.spacetime_upstream.to_string(),
|
||
Some(ProxyTarget::Gitea) => self
|
||
.config
|
||
.gitea_upstream
|
||
.map(|upstream| upstream.to_string())
|
||
.unwrap_or_default(),
|
||
None => String::new(),
|
||
};
|
||
let route = format!("{:?}", ctx.route);
|
||
if let Some(key) = ctx.protection_key.take() {
|
||
self.protection.release(&key);
|
||
}
|
||
|
||
write_access_log(
|
||
self.config.access_log_file.as_deref(),
|
||
AccessLogRecord {
|
||
request_id: &ctx.request_id,
|
||
method,
|
||
path,
|
||
uri: &uri,
|
||
host: &host,
|
||
client_ip: &client_ip,
|
||
status,
|
||
route: &route,
|
||
proxy_target: &proxy_target,
|
||
upstream: &upstream,
|
||
content_length,
|
||
body_bytes_seen: ctx.request_body_bytes_seen,
|
||
protection_class,
|
||
protection_client,
|
||
elapsed_ms,
|
||
error: error.map(ToString::to_string),
|
||
},
|
||
);
|
||
|
||
if let Some(error) = error {
|
||
if is_payload_too_large_error(error) {
|
||
warn!(
|
||
request_id = %ctx.request_id,
|
||
method,
|
||
path,
|
||
uri = %uri,
|
||
host = %host,
|
||
client_ip = %client_ip,
|
||
status,
|
||
route = %route,
|
||
proxy_target = %proxy_target,
|
||
upstream = %upstream,
|
||
content_length = ?content_length,
|
||
body_bytes_seen = ctx.request_body_bytes_seen,
|
||
protection_class = ?protection_class,
|
||
protection_client,
|
||
elapsed_ms,
|
||
%error,
|
||
"Pingora gateway request rejected"
|
||
);
|
||
} else {
|
||
error!(
|
||
request_id = %ctx.request_id,
|
||
method,
|
||
path,
|
||
uri = %uri,
|
||
host = %host,
|
||
client_ip = %client_ip,
|
||
status,
|
||
route = %route,
|
||
proxy_target = %proxy_target,
|
||
upstream = %upstream,
|
||
content_length = ?content_length,
|
||
body_bytes_seen = ctx.request_body_bytes_seen,
|
||
protection_class = ?protection_class,
|
||
protection_client,
|
||
elapsed_ms,
|
||
%error,
|
||
"Pingora gateway request failed"
|
||
);
|
||
}
|
||
} else {
|
||
info!(
|
||
request_id = %ctx.request_id,
|
||
method,
|
||
path,
|
||
uri = %uri,
|
||
host = %host,
|
||
client_ip = %client_ip,
|
||
status,
|
||
route = %route,
|
||
proxy_target = %proxy_target,
|
||
upstream = %upstream,
|
||
content_length = ?content_length,
|
||
body_bytes_seen = ctx.request_body_bytes_seen,
|
||
protection_class = ?protection_class,
|
||
protection_client,
|
||
elapsed_ms,
|
||
"Pingora gateway request completed"
|
||
);
|
||
}
|
||
}
|
||
|
||
fn suppress_error_log(
|
||
&self,
|
||
_session: &Session,
|
||
_ctx: &Self::CTX,
|
||
error: &pingora::Error,
|
||
) -> bool {
|
||
is_payload_too_large_error(error)
|
||
}
|
||
|
||
async fn fail_to_proxy(
|
||
&self,
|
||
session: &mut Session,
|
||
error: &pingora::Error,
|
||
ctx: &mut Self::CTX,
|
||
) -> FailToProxy
|
||
where
|
||
Self::CTX: Send + Sync,
|
||
{
|
||
if is_payload_too_large_error(error) {
|
||
if session.response_written().is_none() {
|
||
respond_json(session, 413, payload_too_large_body())
|
||
.await
|
||
.unwrap_or_else(|error| {
|
||
warn!(%error, "请求体超限响应写入失败");
|
||
});
|
||
}
|
||
|
||
return FailToProxy {
|
||
error_code: 413,
|
||
can_reuse_downstream: false,
|
||
};
|
||
}
|
||
|
||
let status = match error.etype() {
|
||
ConnectTimedout | ReadTimedout | WriteTimedout => 504,
|
||
HTTPStatus(code) => *code,
|
||
_ => match error.esource() {
|
||
ErrorSource::Upstream => 502,
|
||
ErrorSource::Downstream => match error.etype() {
|
||
WriteError | ReadError | ConnectionClosed => 0,
|
||
_ => 400,
|
||
},
|
||
ErrorSource::Internal | ErrorSource::Unset => 500,
|
||
},
|
||
};
|
||
if status > 0 {
|
||
let response_result = if ctx.route.proxy_target().is_some() {
|
||
respond_gateway_proxy_error(session, status).await
|
||
} else {
|
||
session.respond_error(status).await
|
||
};
|
||
response_result.unwrap_or_else(|error| {
|
||
warn!(%error, "Pingora 错误响应写入失败");
|
||
});
|
||
}
|
||
|
||
FailToProxy {
|
||
error_code: status,
|
||
can_reuse_downstream: false,
|
||
}
|
||
}
|
||
}
|
||
|
||
struct AccessLogRecord<'a> {
|
||
request_id: &'a str,
|
||
method: &'a str,
|
||
path: &'a str,
|
||
uri: &'a str,
|
||
host: &'a str,
|
||
client_ip: &'a str,
|
||
status: u16,
|
||
route: &'a str,
|
||
proxy_target: &'a str,
|
||
upstream: &'a str,
|
||
content_length: Option<u64>,
|
||
body_bytes_seen: u64,
|
||
protection_class: Option<&'static str>,
|
||
protection_client: &'a str,
|
||
elapsed_ms: u128,
|
||
error: Option<String>,
|
||
}
|
||
|
||
fn write_access_log(path: Option<&Path>, record: AccessLogRecord<'_>) {
|
||
let Some(path) = path else {
|
||
return;
|
||
};
|
||
|
||
if let Some(parent) = path.parent()
|
||
&& let Err(error) = std::fs::create_dir_all(parent)
|
||
{
|
||
warn!(
|
||
file = %path.display(),
|
||
%error,
|
||
"Pingora access log 目录创建失败"
|
||
);
|
||
return;
|
||
}
|
||
|
||
let line = format_access_log_line(&record);
|
||
match OpenOptions::new().create(true).append(true).open(path) {
|
||
Ok(mut file) => {
|
||
if let Err(error) = file.write_all(line.as_bytes()) {
|
||
warn!(
|
||
file = %path.display(),
|
||
%error,
|
||
"Pingora access log 写入失败"
|
||
);
|
||
}
|
||
}
|
||
Err(error) => {
|
||
warn!(
|
||
file = %path.display(),
|
||
%error,
|
||
"Pingora access log 打开失败"
|
||
);
|
||
}
|
||
}
|
||
}
|
||
|
||
fn format_access_log_line(record: &AccessLogRecord<'_>) -> String {
|
||
format!(
|
||
concat!(
|
||
"request_id={}\tmethod={}\tpath={}\turi={}\thost={}\tclient_ip={}\tstatus={}\t",
|
||
"route={}\tproxy_target={}\tupstream={}\tcontent_length={}\tbody_bytes_seen={}\t",
|
||
"protection_class={}\tprotection_client={}\telapsed_ms={}\terror={}\n"
|
||
),
|
||
escape_access_log_value(record.request_id),
|
||
escape_access_log_value(record.method),
|
||
escape_access_log_value(record.path),
|
||
escape_access_log_value(record.uri),
|
||
escape_access_log_value(record.host),
|
||
escape_access_log_value(record.client_ip),
|
||
record.status,
|
||
escape_access_log_value(record.route),
|
||
escape_access_log_value(record.proxy_target),
|
||
escape_access_log_value(record.upstream),
|
||
record
|
||
.content_length
|
||
.map(|value| value.to_string())
|
||
.unwrap_or_else(|| "-".to_string()),
|
||
record.body_bytes_seen,
|
||
record.protection_class.unwrap_or("-"),
|
||
escape_access_log_value(record.protection_client),
|
||
record.elapsed_ms,
|
||
record
|
||
.error
|
||
.as_deref()
|
||
.map(escape_access_log_value)
|
||
.unwrap_or_else(|| "-".to_string())
|
||
)
|
||
}
|
||
|
||
fn escape_access_log_value(value: &str) -> String {
|
||
value
|
||
.replace('\\', "\\\\")
|
||
.replace('\t', "\\t")
|
||
.replace('\n', "\\n")
|
||
.replace('\r', "\\r")
|
||
}
|
||
|
||
fn should_disable_accel_buffering(route: &RouteDecision) -> bool {
|
||
matches!(
|
||
route.proxy_target(),
|
||
Some(ProxyTarget::Api | ProxyTarget::Gitea)
|
||
)
|
||
}
|
||
|
||
fn is_long_upstream_route(path: &str) -> bool {
|
||
is_spacetime_subscribe_path(path)
|
||
}
|
||
|
||
fn is_generic_api_proxy_path(path: &str) -> bool {
|
||
path == "/api" || path.starts_with("/api/")
|
||
}
|
||
|
||
fn protection_class_for_route(route: &RouteDecision, path: &str) -> Option<ProtectionClass> {
|
||
match route {
|
||
RouteDecision::Proxy {
|
||
target: ProxyTarget::Api,
|
||
..
|
||
} if path.starts_with("/admin/api/") => Some(ProtectionClass::AdminApi),
|
||
RouteDecision::Proxy {
|
||
target: ProxyTarget::Api,
|
||
..
|
||
} => Some(ProtectionClass::Api),
|
||
RouteDecision::Proxy {
|
||
target: ProxyTarget::Spacetime,
|
||
..
|
||
} => Some(ProtectionClass::Spacetime),
|
||
RouteDecision::Proxy {
|
||
target: ProxyTarget::Gitea,
|
||
..
|
||
} => None,
|
||
RouteDecision::Local(_) => None,
|
||
}
|
||
}
|
||
|
||
fn validate_shadow_probe_token(token: &str) -> io::Result<()> {
|
||
if token == "__GENARRATIVE_PINGORA_PROBE_TOKEN__"
|
||
|| token.eq_ignore_ascii_case("changeme")
|
||
|| token.len() < 16
|
||
{
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_PROBE_TOKEN 必须使用 16 字符以上的非占位 token",
|
||
));
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
fn validate_access_log_file(path: &Path) -> io::Result<()> {
|
||
if path.as_os_str().is_empty() {
|
||
return Ok(());
|
||
}
|
||
|
||
if path.exists() && path.is_dir() {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!(
|
||
"GENARRATIVE_PINGORA_GATEWAY_ACCESS_LOG_FILE 不能指向目录:{}",
|
||
path.display()
|
||
),
|
||
));
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
fn validate_instance_protection_boundary(
|
||
instance_count: u32,
|
||
protection_enabled: bool,
|
||
shared_protection_confirmed: bool,
|
||
) -> io::Result<()> {
|
||
if instance_count == 0 {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_INSTANCE_COUNT 必须大于 0",
|
||
));
|
||
}
|
||
|
||
if protection_enabled && instance_count > 1 && !shared_protection_confirmed {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"Pingora 接流保护当前默认是进程内状态;GENARRATIVE_PINGORA_GATEWAY_INSTANCE_COUNT>1 时必须设置 GENARRATIVE_PINGORA_GATEWAY_SHARED_PROTECTION_CONFIRMED=true,或关闭网关保护并由前置层承担",
|
||
));
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
fn parse_socket_addr_config(key: &str, value: &str) -> io::Result<SocketAddr> {
|
||
value.parse().map_err(|error| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 不是有效 socket 地址 {value:?}:{error}"),
|
||
)
|
||
})
|
||
}
|
||
|
||
fn validate_distinct_listen_addr(
|
||
key: &str,
|
||
listen_addr: SocketAddr,
|
||
primary_listen_addr: SocketAddr,
|
||
api_upstream: SocketAddr,
|
||
spacetime_upstream: SocketAddr,
|
||
gitea_upstream: Option<SocketAddr>,
|
||
) -> io::Result<()> {
|
||
if listen_addr == primary_listen_addr {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 不能与 GENARRATIVE_PINGORA_GATEWAY_LISTEN 相同"),
|
||
));
|
||
}
|
||
|
||
if listen_addr == api_upstream
|
||
|| listen_addr == spacetime_upstream
|
||
|| Some(listen_addr) == gitea_upstream
|
||
{
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 不能与上游地址相同"),
|
||
));
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
fn validate_tls_file_config(
|
||
cert_file: &Option<PathBuf>,
|
||
key_file: &Option<PathBuf>,
|
||
) -> io::Result<()> {
|
||
let cert_file = cert_file.as_ref().ok_or_else(|| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"配置 GENARRATIVE_PINGORA_GATEWAY_TLS_LISTEN 时必须设置 GENARRATIVE_PINGORA_GATEWAY_TLS_CERT_FILE",
|
||
)
|
||
})?;
|
||
let key_file = key_file.as_ref().ok_or_else(|| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"配置 GENARRATIVE_PINGORA_GATEWAY_TLS_LISTEN 时必须设置 GENARRATIVE_PINGORA_GATEWAY_TLS_KEY_FILE",
|
||
)
|
||
})?;
|
||
|
||
validate_existing_file(
|
||
"GENARRATIVE_PINGORA_GATEWAY_TLS_CERT_FILE",
|
||
cert_file.as_path(),
|
||
)?;
|
||
validate_existing_file(
|
||
"GENARRATIVE_PINGORA_GATEWAY_TLS_KEY_FILE",
|
||
key_file.as_path(),
|
||
)?;
|
||
Ok(())
|
||
}
|
||
|
||
fn validate_existing_file(key: &str, path: &Path) -> io::Result<()> {
|
||
let metadata = std::fs::metadata(path).map_err(|error| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 无法读取 {}:{error}", path.display()),
|
||
)
|
||
})?;
|
||
if !metadata.is_file() {
|
||
return Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 必须指向文件:{}", path.display()),
|
||
));
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
fn validate_redirect_target_scheme(scheme: &str) -> io::Result<()> {
|
||
if scheme == "https" {
|
||
return Ok(());
|
||
}
|
||
|
||
Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
"GENARRATIVE_PINGORA_GATEWAY_HTTP_REDIRECT_TARGET_SCHEME 当前只允许 https",
|
||
))
|
||
}
|
||
|
||
impl GenarrativeGateway {
|
||
fn is_maintenance_enabled(&self) -> bool {
|
||
self.config.maintenance_file.exists()
|
||
}
|
||
|
||
fn is_shadow_probe_authorized(&self, session: &Session) -> bool {
|
||
let Some(expected) = self.config.shadow_probe_token.as_deref() else {
|
||
return false;
|
||
};
|
||
|
||
session
|
||
.req_header()
|
||
.headers
|
||
.get(SHADOW_PROBE_HEADER)
|
||
.and_then(|value| value.to_str().ok())
|
||
.is_some_and(|actual| actual == expected)
|
||
}
|
||
|
||
fn classify_request(&self, host: Option<&str>, path: &str) -> RouteDecision {
|
||
match self.mode {
|
||
GatewayMode::Proxy if self.config.matches_gitea_host(host) => RouteDecision::Proxy {
|
||
target: ProxyTarget::Gitea,
|
||
body_limit: None,
|
||
},
|
||
GatewayMode::Proxy => classify_path(path),
|
||
GatewayMode::HttpRedirect => classify_http_redirect_path(path),
|
||
}
|
||
}
|
||
}
|
||
|
||
fn classify_path(path: &str) -> RouteDecision {
|
||
if path == SHADOW_PROBE_PATH {
|
||
return RouteDecision::Local(LocalResponse::ShadowProbe);
|
||
}
|
||
|
||
if path == "/admin" {
|
||
return RouteDecision::Local(LocalResponse::RedirectPermanent {
|
||
location: "/admin/",
|
||
});
|
||
}
|
||
|
||
if path.starts_with("/.well-known/acme-challenge/") {
|
||
return RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Acme,
|
||
mode: StaticMode::Exact,
|
||
});
|
||
}
|
||
|
||
if path.starts_with("/admin/api/") {
|
||
return RouteDecision::Proxy {
|
||
target: ProxyTarget::Api,
|
||
body_limit: None,
|
||
};
|
||
}
|
||
|
||
if path == "/api" || path.starts_with("/api/") {
|
||
return RouteDecision::Proxy {
|
||
target: ProxyTarget::Api,
|
||
body_limit: Some(DEFAULT_MAX_API_BODY_BYTES),
|
||
};
|
||
}
|
||
|
||
// 中文注释:公网只转发前端 SDK 必需的 SpacetimeDB subscribe / identity 路由,其它 /v1 继续关闭。
|
||
if is_spacetime_subscribe_path(path) || path.starts_with("/v1/identity") {
|
||
return RouteDecision::Proxy {
|
||
target: ProxyTarget::Spacetime,
|
||
body_limit: None,
|
||
};
|
||
}
|
||
|
||
if path.starts_with("/v1/")
|
||
|| path.starts_with("/generated-")
|
||
|| path.starts_with("/healthz")
|
||
|| path.starts_with("/readyz")
|
||
{
|
||
return RouteDecision::Local(LocalResponse::NotFound);
|
||
}
|
||
|
||
if path.starts_with("/admin/assets/") || path.starts_with("/assets/") {
|
||
return RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Web,
|
||
mode: StaticMode::Exact,
|
||
});
|
||
}
|
||
|
||
if path.starts_with("/admin/") {
|
||
return RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Web,
|
||
mode: StaticMode::SpaFallback,
|
||
});
|
||
}
|
||
|
||
if is_main_spa_path(path) {
|
||
return RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Web,
|
||
mode: StaticMode::SpaFallback,
|
||
});
|
||
}
|
||
|
||
RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Web,
|
||
mode: StaticMode::Exact,
|
||
})
|
||
}
|
||
|
||
fn is_main_spa_path(path: &str) -> bool {
|
||
let normalized = if path.len() > 1 {
|
||
path.strip_suffix('/').unwrap_or(path)
|
||
} else {
|
||
path
|
||
};
|
||
|
||
MAIN_SPA_PATHS
|
||
.iter()
|
||
.any(|candidate| normalized.eq_ignore_ascii_case(candidate))
|
||
}
|
||
|
||
fn is_maintenance_page_asset(path: &str) -> bool {
|
||
matches!(
|
||
path,
|
||
"/branding/taonier-maintenance-page.png" | "/branding/taonier-product-ip.png"
|
||
)
|
||
}
|
||
|
||
fn classify_http_redirect_path(path: &str) -> RouteDecision {
|
||
if path.starts_with("/.well-known/acme-challenge/") {
|
||
return RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Acme,
|
||
mode: StaticMode::Exact,
|
||
});
|
||
}
|
||
|
||
RouteDecision::Local(LocalResponse::HttpToHttpsRedirect)
|
||
}
|
||
|
||
fn apply_configured_body_limit(route: &mut RouteDecision, max_api_body_bytes: u64) {
|
||
if let RouteDecision::Proxy {
|
||
target: ProxyTarget::Api,
|
||
body_limit: Some(body_limit),
|
||
} = route
|
||
{
|
||
*body_limit = max_api_body_bytes;
|
||
}
|
||
}
|
||
|
||
fn payload_too_large_body() -> &'static str {
|
||
r#"{"ok":false,"error":{"code":"PAYLOAD_TOO_LARGE","message":"请求体过大"}}"#
|
||
}
|
||
|
||
fn gateway_proxy_error_body(status: u16) -> &'static str {
|
||
match status {
|
||
502 => {
|
||
r#"{"ok":false,"error":{"code":"GATEWAY_UPSTREAM_ERROR","message":"上游服务不可用"}}"#
|
||
}
|
||
504 => {
|
||
r#"{"ok":false,"error":{"code":"GATEWAY_UPSTREAM_TIMEOUT","message":"上游服务请求超时"}}"#
|
||
}
|
||
_ => r#"{"ok":false,"error":{"code":"GATEWAY_PROXY_ERROR","message":"网关代理失败"}}"#,
|
||
}
|
||
}
|
||
|
||
fn is_payload_too_large_error(error: &pingora::Error) -> bool {
|
||
matches!(error.etype(), HTTPStatus(413))
|
||
&& error
|
||
.context
|
||
.as_ref()
|
||
.is_some_and(|context| context.as_str() == PAYLOAD_TOO_LARGE_CONTEXT)
|
||
}
|
||
|
||
fn is_spacetime_subscribe_path(path: &str) -> bool {
|
||
let parts: Vec<_> = path.trim_matches('/').split('/').collect();
|
||
matches!(parts.as_slice(), ["v1", "database", _, "subscribe"])
|
||
}
|
||
|
||
struct StaticCandidate {
|
||
path: PathBuf,
|
||
cache_kind: StaticCacheKind,
|
||
}
|
||
|
||
struct StaticResponseMetadata {
|
||
len: u64,
|
||
etag: String,
|
||
last_modified: Option<SystemTime>,
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum StaticRangeDecision {
|
||
Full,
|
||
Partial { start: u64, end: u64 },
|
||
Unsatisfiable,
|
||
}
|
||
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||
enum StaticCacheKind {
|
||
Html,
|
||
FingerprintedAsset,
|
||
Other,
|
||
}
|
||
|
||
async fn serve_static(
|
||
session: &mut Session,
|
||
root: &Path,
|
||
root_kind: StaticRoot,
|
||
mode: StaticMode,
|
||
cache_config: &StaticCacheConfig,
|
||
) -> PingoraResult<()> {
|
||
let path = session.req_header().uri.path();
|
||
let Some(candidate) = resolve_static_candidate(root, root_kind, path, mode).await else {
|
||
match root_kind {
|
||
StaticRoot::Web => respond_page_not_found(session, root).await?,
|
||
StaticRoot::Acme => respond_not_found(session).await?,
|
||
}
|
||
return Ok(());
|
||
};
|
||
if !is_static_read_method(&session.req_header().method) {
|
||
respond_static_method_not_allowed(session).await?;
|
||
return Ok(());
|
||
}
|
||
|
||
let metadata = match static_response_metadata(&candidate.path).await {
|
||
Ok(metadata) => metadata,
|
||
Err(error) => {
|
||
warn!(
|
||
path = %path,
|
||
file = %candidate.path.display(),
|
||
%error,
|
||
"静态文件 metadata 读取失败"
|
||
);
|
||
respond_not_found(session).await?;
|
||
return Ok(());
|
||
}
|
||
};
|
||
let content_type = mime_guess::from_path(&candidate.path)
|
||
.first_or_octet_stream()
|
||
.essence_str()
|
||
.to_string();
|
||
let cache_control = cache_control_for_static_candidate(&candidate, cache_config);
|
||
let last_modified = metadata.last_modified.map(fmt_http_date);
|
||
let mut headers = vec![
|
||
("Cache-Control".to_string(), cache_control.to_string()),
|
||
("ETag".to_string(), metadata.etag.clone()),
|
||
("Accept-Ranges".to_string(), "bytes".to_string()),
|
||
];
|
||
if let Some(last_modified) = last_modified.as_deref() {
|
||
headers.push(("Last-Modified".to_string(), last_modified.to_string()));
|
||
}
|
||
|
||
if static_not_modified(session, &metadata) {
|
||
return respond_empty_with_headers(session, 304, Some(&content_type), &headers).await;
|
||
}
|
||
|
||
match static_range_decision(session, &metadata) {
|
||
StaticRangeDecision::Full => {}
|
||
StaticRangeDecision::Partial { start, end } => {
|
||
let range_len = end - start + 1;
|
||
let mut range_headers = headers.clone();
|
||
range_headers.push((
|
||
"Content-Range".to_string(),
|
||
format!("bytes {start}-{end}/{}", metadata.len),
|
||
));
|
||
if session.req_header().method == Method::HEAD {
|
||
return respond_head_with_headers(
|
||
session,
|
||
206,
|
||
Some(&content_type),
|
||
&range_headers,
|
||
range_len,
|
||
)
|
||
.await;
|
||
}
|
||
|
||
let body = match read_static_range(&candidate.path, start, range_len).await {
|
||
Ok(bytes) => bytes,
|
||
Err(error) => {
|
||
warn!(
|
||
path = %path,
|
||
file = %candidate.path.display(),
|
||
%error,
|
||
"静态文件 Range 读取失败"
|
||
);
|
||
respond_not_found(session).await?;
|
||
return Ok(());
|
||
}
|
||
};
|
||
|
||
return respond_bytes_with_headers(
|
||
session,
|
||
206,
|
||
Some(&content_type),
|
||
&range_headers,
|
||
body,
|
||
)
|
||
.await;
|
||
}
|
||
StaticRangeDecision::Unsatisfiable => {
|
||
let mut range_headers = headers.clone();
|
||
range_headers.push((
|
||
"Content-Range".to_string(),
|
||
format!("bytes */{}", metadata.len),
|
||
));
|
||
return respond_head_with_headers(session, 416, Some(&content_type), &range_headers, 0)
|
||
.await;
|
||
}
|
||
}
|
||
|
||
if session.req_header().method == Method::HEAD {
|
||
return respond_head_with_headers(
|
||
session,
|
||
200,
|
||
Some(&content_type),
|
||
&headers,
|
||
metadata.len,
|
||
)
|
||
.await;
|
||
}
|
||
|
||
let body = match tokio::fs::read(&candidate.path).await {
|
||
Ok(bytes) => Bytes::from(bytes),
|
||
Err(error) => {
|
||
warn!(
|
||
path = %path,
|
||
file = %candidate.path.display(),
|
||
%error,
|
||
"静态文件读取失败"
|
||
);
|
||
respond_not_found(session).await?;
|
||
return Ok(());
|
||
}
|
||
};
|
||
|
||
respond_bytes_with_headers(session, 200, Some(&content_type), &headers, body).await
|
||
}
|
||
|
||
fn is_static_read_method(method: &Method) -> bool {
|
||
*method == Method::GET || *method == Method::HEAD
|
||
}
|
||
|
||
async fn read_static_range(path: &Path, start: u64, len: u64) -> io::Result<Bytes> {
|
||
let mut file = tokio::fs::File::open(path).await?;
|
||
file.seek(SeekFrom::Start(start)).await?;
|
||
let mut reader = file.take(len);
|
||
let capacity = usize::try_from(len)
|
||
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "Range 长度超过平台上限"))?;
|
||
let mut body = Vec::with_capacity(capacity);
|
||
reader.read_to_end(&mut body).await?;
|
||
Ok(Bytes::from(body))
|
||
}
|
||
|
||
async fn static_response_metadata(path: &Path) -> io::Result<StaticResponseMetadata> {
|
||
let metadata = tokio::fs::metadata(path).await?;
|
||
let len = metadata.len();
|
||
let last_modified = metadata.modified().ok();
|
||
Ok(StaticResponseMetadata {
|
||
len,
|
||
etag: build_static_etag(len, last_modified),
|
||
last_modified,
|
||
})
|
||
}
|
||
|
||
fn build_static_etag(len: u64, last_modified: Option<SystemTime>) -> String {
|
||
let modified_secs = last_modified
|
||
.and_then(|modified| modified.duration_since(UNIX_EPOCH).ok())
|
||
.map_or(0, |duration| duration.as_secs());
|
||
format!("W/\"{len:x}-{modified_secs:x}\"")
|
||
}
|
||
|
||
fn static_not_modified(session: &Session, metadata: &StaticResponseMetadata) -> bool {
|
||
if session.req_header().method != Method::GET && session.req_header().method != Method::HEAD {
|
||
return false;
|
||
}
|
||
|
||
if let Some(if_none_match) = header_value(session, "if-none-match") {
|
||
return if_none_match_matches(&if_none_match, &metadata.etag);
|
||
}
|
||
|
||
let Some(last_modified) = metadata.last_modified else {
|
||
return false;
|
||
};
|
||
header_value(session, "if-modified-since")
|
||
.and_then(|value| parse_http_date(&value).ok())
|
||
.is_some_and(|since| truncate_system_time_to_seconds(last_modified) <= since)
|
||
}
|
||
|
||
fn static_range_decision(
|
||
session: &Session,
|
||
metadata: &StaticResponseMetadata,
|
||
) -> StaticRangeDecision {
|
||
if session.req_header().method != Method::GET && session.req_header().method != Method::HEAD {
|
||
return StaticRangeDecision::Full;
|
||
}
|
||
|
||
let Some(range) = header_value(session, "range") else {
|
||
return StaticRangeDecision::Full;
|
||
};
|
||
if !static_if_range_matches(session, metadata) {
|
||
return StaticRangeDecision::Full;
|
||
}
|
||
parse_static_range(&range, metadata.len)
|
||
}
|
||
|
||
fn static_if_range_matches(session: &Session, metadata: &StaticResponseMetadata) -> bool {
|
||
let Some(if_range) = header_value(session, "if-range") else {
|
||
return true;
|
||
};
|
||
static_if_range_value_matches(&if_range, metadata)
|
||
}
|
||
|
||
fn static_if_range_value_matches(if_range: &str, metadata: &StaticResponseMetadata) -> bool {
|
||
if let Ok(if_range_date) = parse_http_date(&if_range) {
|
||
return metadata.last_modified.is_some_and(|last_modified| {
|
||
truncate_system_time_to_seconds(last_modified) <= if_range_date
|
||
});
|
||
}
|
||
|
||
strong_etag_opaque_value(&if_range)
|
||
.is_some_and(|candidate| strong_etag_opaque_value(&metadata.etag) == Some(candidate))
|
||
}
|
||
|
||
fn parse_static_range(range: &str, file_len: u64) -> StaticRangeDecision {
|
||
let range = range.trim();
|
||
let Some(unit) = range.get(..6) else {
|
||
return StaticRangeDecision::Full;
|
||
};
|
||
if !unit.eq_ignore_ascii_case("bytes=") {
|
||
return StaticRangeDecision::Full;
|
||
}
|
||
let spec = &range[6..];
|
||
if spec.contains(',') {
|
||
return StaticRangeDecision::Full;
|
||
}
|
||
if file_len == 0 {
|
||
return StaticRangeDecision::Unsatisfiable;
|
||
}
|
||
|
||
let Some((start, end)) = spec.split_once('-') else {
|
||
return StaticRangeDecision::Full;
|
||
};
|
||
let start = start.trim();
|
||
let end = end.trim();
|
||
|
||
if start.is_empty() {
|
||
let Ok(suffix_len) = end.parse::<u64>() else {
|
||
return StaticRangeDecision::Full;
|
||
};
|
||
if suffix_len == 0 {
|
||
return StaticRangeDecision::Unsatisfiable;
|
||
}
|
||
let start = file_len.saturating_sub(suffix_len);
|
||
return StaticRangeDecision::Partial {
|
||
start,
|
||
end: file_len - 1,
|
||
};
|
||
}
|
||
|
||
let Ok(start) = start.parse::<u64>() else {
|
||
return StaticRangeDecision::Full;
|
||
};
|
||
if start >= file_len {
|
||
return StaticRangeDecision::Unsatisfiable;
|
||
}
|
||
if end.is_empty() {
|
||
return StaticRangeDecision::Partial {
|
||
start,
|
||
end: file_len - 1,
|
||
};
|
||
}
|
||
|
||
let Ok(mut end) = end.parse::<u64>() else {
|
||
return StaticRangeDecision::Full;
|
||
};
|
||
if start > end {
|
||
return StaticRangeDecision::Unsatisfiable;
|
||
}
|
||
end = end.min(file_len - 1);
|
||
StaticRangeDecision::Partial { start, end }
|
||
}
|
||
|
||
fn if_none_match_matches(value: &str, etag: &str) -> bool {
|
||
let expected = weak_etag_opaque_value(etag);
|
||
value
|
||
.split(',')
|
||
.map(str::trim)
|
||
.any(|candidate| candidate == "*" || weak_etag_opaque_value(candidate) == expected)
|
||
}
|
||
|
||
fn weak_etag_opaque_value(value: &str) -> &str {
|
||
value.strip_prefix("W/").unwrap_or(value)
|
||
}
|
||
|
||
fn strong_etag_opaque_value(value: &str) -> Option<&str> {
|
||
if value.trim().starts_with("W/") {
|
||
return None;
|
||
}
|
||
Some(value.trim())
|
||
}
|
||
|
||
fn truncate_system_time_to_seconds(value: SystemTime) -> SystemTime {
|
||
value
|
||
.duration_since(UNIX_EPOCH)
|
||
.map(|duration| UNIX_EPOCH + Duration::from_secs(duration.as_secs()))
|
||
.unwrap_or(UNIX_EPOCH)
|
||
}
|
||
|
||
async fn resolve_static_candidate(
|
||
root: &Path,
|
||
root_kind: StaticRoot,
|
||
path: &str,
|
||
mode: StaticMode,
|
||
) -> Option<StaticCandidate> {
|
||
let candidate = sanitize_static_path(root, path)?;
|
||
if is_regular_file(&candidate).await {
|
||
return Some(StaticCandidate {
|
||
cache_kind: classify_static_cache_kind(root_kind, path, &candidate),
|
||
path: candidate,
|
||
});
|
||
}
|
||
|
||
if is_directory(&candidate).await {
|
||
let index = candidate.join("index.html");
|
||
if is_regular_file(&index).await {
|
||
return Some(StaticCandidate {
|
||
path: index,
|
||
cache_kind: StaticCacheKind::Html,
|
||
});
|
||
}
|
||
}
|
||
|
||
match mode {
|
||
StaticMode::Exact => None,
|
||
StaticMode::SpaFallback if path.starts_with("/admin/") => {
|
||
let fallback = root.join("admin").join("index.html");
|
||
is_regular_file(&fallback).await.then_some(StaticCandidate {
|
||
path: fallback,
|
||
cache_kind: StaticCacheKind::Html,
|
||
})
|
||
}
|
||
StaticMode::SpaFallback => {
|
||
let fallback = root.join("index.html");
|
||
is_regular_file(&fallback).await.then_some(StaticCandidate {
|
||
path: fallback,
|
||
cache_kind: StaticCacheKind::Html,
|
||
})
|
||
}
|
||
}
|
||
}
|
||
|
||
fn classify_static_cache_kind(
|
||
root_kind: StaticRoot,
|
||
request_path: &str,
|
||
file_path: &Path,
|
||
) -> StaticCacheKind {
|
||
if is_html_path(file_path) {
|
||
return StaticCacheKind::Html;
|
||
}
|
||
|
||
if root_kind == StaticRoot::Web
|
||
&& (request_path.starts_with("/assets/") || request_path.starts_with("/admin/assets/"))
|
||
&& has_fingerprinted_file_name(file_path)
|
||
{
|
||
return StaticCacheKind::FingerprintedAsset;
|
||
}
|
||
|
||
StaticCacheKind::Other
|
||
}
|
||
|
||
fn cache_control_for_static_candidate<'a>(
|
||
candidate: &StaticCandidate,
|
||
config: &'a StaticCacheConfig,
|
||
) -> &'a str {
|
||
match candidate.cache_kind {
|
||
StaticCacheKind::Html => config.html_cache_control.as_str(),
|
||
StaticCacheKind::FingerprintedAsset => config.asset_cache_control.as_str(),
|
||
StaticCacheKind::Other => config.static_cache_control.as_str(),
|
||
}
|
||
}
|
||
|
||
fn is_html_path(path: &Path) -> bool {
|
||
path.extension()
|
||
.and_then(|extension| extension.to_str())
|
||
.is_some_and(|extension| extension.eq_ignore_ascii_case("html"))
|
||
}
|
||
|
||
fn has_fingerprinted_file_name(path: &Path) -> bool {
|
||
let Some(file_stem) = path.file_stem().and_then(|value| value.to_str()) else {
|
||
return false;
|
||
};
|
||
file_stem
|
||
.rsplit_once('-')
|
||
.is_some_and(|(_, suffix)| is_asset_fingerprint(suffix))
|
||
}
|
||
|
||
fn is_asset_fingerprint(value: &str) -> bool {
|
||
value.len() >= 8 && value.bytes().all(|byte| byte.is_ascii_alphanumeric())
|
||
}
|
||
|
||
fn sanitize_static_path(root: &Path, path: &str) -> Option<PathBuf> {
|
||
let mut result = root.to_path_buf();
|
||
for raw_segment in path.trim_start_matches('/').split('/') {
|
||
if raw_segment.is_empty() {
|
||
continue;
|
||
}
|
||
|
||
let segment = decode_static_path_segment(raw_segment)?;
|
||
match segment.as_str() {
|
||
"." => continue,
|
||
".." => return None,
|
||
_ => result.push(segment),
|
||
}
|
||
}
|
||
Some(result)
|
||
}
|
||
|
||
fn decode_static_path_segment(segment: &str) -> Option<String> {
|
||
let bytes = segment.as_bytes();
|
||
let mut decoded = Vec::with_capacity(bytes.len());
|
||
let mut index = 0;
|
||
|
||
while index < bytes.len() {
|
||
if bytes[index] == b'%' {
|
||
let high = *bytes.get(index + 1)?;
|
||
let low = *bytes.get(index + 2)?;
|
||
let value = decode_percent_hex(high)? << 4 | decode_percent_hex(low)?;
|
||
if value == b'/' || value == b'\\' || value == b'\0' {
|
||
return None;
|
||
}
|
||
decoded.push(value);
|
||
index += 3;
|
||
} else {
|
||
if bytes[index] == b'\\' || bytes[index] == b'\0' {
|
||
return None;
|
||
}
|
||
decoded.push(bytes[index]);
|
||
index += 1;
|
||
}
|
||
}
|
||
|
||
String::from_utf8(decoded).ok()
|
||
}
|
||
|
||
fn decode_percent_hex(byte: u8) -> Option<u8> {
|
||
match byte {
|
||
b'0'..=b'9' => Some(byte - b'0'),
|
||
b'a'..=b'f' => Some(byte - b'a' + 10),
|
||
b'A'..=b'F' => Some(byte - b'A' + 10),
|
||
_ => None,
|
||
}
|
||
}
|
||
|
||
async fn is_regular_file(path: &Path) -> bool {
|
||
tokio::fs::metadata(path)
|
||
.await
|
||
.map(|metadata| metadata.is_file())
|
||
.unwrap_or(false)
|
||
}
|
||
|
||
async fn is_directory(path: &Path) -> bool {
|
||
tokio::fs::metadata(path)
|
||
.await
|
||
.map(|metadata| metadata.is_dir())
|
||
.unwrap_or(false)
|
||
}
|
||
|
||
async fn respond_maintenance(
|
||
session: &mut Session,
|
||
api_like: bool,
|
||
maintenance_page_file: &Path,
|
||
web_root: &Path,
|
||
) -> PingoraResult<()> {
|
||
if api_like {
|
||
return respond_json(
|
||
session,
|
||
503,
|
||
r#"{"ok":false,"error":{"code":"MAINTENANCE","message":"服务维护中"}}"#,
|
||
)
|
||
.await;
|
||
}
|
||
|
||
let default_maintenance_page = web_root.join("maintenance.html");
|
||
let maintenance_page = if is_regular_file(maintenance_page_file).await {
|
||
Some(maintenance_page_file)
|
||
} else if is_regular_file(&default_maintenance_page).await {
|
||
Some(default_maintenance_page.as_path())
|
||
} else {
|
||
None
|
||
};
|
||
|
||
if let Some(maintenance_page) = maintenance_page {
|
||
let body = Bytes::from(tokio::fs::read(maintenance_page).await.unwrap_or_default());
|
||
respond_bytes(session, 503, Some("text/html; charset=utf-8"), body).await
|
||
} else {
|
||
respond_bytes(
|
||
session,
|
||
503,
|
||
Some("text/plain; charset=utf-8"),
|
||
Bytes::from_static("服务维护中".as_bytes()),
|
||
)
|
||
.await
|
||
}
|
||
}
|
||
|
||
async fn respond_redirect(session: &mut Session, location: &str) -> PingoraResult<()> {
|
||
let mut header = ResponseHeader::build(301, Some(0))?;
|
||
header.insert_header("Location", location)?;
|
||
header.insert_header("Content-Length", "0")?;
|
||
session.write_response_header(Box::new(header), true).await
|
||
}
|
||
|
||
async fn respond_http_to_https_redirect(
|
||
session: &mut Session,
|
||
target_scheme: &str,
|
||
) -> PingoraResult<()> {
|
||
let Some(host) = host_or_authority(session).filter(|value| is_valid_redirect_host(value))
|
||
else {
|
||
respond_json(
|
||
session,
|
||
400,
|
||
r#"{"ok":false,"error":{"code":"BAD_REQUEST","message":"无效请求"}}"#,
|
||
)
|
||
.await?;
|
||
return Ok(());
|
||
};
|
||
let location = format!(
|
||
"{target_scheme}://{host}{}",
|
||
session
|
||
.req_header()
|
||
.uri
|
||
.path_and_query()
|
||
.map_or("/", |value| value.as_str())
|
||
);
|
||
let mut header = ResponseHeader::build(301, Some(0))?;
|
||
header.insert_header("Location", location.as_str())?;
|
||
header.insert_header("Content-Length", "0")?;
|
||
header.insert_header("X-Genarrative-Gateway", "pingora-shadow")?;
|
||
session.write_response_header(Box::new(header), true).await
|
||
}
|
||
|
||
async fn respond_json(session: &mut Session, status: u16, body: &str) -> PingoraResult<()> {
|
||
respond_bytes(
|
||
session,
|
||
status,
|
||
Some("application/json; charset=utf-8"),
|
||
Bytes::copy_from_slice(body.as_bytes()),
|
||
)
|
||
.await
|
||
}
|
||
|
||
async fn respond_not_found(session: &mut Session) -> PingoraResult<()> {
|
||
respond_bytes(session, 404, None, Bytes::new()).await
|
||
}
|
||
|
||
async fn respond_page_not_found(session: &mut Session, web_root: &Path) -> PingoraResult<()> {
|
||
let accepts_html = session
|
||
.req_header()
|
||
.headers
|
||
.get("accept")
|
||
.and_then(|value| value.to_str().ok())
|
||
.is_some_and(|value| {
|
||
value
|
||
.split(',')
|
||
.any(|entry| entry.trim().starts_with("text/html"))
|
||
});
|
||
let not_found_page = web_root.join("404.html");
|
||
if accepts_html && is_regular_file(¬_found_page).await {
|
||
let body = Bytes::from(tokio::fs::read(¬_found_page).await.unwrap_or_default());
|
||
return respond_bytes_with_headers(
|
||
session,
|
||
404,
|
||
Some("text/html; charset=utf-8"),
|
||
&[("Cache-Control".to_string(), "no-store".to_string())],
|
||
body,
|
||
)
|
||
.await;
|
||
}
|
||
|
||
respond_not_found(session).await
|
||
}
|
||
|
||
async fn respond_static_method_not_allowed(session: &mut Session) -> PingoraResult<()> {
|
||
respond_head_with_headers(
|
||
session,
|
||
405,
|
||
None,
|
||
&[("Allow".to_string(), "GET, HEAD".to_string())],
|
||
0,
|
||
)
|
||
.await
|
||
}
|
||
|
||
async fn respond_gateway_proxy_error(session: &mut Session, status: u16) -> PingoraResult<()> {
|
||
respond_json(session, status, gateway_proxy_error_body(status)).await
|
||
}
|
||
|
||
async fn respond_too_many_requests(
|
||
session: &mut Session,
|
||
reason: RejectReason,
|
||
) -> PingoraResult<()> {
|
||
let body = format!(
|
||
r#"{{"ok":false,"error":{{"code":"{}","message":"{}"}}}}"#,
|
||
reason.code(),
|
||
reason.message()
|
||
);
|
||
let is_head = session.req_header().method == Method::HEAD;
|
||
let mut header = ResponseHeader::build(429, Some(body.len()))?;
|
||
header.insert_header("Content-Length", body.len().to_string())?;
|
||
header.insert_header("Content-Type", "application/json; charset=utf-8")?;
|
||
header.insert_header("Retry-After", "1")?;
|
||
header.insert_header("X-Genarrative-Gateway", "pingora-shadow")?;
|
||
session
|
||
.write_response_header(Box::new(header), is_head)
|
||
.await?;
|
||
if !is_head {
|
||
session
|
||
.write_response_body(Some(Bytes::from(body)), true)
|
||
.await?;
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
async fn respond_shadow_probe(
|
||
session: &mut Session,
|
||
maintenance_enabled: bool,
|
||
) -> PingoraResult<()> {
|
||
let maintenance = if maintenance_enabled { "true" } else { "false" };
|
||
let body = format!(r#"{{"ok":true,"gateway":"pingora-shadow","maintenance":{maintenance}}}"#);
|
||
respond_bytes(
|
||
session,
|
||
200,
|
||
Some("application/json; charset=utf-8"),
|
||
Bytes::from(body),
|
||
)
|
||
.await
|
||
}
|
||
|
||
async fn respond_bytes(
|
||
session: &mut Session,
|
||
status: u16,
|
||
content_type: Option<&str>,
|
||
body: Bytes,
|
||
) -> PingoraResult<()> {
|
||
respond_bytes_with_headers(session, status, content_type, &[], body).await
|
||
}
|
||
|
||
async fn respond_bytes_with_headers(
|
||
session: &mut Session,
|
||
status: u16,
|
||
content_type: Option<&str>,
|
||
headers: &[(String, String)],
|
||
body: Bytes,
|
||
) -> PingoraResult<()> {
|
||
let is_head = session.req_header().method == Method::HEAD;
|
||
let mut header = ResponseHeader::build(status, Some(body.len()))?;
|
||
header.insert_header("Content-Length", body.len().to_string())?;
|
||
if let Some(content_type) = content_type {
|
||
header.insert_header("Content-Type", content_type)?;
|
||
}
|
||
for (name, value) in headers {
|
||
if !value.is_empty() {
|
||
header.insert_header(name.clone(), value.as_str())?;
|
||
}
|
||
}
|
||
header.insert_header("X-Genarrative-Gateway", "pingora-shadow")?;
|
||
session
|
||
.write_response_header(Box::new(header), is_head || body.is_empty())
|
||
.await?;
|
||
if !is_head && !body.is_empty() {
|
||
session.write_response_body(Some(body), true).await?;
|
||
}
|
||
Ok(())
|
||
}
|
||
|
||
async fn respond_head_with_headers(
|
||
session: &mut Session,
|
||
status: u16,
|
||
content_type: Option<&str>,
|
||
headers: &[(String, String)],
|
||
content_length: u64,
|
||
) -> PingoraResult<()> {
|
||
let mut header = ResponseHeader::build(status, Some(headers.len() + 3))?;
|
||
header.insert_header("Content-Length", content_length.to_string())?;
|
||
if let Some(content_type) = content_type {
|
||
header.insert_header("Content-Type", content_type)?;
|
||
}
|
||
for (name, value) in headers {
|
||
if !value.is_empty() {
|
||
header.insert_header(name.clone(), value.as_str())?;
|
||
}
|
||
}
|
||
header.insert_header("X-Genarrative-Gateway", "pingora-shadow")?;
|
||
session.write_response_header(Box::new(header), true).await
|
||
}
|
||
|
||
async fn respond_empty_with_headers(
|
||
session: &mut Session,
|
||
status: u16,
|
||
content_type: Option<&str>,
|
||
headers: &[(String, String)],
|
||
) -> PingoraResult<()> {
|
||
let mut header = ResponseHeader::build(status, Some(headers.len() + 2))?;
|
||
if let Some(content_type) = content_type {
|
||
header.insert_header("Content-Type", content_type)?;
|
||
}
|
||
for (name, value) in headers {
|
||
if !value.is_empty() {
|
||
header.insert_header(name.clone(), value.as_str())?;
|
||
}
|
||
}
|
||
header.insert_header("X-Genarrative-Gateway", "pingora-shadow")?;
|
||
session.write_response_header(Box::new(header), true).await
|
||
}
|
||
|
||
fn resolve_request_id(session: &Session) -> String {
|
||
session
|
||
.req_header()
|
||
.headers
|
||
.get("x-request-id")
|
||
.and_then(|value| value.to_str().ok())
|
||
.filter(|value| !value.trim().is_empty())
|
||
.map(ToOwned::to_owned)
|
||
.unwrap_or_else(|| Uuid::new_v4().to_string())
|
||
}
|
||
|
||
fn content_length(session: &Session) -> Option<u64> {
|
||
header_value(session, "content-length")?.parse().ok()
|
||
}
|
||
|
||
fn header_value(session: &Session, name: &str) -> Option<String> {
|
||
session
|
||
.req_header()
|
||
.headers
|
||
.get(name)
|
||
.and_then(|value| value.to_str().ok())
|
||
.filter(|value| !value.trim().is_empty())
|
||
.map(ToOwned::to_owned)
|
||
}
|
||
|
||
fn host_or_authority(session: &Session) -> Option<String> {
|
||
header_value(session, "host").or_else(|| {
|
||
session
|
||
.req_header()
|
||
.uri
|
||
.authority()
|
||
.map(|authority| authority.as_str().to_string())
|
||
})
|
||
}
|
||
|
||
fn is_valid_redirect_host(host: &str) -> bool {
|
||
!host.is_empty()
|
||
&& host.len() <= 255
|
||
&& host.bytes().all(|byte| {
|
||
byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'-' | b':' | b'[' | b']')
|
||
})
|
||
}
|
||
|
||
fn client_ip(session: &Session) -> Option<String> {
|
||
session
|
||
.as_downstream()
|
||
.client_addr()
|
||
.and_then(|addr| addr.as_inet())
|
||
.map(|addr| addr.ip().to_string())
|
||
}
|
||
|
||
fn request_source_ip(session: &Session) -> Option<IpAddr> {
|
||
let peer_ip = session
|
||
.as_downstream()
|
||
.client_addr()
|
||
.and_then(|addr| addr.as_inet())
|
||
.map(|addr| addr.ip())?;
|
||
if peer_ip.is_loopback()
|
||
&& let Some(real_ip) =
|
||
header_value(session, "x-real-ip").and_then(|value| value.trim().parse::<IpAddr>().ok())
|
||
{
|
||
return Some(real_ip);
|
||
}
|
||
Some(peer_ip)
|
||
}
|
||
|
||
fn is_admin_request_path(path: &str) -> bool {
|
||
path == "/admin" || path.starts_with("/admin/")
|
||
}
|
||
|
||
fn is_internal_network_ip(ip: &IpAddr) -> bool {
|
||
match ip {
|
||
IpAddr::V4(ip) => ip.is_private() || ip.is_loopback() || ip.is_link_local(),
|
||
IpAddr::V6(ip) => ip.is_loopback() || ip.is_unique_local() || ip.is_unicast_link_local(),
|
||
}
|
||
}
|
||
|
||
fn allows_internal_maintenance_bypass(client_ip: Option<&IpAddr>) -> bool {
|
||
client_ip.is_some_and(is_internal_network_ip)
|
||
}
|
||
|
||
fn append_forwarded_for(session: &Session, client_ip: &str) -> String {
|
||
session
|
||
.req_header()
|
||
.headers
|
||
.get("x-forwarded-for")
|
||
.and_then(|value| value.to_str().ok())
|
||
.filter(|value| !value.trim().is_empty())
|
||
.map(|existing| format!("{existing}, {client_ip}"))
|
||
.unwrap_or_else(|| client_ip.to_string())
|
||
}
|
||
|
||
fn protection_client_id(session: &Session, trust_x_forwarded_for: bool) -> String {
|
||
if trust_x_forwarded_for && let Some(forwarded_for) = first_forwarded_for(session) {
|
||
return forwarded_for;
|
||
}
|
||
|
||
client_ip(session).unwrap_or_else(|| "unknown".to_string())
|
||
}
|
||
|
||
fn first_forwarded_for(session: &Session) -> Option<String> {
|
||
session
|
||
.req_header()
|
||
.headers
|
||
.get("x-forwarded-for")
|
||
.and_then(|value| value.to_str().ok())
|
||
.and_then(|value| value.split(',').next())
|
||
.map(str::trim)
|
||
.filter(|value| !value.is_empty())
|
||
.map(ToOwned::to_owned)
|
||
}
|
||
|
||
fn is_upgrade_request(session: &Session) -> bool {
|
||
session
|
||
.req_header()
|
||
.headers
|
||
.get("upgrade")
|
||
.and_then(|value| value.to_str().ok())
|
||
.is_some_and(|value| value.eq_ignore_ascii_case("websocket"))
|
||
}
|
||
|
||
fn normalize_accept_encoding_for_gateway_compression(
|
||
req: &mut RequestHeader,
|
||
config: CompressionRequestConfig,
|
||
) -> PingoraResult<()> {
|
||
let algorithm = if config.enabled {
|
||
req.headers
|
||
.get("accept-encoding")
|
||
.and_then(|value| value.to_str().ok())
|
||
.and_then(|value| select_compression_algorithm(value, config))
|
||
} else {
|
||
None
|
||
};
|
||
|
||
if let Some(algorithm) = algorithm {
|
||
req.insert_header("Accept-Encoding", algorithm.header_value())?;
|
||
} else {
|
||
req.remove_header("accept-encoding");
|
||
}
|
||
|
||
Ok(())
|
||
}
|
||
|
||
fn should_allow_gateway_compression_for_response(
|
||
resp: &ResponseHeader,
|
||
min_length_bytes: u64,
|
||
) -> bool {
|
||
if resp.status.as_u16() != 200
|
||
|| resp.headers.get("content-range").is_some()
|
||
|| resp.headers.get("content-encoding").is_some()
|
||
{
|
||
return false;
|
||
}
|
||
|
||
resp.headers
|
||
.get("content-length")
|
||
.and_then(|value| value.to_str().ok())
|
||
.and_then(|value| value.parse::<u64>().ok())
|
||
.is_none_or(|content_length| content_length >= min_length_bytes)
|
||
}
|
||
|
||
fn select_compression_algorithm(
|
||
value: &str,
|
||
config: CompressionRequestConfig,
|
||
) -> Option<CompressionAlgorithm> {
|
||
parse_accepted_encodings(value)
|
||
.into_iter()
|
||
.find_map(|coding| {
|
||
if coding.eq_ignore_ascii_case("gzip") && config.gzip {
|
||
Some(CompressionAlgorithm::Gzip)
|
||
} else {
|
||
None
|
||
}
|
||
})
|
||
}
|
||
|
||
fn parse_accepted_encodings(value: &str) -> Vec<String> {
|
||
value
|
||
.split(',')
|
||
.filter_map(|item| {
|
||
let mut parts = item.split(';');
|
||
let coding = parts.next()?.trim();
|
||
if coding.is_empty() {
|
||
return None;
|
||
}
|
||
|
||
let enabled = parts
|
||
.map(str::trim)
|
||
.filter_map(|param| param.split_once('='))
|
||
.filter(|(name, _)| name.trim().eq_ignore_ascii_case("q"))
|
||
.all(|(_, value)| value.trim() != "0" && value.trim() != "0.0");
|
||
if enabled {
|
||
Some(coding.to_string())
|
||
} else {
|
||
None
|
||
}
|
||
})
|
||
.collect()
|
||
}
|
||
|
||
fn read_env_or_default(key: &str, default_value: &str) -> String {
|
||
env::var(key)
|
||
.ok()
|
||
.filter(|value| !value.trim().is_empty())
|
||
.unwrap_or_else(|| default_value.to_string())
|
||
}
|
||
|
||
fn read_optional_env(key: &str) -> Option<String> {
|
||
env::var(key)
|
||
.ok()
|
||
.map(|value| value.trim().to_string())
|
||
.filter(|value| !value.is_empty())
|
||
}
|
||
|
||
fn read_bool_env(key: &str, default_value: bool) -> bool {
|
||
env::var(key)
|
||
.ok()
|
||
.and_then(|value| match value.trim().to_ascii_lowercase().as_str() {
|
||
"1" | "true" | "yes" | "on" => Some(true),
|
||
"0" | "false" | "no" | "off" => Some(false),
|
||
_ => None,
|
||
})
|
||
.unwrap_or(default_value)
|
||
}
|
||
|
||
fn read_u32_env(key: &str, default_value: u32) -> io::Result<u32> {
|
||
match env::var(key) {
|
||
Ok(value) if !value.trim().is_empty() => value.trim().parse().map_err(|error| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 不是有效整数:{error}"),
|
||
)
|
||
}),
|
||
_ => Ok(default_value),
|
||
}
|
||
}
|
||
|
||
fn read_u64_env(key: &str, default_value: u64) -> io::Result<u64> {
|
||
match env::var(key) {
|
||
Ok(value) if !value.trim().is_empty() => value.trim().parse().map_err(|error| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 不是有效整数:{error}"),
|
||
)
|
||
}),
|
||
_ => Ok(default_value),
|
||
}
|
||
}
|
||
|
||
fn read_socket_addr_env(key: &str, default_value: &str) -> io::Result<SocketAddr> {
|
||
let value = read_env_or_default(key, default_value);
|
||
value.parse().map_err(|error| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 不是有效 socket 地址 {value:?}:{error}"),
|
||
)
|
||
})
|
||
}
|
||
|
||
fn read_optional_socket_addr_env(key: &str) -> io::Result<Option<SocketAddr>> {
|
||
read_optional_env(key)
|
||
.map(|value| {
|
||
value.parse().map_err(|error| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 不是有效 socket 地址 {value:?}:{error}"),
|
||
)
|
||
})
|
||
})
|
||
.transpose()
|
||
}
|
||
|
||
fn read_host_list_env(key: &str) -> io::Result<Vec<String>> {
|
||
let Some(raw) = read_optional_env(key) else {
|
||
return Ok(Vec::new());
|
||
};
|
||
let mut result = Vec::new();
|
||
for value in raw
|
||
.split(',')
|
||
.map(str::trim)
|
||
.filter(|value| !value.is_empty())
|
||
{
|
||
let host = normalize_gateway_host(value).ok_or_else(|| {
|
||
io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 包含无效 Host:{value:?}"),
|
||
)
|
||
})?;
|
||
if !result.contains(&host) {
|
||
result.push(host);
|
||
}
|
||
}
|
||
Ok(result)
|
||
}
|
||
|
||
fn normalize_gateway_host(value: &str) -> Option<String> {
|
||
let value = value.trim();
|
||
if value.is_empty() || value.contains(['/', '\\', '\n', '\r', '\0']) {
|
||
return None;
|
||
}
|
||
if value.starts_with('[') {
|
||
return None;
|
||
}
|
||
let host = value
|
||
.rsplit_once(':')
|
||
.and_then(|(host, port)| {
|
||
(!port.is_empty() && port.bytes().all(|byte| byte.is_ascii_digit())).then_some(host)
|
||
})
|
||
.unwrap_or(value)
|
||
.trim_end_matches('.');
|
||
if host.is_empty() || host.len() > 253 {
|
||
return None;
|
||
}
|
||
if !host
|
||
.bytes()
|
||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'-'))
|
||
{
|
||
return None;
|
||
}
|
||
Some(host.to_ascii_lowercase())
|
||
}
|
||
|
||
fn validate_configured_gateway_host(key: &str, host: &str) -> io::Result<()> {
|
||
if normalize_gateway_host(host).as_deref() == Some(host) {
|
||
return Ok(());
|
||
}
|
||
|
||
Err(io::Error::new(
|
||
io::ErrorKind::InvalidInput,
|
||
format!("{key} 包含无效 Host:{host:?}"),
|
||
))
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
use serde::Deserialize;
|
||
|
||
const ROUTE_PARITY_MATRIX_JSON: &str =
|
||
include_str!("../../../../deploy/pingora/nginx-route-parity.matrix.json");
|
||
|
||
#[derive(Deserialize)]
|
||
struct RouteParityMatrix {
|
||
version: u32,
|
||
routes: Vec<RouteParityCase>,
|
||
}
|
||
|
||
#[derive(Deserialize)]
|
||
#[serde(rename_all = "camelCase")]
|
||
struct RouteParityCase {
|
||
id: String,
|
||
sample_path: String,
|
||
expect: RouteExpectation,
|
||
}
|
||
|
||
#[derive(Deserialize)]
|
||
#[serde(rename_all = "camelCase")]
|
||
struct RouteExpectation {
|
||
kind: String,
|
||
target: Option<String>,
|
||
body_limit: Option<BodyLimitExpectation>,
|
||
root: Option<String>,
|
||
mode: Option<String>,
|
||
location: Option<String>,
|
||
protection_class: Option<String>,
|
||
}
|
||
|
||
#[derive(Deserialize)]
|
||
#[serde(untagged)]
|
||
enum BodyLimitExpectation {
|
||
Named(String),
|
||
Bytes(u64),
|
||
}
|
||
|
||
fn assert_route_matches_expectation(case: &RouteParityCase, route: &RouteDecision) {
|
||
match (case.expect.kind.as_str(), route) {
|
||
("proxy", RouteDecision::Proxy { target, body_limit }) => {
|
||
assert_eq!(
|
||
Some(proxy_target_name(*target)),
|
||
case.expect.target.as_deref(),
|
||
"route parity proxy target mismatch: {} {}",
|
||
case.id,
|
||
case.sample_path
|
||
);
|
||
assert_eq!(
|
||
*body_limit,
|
||
expected_body_limit(case.expect.body_limit.as_ref()),
|
||
"route parity body limit mismatch: {} {}",
|
||
case.id,
|
||
case.sample_path
|
||
);
|
||
}
|
||
("static", RouteDecision::Local(LocalResponse::Static { root, mode })) => {
|
||
assert_eq!(
|
||
Some(static_root_name(*root)),
|
||
case.expect.root.as_deref(),
|
||
"route parity static root mismatch: {} {}",
|
||
case.id,
|
||
case.sample_path
|
||
);
|
||
assert_eq!(
|
||
Some(static_mode_name(*mode)),
|
||
case.expect.mode.as_deref(),
|
||
"route parity static mode mismatch: {} {}",
|
||
case.id,
|
||
case.sample_path
|
||
);
|
||
}
|
||
(
|
||
"redirect_permanent",
|
||
RouteDecision::Local(LocalResponse::RedirectPermanent { location }),
|
||
) => {
|
||
assert_eq!(
|
||
Some(*location),
|
||
case.expect.location.as_deref(),
|
||
"route parity redirect location mismatch: {} {}",
|
||
case.id,
|
||
case.sample_path
|
||
);
|
||
}
|
||
("shadow_probe", RouteDecision::Local(LocalResponse::ShadowProbe)) => {}
|
||
("not_found", RouteDecision::Local(LocalResponse::NotFound)) => {}
|
||
_ => panic!(
|
||
"route parity kind mismatch: {} {} expected {} got {:?}",
|
||
case.id, case.sample_path, case.expect.kind, route
|
||
),
|
||
}
|
||
}
|
||
|
||
fn proxy_target_name(target: ProxyTarget) -> &'static str {
|
||
match target {
|
||
ProxyTarget::Api => "api",
|
||
ProxyTarget::Spacetime => "spacetime",
|
||
ProxyTarget::Gitea => "gitea",
|
||
}
|
||
}
|
||
|
||
fn static_root_name(root: StaticRoot) -> &'static str {
|
||
match root {
|
||
StaticRoot::Web => "web",
|
||
StaticRoot::Acme => "acme",
|
||
}
|
||
}
|
||
|
||
fn static_mode_name(mode: StaticMode) -> &'static str {
|
||
match mode {
|
||
StaticMode::Exact => "exact",
|
||
StaticMode::SpaFallback => "spa_fallback",
|
||
}
|
||
}
|
||
|
||
fn expected_body_limit(expectation: Option<&BodyLimitExpectation>) -> Option<u64> {
|
||
match expectation {
|
||
None => None,
|
||
Some(BodyLimitExpectation::Bytes(bytes)) => Some(*bytes),
|
||
Some(BodyLimitExpectation::Named(name)) if name == "default" => {
|
||
Some(DEFAULT_MAX_API_BODY_BYTES)
|
||
}
|
||
Some(BodyLimitExpectation::Named(name)) => {
|
||
panic!("unknown route parity body limit expectation: {name}")
|
||
}
|
||
}
|
||
}
|
||
|
||
fn test_gateway_config() -> GatewayConfig {
|
||
GatewayConfig {
|
||
listen_addr: "127.0.0.1:18081".to_string(),
|
||
tls_listen_addr: None,
|
||
tls_cert_file: None,
|
||
tls_key_file: None,
|
||
http_redirect_listen_addr: None,
|
||
http_redirect_target_scheme: "https".to_string(),
|
||
api_upstream: "127.0.0.1:8082".parse().unwrap(),
|
||
spacetime_upstream: "127.0.0.1:3101".parse().unwrap(),
|
||
gitea_hosts: Vec::new(),
|
||
gitea_upstream: None,
|
||
web_root: PathBuf::from("/srv/genarrative/web"),
|
||
acme_root: PathBuf::from("/var/www/html"),
|
||
maintenance_file: PathBuf::from("/var/lib/genarrative/maintenance/enabled"),
|
||
maintenance_page_file: PathBuf::from("/var/lib/genarrative/maintenance/page.html"),
|
||
forwarded_proto: "http".to_string(),
|
||
max_api_body_bytes: DEFAULT_MAX_API_BODY_BYTES,
|
||
gzip_enabled: true,
|
||
gzip_level: DEFAULT_GZIP_LEVEL,
|
||
gzip_min_length_bytes: DEFAULT_GZIP_MIN_LENGTH_BYTES,
|
||
compression: CompressionConfig {
|
||
gzip: true,
|
||
raw_algorithms: "gzip".to_string(),
|
||
},
|
||
static_cache: StaticCacheConfig {
|
||
html_cache_control: DEFAULT_HTML_CACHE_CONTROL.to_string(),
|
||
asset_cache_control: DEFAULT_ASSET_CACHE_CONTROL.to_string(),
|
||
static_cache_control: DEFAULT_STATIC_CACHE_CONTROL.to_string(),
|
||
},
|
||
upstream_timeouts: UpstreamTimeoutConfig {
|
||
connect_timeout_ms: DEFAULT_UPSTREAM_CONNECT_TIMEOUT_MS,
|
||
default_read_timeout_seconds: DEFAULT_UPSTREAM_DEFAULT_READ_TIMEOUT_SECONDS,
|
||
api_read_timeout_seconds: DEFAULT_UPSTREAM_API_READ_TIMEOUT_SECONDS,
|
||
long_read_timeout_seconds: DEFAULT_UPSTREAM_LONG_READ_TIMEOUT_SECONDS,
|
||
write_timeout_seconds: DEFAULT_UPSTREAM_WRITE_TIMEOUT_SECONDS,
|
||
},
|
||
shadow_probe_token: None,
|
||
trust_x_forwarded_for: false,
|
||
instance_count: DEFAULT_INSTANCE_COUNT,
|
||
shared_protection_confirmed: false,
|
||
protection: ProtectionConfig {
|
||
enabled: true,
|
||
admin_api: ProtectionClassConfig {
|
||
max_concurrent: 64,
|
||
rate_per_second: 30,
|
||
burst: 16,
|
||
},
|
||
api: ProtectionClassConfig {
|
||
max_concurrent: 64,
|
||
rate_per_second: 300,
|
||
burst: 64,
|
||
},
|
||
spacetime: ProtectionClassConfig {
|
||
max_concurrent: 256,
|
||
rate_per_second: 1000,
|
||
burst: 256,
|
||
},
|
||
},
|
||
log_filter: DEFAULT_LOG_FILTER.to_string(),
|
||
access_log_file: None,
|
||
otel_enabled: false,
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn matches_nginx_route_parity_matrix() {
|
||
let matrix: RouteParityMatrix = serde_json::from_str(ROUTE_PARITY_MATRIX_JSON)
|
||
.expect("route parity matrix should be valid JSON");
|
||
assert_eq!(matrix.version, 1);
|
||
|
||
for case in matrix.routes {
|
||
let route = classify_path(&case.sample_path);
|
||
assert_route_matches_expectation(&case, &route);
|
||
|
||
let expected_protection_class = case.expect.protection_class.as_deref();
|
||
let actual_protection_class =
|
||
protection_class_for_route(&route, &case.sample_path).map(ProtectionClass::as_str);
|
||
assert_eq!(
|
||
actual_protection_class, expected_protection_class,
|
||
"route parity protection class mismatch: {} {}",
|
||
case.id, case.sample_path
|
||
);
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn applies_configured_body_limit_to_generic_api_routes_only() {
|
||
let mut generic_api = classify_path("/api/assets/history");
|
||
apply_configured_body_limit(&mut generic_api, 1024);
|
||
assert_eq!(generic_api.body_limit(), Some(1024));
|
||
|
||
let mut gitea = RouteDecision::Proxy {
|
||
target: ProxyTarget::Gitea,
|
||
body_limit: None,
|
||
};
|
||
apply_configured_body_limit(&mut gitea, 1024);
|
||
assert_eq!(gitea.body_limit(), None);
|
||
}
|
||
|
||
#[test]
|
||
fn routes_configured_gitea_hosts_to_gitea_upstream() {
|
||
let config = GatewayConfig {
|
||
gitea_hosts: vec!["git.genarrative.world".to_string()],
|
||
gitea_upstream: Some("127.0.0.1:3000".parse().unwrap()),
|
||
..test_gateway_config()
|
||
};
|
||
let gateway = GenarrativeGateway {
|
||
config: Arc::new(config),
|
||
protection: Arc::new(GatewayProtection::new(test_gateway_config().protection)),
|
||
mode: GatewayMode::Proxy,
|
||
};
|
||
|
||
let route = RouteDecision::Proxy {
|
||
target: ProxyTarget::Gitea,
|
||
body_limit: None,
|
||
};
|
||
assert_eq!(
|
||
gateway.classify_request(Some("git.genarrative.world"), "/"),
|
||
route
|
||
);
|
||
assert_eq!(
|
||
gateway.classify_request(Some("git.genarrative.world:443"), "/api/v1/repos"),
|
||
route
|
||
);
|
||
assert_eq!(
|
||
gateway.classify_request(Some("GIT.GENARRATIVE.WORLD"), "/assets/app.js"),
|
||
route
|
||
);
|
||
assert_eq!(
|
||
gateway.classify_request(Some("dev.genarrative.world"), "/api/assets/history"),
|
||
classify_path("/api/assets/history")
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn gitea_routes_bypass_maintenance_body_limit_and_protection() {
|
||
let route = RouteDecision::Proxy {
|
||
target: ProxyTarget::Gitea,
|
||
body_limit: None,
|
||
};
|
||
|
||
assert!(!route.applies_maintenance_gate());
|
||
assert!(!route.is_api_like());
|
||
assert_eq!(route.body_limit(), None);
|
||
assert_eq!(protection_class_for_route(&route, "/api/v1/repos"), None);
|
||
assert!(should_disable_accel_buffering(&route));
|
||
}
|
||
|
||
#[test]
|
||
fn maintenance_allows_internal_clients_to_bypass_all_gated_routes() {
|
||
for value in [
|
||
"127.0.0.1",
|
||
"10.35.0.50",
|
||
"172.16.0.50",
|
||
"192.168.35.50",
|
||
"169.254.0.50",
|
||
"::1",
|
||
"fd00::50",
|
||
"fe80::50",
|
||
] {
|
||
let internal_ip: IpAddr = value.parse().unwrap();
|
||
assert!(allows_internal_maintenance_bypass(Some(&internal_ip)));
|
||
}
|
||
for value in ["203.0.113.50", "2001:db8::50"] {
|
||
let public_ip: IpAddr = value.parse().unwrap();
|
||
assert!(!allows_internal_maintenance_bypass(Some(&public_ip)));
|
||
}
|
||
assert!(!allows_internal_maintenance_bypass(None));
|
||
}
|
||
|
||
#[test]
|
||
fn maintenance_only_allows_required_branding_assets() {
|
||
for path in [
|
||
"/branding/taonier-maintenance-page.png",
|
||
"/branding/taonier-product-ip.png",
|
||
] {
|
||
assert!(is_maintenance_page_asset(path));
|
||
}
|
||
for path in [
|
||
"/branding/other.png",
|
||
"/branding/taonier-maintenance-page.png/extra",
|
||
"/assets/app.js",
|
||
] {
|
||
assert!(!is_maintenance_page_asset(path));
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn normalizes_gateway_hosts_for_matching() {
|
||
assert_eq!(
|
||
normalize_gateway_host("Git.Genarrative.World:443"),
|
||
Some("git.genarrative.world".to_string())
|
||
);
|
||
assert_eq!(
|
||
normalize_gateway_host("git.genarrative.world."),
|
||
Some("git.genarrative.world".to_string())
|
||
);
|
||
assert!(normalize_gateway_host("git.genarrative.world/path").is_none());
|
||
assert!(normalize_gateway_host("[::1]:3000").is_none());
|
||
assert!(normalize_gateway_host("bad\r\nhost").is_none());
|
||
}
|
||
|
||
#[test]
|
||
fn validates_gitea_host_routing_config() {
|
||
let config = GatewayConfig {
|
||
gitea_hosts: vec!["git.genarrative.world".to_string()],
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
|
||
let config = GatewayConfig {
|
||
gitea_upstream: Some("127.0.0.1:3000".parse().unwrap()),
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
|
||
let config = GatewayConfig {
|
||
gitea_hosts: vec!["git.genarrative.world".to_string()],
|
||
gitea_upstream: Some("127.0.0.1:8082".parse().unwrap()),
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
|
||
let config = GatewayConfig {
|
||
gitea_hosts: vec!["git.genarrative.world".to_string()],
|
||
gitea_upstream: Some("127.0.0.1:3000".parse().unwrap()),
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_ok());
|
||
}
|
||
|
||
#[test]
|
||
fn http_redirect_mode_redirects_non_acme_routes_only() {
|
||
assert_eq!(
|
||
classify_http_redirect_path("/api/assets/history"),
|
||
RouteDecision::Local(LocalResponse::HttpToHttpsRedirect)
|
||
);
|
||
assert_eq!(
|
||
classify_http_redirect_path("/.well-known/acme-challenge/token"),
|
||
RouteDecision::Local(LocalResponse::Static {
|
||
root: StaticRoot::Acme,
|
||
mode: StaticMode::Exact,
|
||
})
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn identifies_streaming_payload_limit_errors() {
|
||
let error = Error::explain(HTTPStatus(413), PAYLOAD_TOO_LARGE_CONTEXT);
|
||
assert!(is_payload_too_large_error(&error));
|
||
|
||
let unrelated = Error::explain(HTTPStatus(413), "other");
|
||
assert!(!is_payload_too_large_error(&unrelated));
|
||
}
|
||
|
||
#[test]
|
||
fn maps_gateway_proxy_errors_to_stable_json_codes() {
|
||
assert!(gateway_proxy_error_body(502).contains("GATEWAY_UPSTREAM_ERROR"));
|
||
assert!(gateway_proxy_error_body(504).contains("GATEWAY_UPSTREAM_TIMEOUT"));
|
||
assert!(gateway_proxy_error_body(500).contains("GATEWAY_PROXY_ERROR"));
|
||
}
|
||
|
||
#[test]
|
||
fn formats_access_log_as_tab_separated_key_values() {
|
||
let line = format_access_log_line(&AccessLogRecord {
|
||
request_id: "req-1",
|
||
method: "GET",
|
||
path: "/api/test",
|
||
uri: "/api/test?a=1\t2",
|
||
host: "example.test",
|
||
client_ip: "127.0.0.1",
|
||
status: 200,
|
||
route: "Proxy",
|
||
proxy_target: "Api",
|
||
upstream: "127.0.0.1:8082",
|
||
content_length: None,
|
||
body_bytes_seen: 0,
|
||
protection_class: Some("api"),
|
||
protection_client: "127.0.0.1",
|
||
elapsed_ms: 12,
|
||
error: None,
|
||
});
|
||
|
||
assert!(line.contains("request_id=req-1\tmethod=GET\tpath=/api/test"));
|
||
assert!(line.contains("uri=/api/test?a=1\\t2"));
|
||
assert!(line.contains("content_length=-"));
|
||
assert!(line.ends_with("error=-\n"));
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_placeholder_or_short_shadow_probe_tokens() {
|
||
assert!(validate_shadow_probe_token("__GENARRATIVE_PINGORA_PROBE_TOKEN__").is_err());
|
||
assert!(validate_shadow_probe_token("changeme").is_err());
|
||
assert!(validate_shadow_probe_token("too-short").is_err());
|
||
assert!(validate_shadow_probe_token("0123456789abcdef").is_ok());
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_burst_without_rate_limit() {
|
||
let config = ProtectionClassConfig {
|
||
max_concurrent: 1,
|
||
rate_per_second: 0,
|
||
burst: 1,
|
||
};
|
||
assert!(config.validate("api").is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_access_log_directory_paths() {
|
||
let temp_dir = std::env::temp_dir();
|
||
assert!(validate_access_log_file(&temp_dir).is_err());
|
||
assert!(validate_access_log_file(&temp_dir.join("pingora-access.log")).is_ok());
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_incomplete_tls_and_redirect_config() {
|
||
let config = GatewayConfig {
|
||
tls_listen_addr: Some("127.0.0.1:18443".to_string()),
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
|
||
let config = GatewayConfig {
|
||
http_redirect_listen_addr: Some("127.0.0.1:18080".to_string()),
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
|
||
let config = GatewayConfig {
|
||
tls_listen_addr: Some("127.0.0.1:18081".to_string()),
|
||
tls_cert_file: Some(PathBuf::from("/missing/cert.pem")),
|
||
tls_key_file: Some(PathBuf::from("/missing/key.pem")),
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn validates_redirect_target_host_and_scheme() {
|
||
assert!(validate_redirect_target_scheme("https").is_ok());
|
||
assert!(validate_redirect_target_scheme("http").is_err());
|
||
assert!(is_valid_redirect_host("genarrative.world"));
|
||
assert!(is_valid_redirect_host("localhost:8443"));
|
||
assert!(is_valid_redirect_host("[::1]:8443"));
|
||
assert!(!is_valid_redirect_host("evil.test/path"));
|
||
assert!(!is_valid_redirect_host("evil.test\r\nx: y"));
|
||
}
|
||
|
||
#[test]
|
||
fn builds_stable_static_validators() {
|
||
let modified = UNIX_EPOCH + Duration::from_secs(1_700_000_123);
|
||
|
||
assert_eq!(
|
||
build_static_etag(4096, Some(modified)),
|
||
"W/\"1000-6553f17b\""
|
||
);
|
||
assert_eq!(build_static_etag(0, None), "W/\"0-0\"");
|
||
assert_eq!(
|
||
truncate_system_time_to_seconds(modified + Duration::from_millis(900)),
|
||
modified
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn matches_static_if_none_match_values() {
|
||
assert!(if_none_match_matches("*", "W/\"1000-6553f17b\""));
|
||
assert!(if_none_match_matches(
|
||
"\"other\", W/\"1000-6553f17b\"",
|
||
"W/\"1000-6553f17b\""
|
||
));
|
||
assert!(if_none_match_matches(
|
||
"\"1000-6553f17b\"",
|
||
"W/\"1000-6553f17b\""
|
||
));
|
||
assert!(!if_none_match_matches(
|
||
"\"other\", W/\"bad\"",
|
||
"W/\"1000-6553f17b\""
|
||
));
|
||
}
|
||
|
||
#[test]
|
||
fn parses_static_byte_ranges() {
|
||
assert_eq!(
|
||
parse_static_range("bytes=0-3", 10),
|
||
StaticRangeDecision::Partial { start: 0, end: 3 }
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("Bytes=5-", 10),
|
||
StaticRangeDecision::Partial { start: 5, end: 9 }
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=-4", 10),
|
||
StaticRangeDecision::Partial { start: 6, end: 9 }
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=-20", 10),
|
||
StaticRangeDecision::Partial { start: 0, end: 9 }
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=8-20", 10),
|
||
StaticRangeDecision::Partial { start: 8, end: 9 }
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=10-11", 10),
|
||
StaticRangeDecision::Unsatisfiable
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=4-3", 10),
|
||
StaticRangeDecision::Unsatisfiable
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=-0", 10),
|
||
StaticRangeDecision::Unsatisfiable
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=0-1,3-4", 10),
|
||
StaticRangeDecision::Full
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("items=0-1", 10),
|
||
StaticRangeDecision::Full
|
||
);
|
||
assert_eq!(
|
||
parse_static_range("bytes=0-1", 0),
|
||
StaticRangeDecision::Unsatisfiable
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn matches_static_if_range_values() {
|
||
let modified = UNIX_EPOCH + Duration::from_secs(1_700_000_123);
|
||
let metadata = StaticResponseMetadata {
|
||
len: 4096,
|
||
etag: build_static_etag(4096, Some(modified)),
|
||
last_modified: Some(modified),
|
||
};
|
||
|
||
assert!(static_if_range_value_matches(
|
||
&fmt_http_date(modified + Duration::from_secs(30)),
|
||
&metadata
|
||
));
|
||
assert!(!static_if_range_value_matches(
|
||
&fmt_http_date(modified - Duration::from_secs(30)),
|
||
&metadata
|
||
));
|
||
assert!(!static_if_range_value_matches(&metadata.etag, &metadata));
|
||
assert!(!static_if_range_value_matches(
|
||
"\"1000-6553f17b\"",
|
||
&metadata
|
||
));
|
||
}
|
||
|
||
#[test]
|
||
fn only_allows_static_read_methods() {
|
||
assert!(is_static_read_method(&Method::GET));
|
||
assert!(is_static_read_method(&Method::HEAD));
|
||
assert!(!is_static_read_method(&Method::POST));
|
||
assert!(!is_static_read_method(&Method::PUT));
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_invalid_gzip_level() {
|
||
let config = GatewayConfig {
|
||
gzip_level: 10,
|
||
..test_gateway_config()
|
||
};
|
||
|
||
assert!(config.validate().is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_invalid_gzip_min_length() {
|
||
let config = GatewayConfig {
|
||
gzip_min_length_bytes: 0,
|
||
..test_gateway_config()
|
||
};
|
||
|
||
assert!(config.validate().is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_static_cache_control_header_injection() {
|
||
let config = GatewayConfig {
|
||
static_cache: StaticCacheConfig {
|
||
html_cache_control: "no-cache\nX-Bad: yes".to_string(),
|
||
asset_cache_control: DEFAULT_ASSET_CACHE_CONTROL.to_string(),
|
||
static_cache_control: DEFAULT_STATIC_CACHE_CONTROL.to_string(),
|
||
},
|
||
..test_gateway_config()
|
||
};
|
||
|
||
assert!(config.validate().is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_unconfirmed_multi_instance_protection() {
|
||
let config = GatewayConfig {
|
||
instance_count: 0,
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
|
||
let config = GatewayConfig {
|
||
instance_count: 2,
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_err());
|
||
|
||
let config = GatewayConfig {
|
||
instance_count: 2,
|
||
shared_protection_confirmed: true,
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_ok());
|
||
|
||
let config = GatewayConfig {
|
||
instance_count: 2,
|
||
protection: ProtectionConfig {
|
||
enabled: false,
|
||
..test_gateway_config().protection
|
||
},
|
||
..test_gateway_config()
|
||
};
|
||
assert!(config.validate().is_ok());
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_zero_upstream_timeouts() {
|
||
assert!(
|
||
UpstreamTimeoutConfig {
|
||
connect_timeout_ms: 0,
|
||
default_read_timeout_seconds: DEFAULT_UPSTREAM_DEFAULT_READ_TIMEOUT_SECONDS,
|
||
api_read_timeout_seconds: DEFAULT_UPSTREAM_API_READ_TIMEOUT_SECONDS,
|
||
long_read_timeout_seconds: DEFAULT_UPSTREAM_LONG_READ_TIMEOUT_SECONDS,
|
||
write_timeout_seconds: DEFAULT_UPSTREAM_WRITE_TIMEOUT_SECONDS,
|
||
}
|
||
.validate()
|
||
.is_err()
|
||
);
|
||
assert!(
|
||
UpstreamTimeoutConfig {
|
||
connect_timeout_ms: DEFAULT_UPSTREAM_CONNECT_TIMEOUT_MS,
|
||
default_read_timeout_seconds: DEFAULT_UPSTREAM_DEFAULT_READ_TIMEOUT_SECONDS,
|
||
api_read_timeout_seconds: 0,
|
||
long_read_timeout_seconds: DEFAULT_UPSTREAM_LONG_READ_TIMEOUT_SECONDS,
|
||
write_timeout_seconds: DEFAULT_UPSTREAM_WRITE_TIMEOUT_SECONDS,
|
||
}
|
||
.validate()
|
||
.is_err()
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn chooses_long_upstream_timeout_for_spacetime_routes() {
|
||
let timeouts = UpstreamTimeoutConfig {
|
||
connect_timeout_ms: 100,
|
||
default_read_timeout_seconds: 60,
|
||
api_read_timeout_seconds: 10,
|
||
long_read_timeout_seconds: 3600,
|
||
write_timeout_seconds: 20,
|
||
};
|
||
let generic_api = classify_path("/api/assets/history");
|
||
let subscribe = classify_path("/v1/database/genarrative/subscribe");
|
||
|
||
assert_eq!(
|
||
timeouts.read_timeout_for_route(&generic_api, "/api/assets/history"),
|
||
Duration::from_secs(10)
|
||
);
|
||
assert_eq!(
|
||
timeouts.read_timeout_for_route(&subscribe, "/v1/database/genarrative/subscribe"),
|
||
Duration::from_secs(3600)
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_unsupported_compression_algorithms() {
|
||
assert!(CompressionConfig::parse("zstd".to_string()).is_err());
|
||
assert!(CompressionConfig::parse("gzip,br".to_string()).is_err());
|
||
}
|
||
|
||
#[test]
|
||
fn selects_configured_gateway_compression_algorithm() {
|
||
let gzip_only = CompressionRequestConfig {
|
||
enabled: true,
|
||
gzip: true,
|
||
min_length_bytes: DEFAULT_GZIP_MIN_LENGTH_BYTES,
|
||
gzip_level: DEFAULT_GZIP_LEVEL,
|
||
};
|
||
|
||
let mut req = RequestHeader::build(Method::GET, b"/assets/app.js", None).unwrap();
|
||
req.insert_header("Accept-Encoding", "br, gzip").unwrap();
|
||
normalize_accept_encoding_for_gateway_compression(&mut req, gzip_only).unwrap();
|
||
assert_eq!(
|
||
req.headers
|
||
.get("accept-encoding")
|
||
.and_then(|value| value.to_str().ok()),
|
||
Some("gzip")
|
||
);
|
||
|
||
let mut br_only = RequestHeader::build(Method::GET, b"/assets/app.js", None).unwrap();
|
||
br_only.insert_header("Accept-Encoding", "br").unwrap();
|
||
normalize_accept_encoding_for_gateway_compression(&mut br_only, gzip_only).unwrap();
|
||
assert!(br_only.headers.get("accept-encoding").is_none());
|
||
|
||
let mut gzip_disabled = RequestHeader::build(Method::GET, b"/assets/app.js", None).unwrap();
|
||
gzip_disabled
|
||
.insert_header("Accept-Encoding", "gzip")
|
||
.unwrap();
|
||
normalize_accept_encoding_for_gateway_compression(
|
||
&mut gzip_disabled,
|
||
CompressionRequestConfig {
|
||
enabled: false,
|
||
gzip: true,
|
||
min_length_bytes: DEFAULT_GZIP_MIN_LENGTH_BYTES,
|
||
gzip_level: DEFAULT_GZIP_LEVEL,
|
||
},
|
||
)
|
||
.unwrap();
|
||
assert!(gzip_disabled.headers.get("accept-encoding").is_none());
|
||
|
||
assert_eq!(parse_accepted_encodings("br, gzip;q=1"), ["br", "gzip"]);
|
||
assert_eq!(parse_accepted_encodings("br, gzip;q=0"), ["br"]);
|
||
}
|
||
|
||
#[test]
|
||
fn honors_configured_gzip_min_length() {
|
||
let mut small = ResponseHeader::build(200, Some(1023)).unwrap();
|
||
small.insert_header("Content-Length", "1023").unwrap();
|
||
small
|
||
.insert_header("Content-Type", "application/javascript")
|
||
.unwrap();
|
||
assert!(!should_allow_gateway_compression_for_response(
|
||
&small,
|
||
DEFAULT_GZIP_MIN_LENGTH_BYTES
|
||
));
|
||
|
||
let mut exact = ResponseHeader::build(200, Some(1024)).unwrap();
|
||
exact.insert_header("Content-Length", "1024").unwrap();
|
||
exact
|
||
.insert_header("Content-Type", "application/javascript")
|
||
.unwrap();
|
||
assert!(should_allow_gateway_compression_for_response(
|
||
&exact,
|
||
DEFAULT_GZIP_MIN_LENGTH_BYTES
|
||
));
|
||
|
||
let mut unknown_length = ResponseHeader::build(200, None).unwrap();
|
||
unknown_length
|
||
.insert_header("Content-Type", "application/javascript")
|
||
.unwrap();
|
||
assert!(should_allow_gateway_compression_for_response(
|
||
&unknown_length,
|
||
DEFAULT_GZIP_MIN_LENGTH_BYTES
|
||
));
|
||
|
||
let mut partial = ResponseHeader::build(206, Some(1024)).unwrap();
|
||
partial.insert_header("Content-Length", "1024").unwrap();
|
||
partial
|
||
.insert_header("Content-Range", "bytes 0-1023/4096")
|
||
.unwrap();
|
||
assert!(!should_allow_gateway_compression_for_response(
|
||
&partial,
|
||
DEFAULT_GZIP_MIN_LENGTH_BYTES
|
||
));
|
||
|
||
let mut not_modified = ResponseHeader::build(304, Some(0)).unwrap();
|
||
not_modified.insert_header("Content-Length", "0").unwrap();
|
||
assert!(!should_allow_gateway_compression_for_response(
|
||
¬_modified,
|
||
DEFAULT_GZIP_MIN_LENGTH_BYTES
|
||
));
|
||
}
|
||
|
||
#[test]
|
||
fn disables_accel_buffering_for_api_proxy_only() {
|
||
assert!(should_disable_accel_buffering(&classify_path(
|
||
"/api/assets/history"
|
||
)));
|
||
assert!(!should_disable_accel_buffering(&classify_path(
|
||
"/v1/database/genarrative/subscribe"
|
||
)));
|
||
assert!(!should_disable_accel_buffering(&classify_path(
|
||
"/assets/app.js"
|
||
)));
|
||
}
|
||
|
||
#[test]
|
||
fn maps_proxy_routes_to_protection_classes() {
|
||
let cases = [
|
||
("/admin/api/users", Some(ProtectionClass::AdminApi)),
|
||
("/api/assets/history", Some(ProtectionClass::Api)),
|
||
(
|
||
"/v1/database/genarrative/subscribe",
|
||
Some(ProtectionClass::Spacetime),
|
||
),
|
||
("/assets/app.js", None),
|
||
];
|
||
|
||
for (path, expected) in cases {
|
||
assert_eq!(
|
||
protection_class_for_route(&classify_path(path), path),
|
||
expected,
|
||
"path: {path}"
|
||
);
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn protection_rejects_when_concurrency_is_exhausted() {
|
||
let config = ProtectionConfig {
|
||
enabled: true,
|
||
api: ProtectionClassConfig {
|
||
max_concurrent: 1,
|
||
rate_per_second: 0,
|
||
burst: 0,
|
||
},
|
||
admin_api: ProtectionClassConfig {
|
||
max_concurrent: 0,
|
||
rate_per_second: 0,
|
||
burst: 0,
|
||
},
|
||
spacetime: ProtectionClassConfig {
|
||
max_concurrent: 0,
|
||
rate_per_second: 0,
|
||
burst: 0,
|
||
},
|
||
};
|
||
let protection = Arc::new(GatewayProtection::new(config));
|
||
let first = protection
|
||
.try_acquire(ProtectionClass::Api, "127.0.0.1".to_string())
|
||
.expect("first request should be accepted");
|
||
|
||
assert!(matches!(
|
||
protection.try_acquire(ProtectionClass::Api, "127.0.0.1".to_string()),
|
||
Err(RejectReason::Concurrent)
|
||
));
|
||
|
||
if let Some(key) = first {
|
||
protection.release(&key);
|
||
}
|
||
assert!(
|
||
protection
|
||
.try_acquire(ProtectionClass::Api, "127.0.0.1".to_string())
|
||
.expect("request should be accepted after release")
|
||
.is_some()
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn protection_rejects_when_rate_bucket_is_empty() {
|
||
let config = ProtectionConfig {
|
||
enabled: true,
|
||
api: ProtectionClassConfig {
|
||
max_concurrent: 0,
|
||
rate_per_second: 1,
|
||
burst: 0,
|
||
},
|
||
admin_api: ProtectionClassConfig {
|
||
max_concurrent: 0,
|
||
rate_per_second: 0,
|
||
burst: 0,
|
||
},
|
||
spacetime: ProtectionClassConfig {
|
||
max_concurrent: 0,
|
||
rate_per_second: 0,
|
||
burst: 0,
|
||
},
|
||
};
|
||
let protection = Arc::new(GatewayProtection::new(config));
|
||
let first = protection
|
||
.try_acquire(ProtectionClass::Api, "127.0.0.1".to_string())
|
||
.expect("first request should pass");
|
||
if let Some(key) = first {
|
||
protection.release(&key);
|
||
}
|
||
|
||
assert!(matches!(
|
||
protection.try_acquire(ProtectionClass::Api, "127.0.0.1".to_string()),
|
||
Err(RejectReason::Rate)
|
||
));
|
||
assert!(
|
||
protection
|
||
.try_acquire(ProtectionClass::Api, "127.0.0.2".to_string())
|
||
.expect("different client should have an independent bucket")
|
||
.is_some()
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn rejects_path_traversal_for_static_files() {
|
||
assert!(sanitize_static_path(Path::new("/srv/web"), "/assets/app.js").is_some());
|
||
assert!(sanitize_static_path(Path::new("/srv/web"), "/assets/../secret").is_none());
|
||
assert!(sanitize_static_path(Path::new("/srv/web"), "/assets/%2e%2e/secret").is_none(),);
|
||
assert!(sanitize_static_path(Path::new("/srv/web"), "/assets%2fsecret").is_none());
|
||
assert!(sanitize_static_path(Path::new("/srv/web"), "/assets/%GG").is_none());
|
||
}
|
||
|
||
#[test]
|
||
fn decodes_percent_encoded_static_path_segments() {
|
||
assert_eq!(
|
||
sanitize_static_path(
|
||
Path::new("/srv/web"),
|
||
"/Icons/Admurin%27s%20Pixel%20Items/General/Singles/499_Iron_Gear.png",
|
||
),
|
||
Some(PathBuf::from(
|
||
"/srv/web/Icons/Admurin's Pixel Items/General/Singles/499_Iron_Gear.png",
|
||
)),
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn classifies_static_cache_policy() {
|
||
let root = StaticRoot::Web;
|
||
assert_eq!(
|
||
classify_static_cache_kind(root, "/", Path::new("/srv/web/index.html")),
|
||
StaticCacheKind::Html
|
||
);
|
||
assert_eq!(
|
||
classify_static_cache_kind(
|
||
root,
|
||
"/assets/index-B4dmVw0r.js",
|
||
Path::new("/srv/web/assets/index-B4dmVw0r.js"),
|
||
),
|
||
StaticCacheKind::FingerprintedAsset
|
||
);
|
||
assert_eq!(
|
||
classify_static_cache_kind(
|
||
root,
|
||
"/admin/assets/admin-Ca8f3012.css",
|
||
Path::new("/srv/web/admin/assets/admin-Ca8f3012.css"),
|
||
),
|
||
StaticCacheKind::FingerprintedAsset
|
||
);
|
||
assert_eq!(
|
||
classify_static_cache_kind(root, "/assets/app.js", Path::new("/srv/web/assets/app.js")),
|
||
StaticCacheKind::Other
|
||
);
|
||
assert_eq!(
|
||
classify_static_cache_kind(
|
||
StaticRoot::Acme,
|
||
"/.well-known/acme-challenge/token-12345678",
|
||
Path::new("/var/www/html/.well-known/acme-challenge/token-12345678"),
|
||
),
|
||
StaticCacheKind::Other
|
||
);
|
||
}
|
||
}
|