Files
Genarrative/server-rs/crates/pingora-gateway/src/main.rs
T
kdletters 47f00dae4d
Project CI / Repository checks (push) Successful in 1m4s
Project CI / Frontend tests (push) Successful in 2m2s
Project CI / Native shell tests (push) Successful in 2m28s
Project CI / Backend tests (push) Successful in 3m3s
修复维护页品牌图片加载
Nginx 精确放行维护页背景图与品牌图标
Pingora 同步维护资产白名单
补充维护模式网关回归测试
更新生产运维与共享决策文档
2026-07-27 21:57:30 +08:00

3996 lines
128 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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(&not_found_page).await {
let body = Bytes::from(tokio::fs::read(&not_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(
&not_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
);
}
}