refactor(svc): tidy the rust services and build tooling (#2735)

This commit is contained in:
Hampus
2026-09-13 17:38:32 +02:00
committed by GitHub
parent 33737e0f79
commit 6af33c7188
41 changed files with 863 additions and 1304 deletions
Generated
+2
View File
@@ -2011,6 +2011,7 @@ dependencies = [
"axum", "axum",
"base64", "base64",
"fluxer_common", "fluxer_common",
"futures-util",
"hex", "hex",
"rand 0.10.1", "rand 0.10.1",
"reqwest", "reqwest",
@@ -2039,6 +2040,7 @@ dependencies = [
"reqwest", "reqwest",
"serde_json", "serde_json",
"sha2 0.11.0", "sha2 0.11.0",
"tempfile",
"thiserror", "thiserror",
"time", "time",
"tracing", "tracing",
+1
View File
@@ -10,6 +10,7 @@ anyhow = "1.0.104"
axum = { version = "0.8.9", features = ["macros"] } axum = { version = "0.8.9", features = ["macros"] }
base64 = "0.22" base64 = "0.22"
fluxer_common = { path = "../fluxer_common" } fluxer_common = { path = "../fluxer_common" }
futures-util = { version = "0.3.32", default-features = false, features = ["std"] }
hex = "0.4" hex = "0.4"
rand = "0.10" rand = "0.10"
reqwest = { version = "0.13.4", default-features = false, features = ["json", "rustls", "stream", "gzip", "brotli", "deflate"] } reqwest = { version = "0.13.4", default-features = false, features = ["json", "rustls", "stream", "gzip", "brotli", "deflate"] }
+13 -20
View File
@@ -132,7 +132,7 @@ pub fn inject_bootstrap(
let media = media_endpoint.trim_end_matches('/'); let media = media_endpoint.trim_end_matches('/');
let nonced = html.replace("{{CSP_NONCE_PLACEHOLDER}}", nonce); let nonced = html.replace("{{CSP_NONCE_PLACEHOLDER}}", nonce);
let nonced = apply_static_preconnect(&nonced, static_cdn); let nonced = apply_static_preconnect(nonced, static_cdn);
let nonced = nonced.replace("{{STATIC_CDN_ENDPOINT}}", static_cdn); let nonced = nonced.replace("{{STATIC_CDN_ENDPOINT}}", static_cdn);
let nonced = apply_media_preconnect(&nonced, media, static_cdn); let nonced = apply_media_preconnect(&nonced, media, static_cdn);
@@ -143,20 +143,14 @@ pub fn inject_bootstrap(
return nonced.replace("{{FLUXER_BOOTSTRAP}}", script_tag); return nonced.replace("{{FLUXER_BOOTSTRAP}}", script_tag);
} }
if let Some(pos) = nonced.find("<head>") { let insert_at = nonced
let insert_at = pos + "<head>".len(); .find("<head>")
let mut result = String::with_capacity(nonced.len() + script_tag.len() + 3); .map(|pos| pos + "<head>".len())
result.push_str(&nonced[..insert_at]); .or_else(|| {
result.push_str("\n\t\t"); let pos = nonced.find("<head ")?;
result.push_str(script_tag); nonced[pos..].find('>').map(|close| pos + close + 1)
result.push_str(&nonced[insert_at..]); });
return result; if let Some(insert_at) = insert_at {
}
if let Some(pos) = nonced.find("<head ")
&& let Some(close) = nonced[pos..].find('>')
{
let insert_at = pos + close + 1;
let mut result = String::with_capacity(nonced.len() + script_tag.len() + 3); let mut result = String::with_capacity(nonced.len() + script_tag.len() + 3);
result.push_str(&nonced[..insert_at]); result.push_str(&nonced[..insert_at]);
result.push_str("\n\t\t"); result.push_str("\n\t\t");
@@ -168,15 +162,14 @@ pub fn inject_bootstrap(
nonced nonced
} }
fn apply_static_preconnect(html: &str, static_cdn: &str) -> String { fn apply_static_preconnect(mut html: String, static_cdn: &str) -> String {
if !static_cdn.is_empty() { if !static_cdn.is_empty() {
return html.to_owned(); return html;
} }
let mut stripped = html.to_owned();
for tag in STATIC_PRECONNECT_TAGS { for tag in STATIC_PRECONNECT_TAGS {
stripped = stripped.replace(&format!("{tag}\n"), "").replace(tag, ""); html = html.replace(&format!("{tag}\n"), "").replace(tag, "");
} }
stripped html
} }
fn apply_media_preconnect(html: &str, media: &str, static_cdn: &str) -> String { fn apply_media_preconnect(html: &str, media: &str, static_cdn: &str) -> String {
+17 -36
View File
@@ -309,30 +309,20 @@ impl fmt::Display for CspReportUri {
} }
} }
fn warn_invalid(error: InvalidAppProxyEnvironmentError) { fn warn_invalid(error: &InvalidAppProxyEnvironmentError) {
tracing::warn!(%error, "ignoring invalid app proxy environment value"); tracing::warn!(%error, "ignoring invalid app proxy environment value");
} }
fn parse_optional_http_url(name: &'static str, value: Option<String>) -> Option<HttpUrl> { fn parse_optional_http_url(name: &'static str, value: Option<String>) -> Option<HttpUrl> {
let value = value?; let value = value?;
match HttpUrl::parse(name, &value) { HttpUrl::parse(name, &value).inspect_err(warn_invalid).ok()
Ok(url) => Some(url),
Err(error) => {
warn_invalid(error);
None
}
}
} }
fn parse_optional_http_endpoint(name: &'static str, value: Option<String>) -> Option<HttpEndpoint> { fn parse_optional_http_endpoint(name: &'static str, value: Option<String>) -> Option<HttpEndpoint> {
let value = value?; let value = value?;
match HttpEndpoint::parse(name, &value) { HttpEndpoint::parse(name, &value)
Ok(endpoint) => Some(endpoint), .inspect_err(warn_invalid)
Err(error) => { .ok()
warn_invalid(error);
None
}
}
} }
fn parse_env_or_warn<T: std::str::FromStr>(name: &str, raw: &str, default: T) -> T { fn parse_env_or_warn<T: std::str::FromStr>(name: &str, raw: &str, default: T) -> T {
@@ -434,25 +424,19 @@ fn read_csp_sources(name: &'static str) -> Vec<CspSource> {
.split([',', ' ', '\t', '\n']) .split([',', ' ', '\t', '\n'])
.map(str::trim) .map(str::trim)
.filter(|source| !source.is_empty()) .filter(|source| !source.is_empty())
.filter_map(|source| match CspSource::parse(name, source) { .filter_map(|source| {
Ok(source) => Some(source), CspSource::parse(name, source)
Err(error) => { .inspect_err(warn_invalid)
warn_invalid(error); .ok()
None
}
}) })
.collect() .collect()
} }
fn read_csp_report_uri(name: &'static str) -> Option<CspReportUri> { fn read_csp_report_uri(name: &'static str) -> Option<CspReportUri> {
let value = cfg::non_empty_env(name)?; let value = cfg::non_empty_env(name)?;
match CspReportUri::parse(name, &value) { CspReportUri::parse(name, &value)
Ok(report_uri) => Some(report_uri), .inspect_err(warn_invalid)
Err(error) => { .ok()
warn_invalid(error);
None
}
}
} }
impl AppProxyConfig { impl AppProxyConfig {
@@ -474,13 +458,10 @@ impl AppProxyConfig {
); );
let s3_uploads_bucket = cfg::read_env("FLUXER_S3_BUCKET_UPLOADS", "fluxer-uploads"); let s3_uploads_bucket = cfg::read_env("FLUXER_S3_BUCKET_UPLOADS", "fluxer-uploads");
let s3_uploads_endpoint = s3_public_endpoint.as_ref().and_then(|endpoint| { let s3_uploads_endpoint = s3_public_endpoint.as_ref().and_then(|endpoint| {
match endpoint.with_host_prefix("FLUXER_S3_BUCKET_UPLOADS", s3_uploads_bucket.trim()) { endpoint
Ok(endpoint) => Some(endpoint), .with_host_prefix("FLUXER_S3_BUCKET_UPLOADS", s3_uploads_bucket.trim())
Err(error) => { .inspect_err(warn_invalid)
warn_invalid(error); .ok()
None
}
}
}); });
Self { Self {
@@ -499,7 +480,7 @@ impl AppProxyConfig {
"FLUXER_STATIC_CDN_ENDPOINT", "FLUXER_STATIC_CDN_ENDPOINT",
cfg::non_empty_env("FLUXER_STATIC_CDN_ENDPOINT"), cfg::non_empty_env("FLUXER_STATIC_CDN_ENDPOINT"),
), ),
s3_public_endpoint: s3_public_endpoint.clone(), s3_public_endpoint,
s3_uploads_endpoint, s3_uploads_endpoint,
discovery_upstream_url: resolve_discovery_upstream_url_from_env(), discovery_upstream_url: resolve_discovery_upstream_url_from_env(),
discovery_refresh_interval_ms: parse_env_or_warn( discovery_refresh_interval_ms: parse_env_or_warn(
+5 -9
View File
@@ -287,15 +287,11 @@ fn extend_runtime_s3_sources(target: &mut Vec<String>, runtime_sources: &Runtime
} }
fn extend_from(target: &mut Vec<String>, extra: &[CspSource], defaults: &[&str]) { fn extend_from(target: &mut Vec<String>, extra: &[CspSource], defaults: &[&str]) {
for source in defaults { for source in defaults
if target.iter().any(|existing| existing == source) { .iter()
continue; .copied()
} .chain(extra.iter().map(CspSource::as_str))
target.push((*source).to_owned()); {
}
for source in extra {
let source = source.as_str();
if target.iter().any(|existing| existing == source) { if target.iter().any(|existing| existing == source) {
continue; continue;
} }
@@ -2,47 +2,33 @@
use axum::{ use axum::{
Json, Json,
http::{HeaderValue, header}, http::header,
response::{IntoResponse, Response}, response::{IntoResponse, Response},
}; };
use serde_json::json; use serde_json::json;
pub async fn assetlinks() -> Response { pub async fn assetlinks() -> Response {
let body = json!([ let body = ["com.fluxer", "com.fluxer.canary"].map(|package_name| {
{ json!({
"relation": [ "relation": [
"delegate_permission/common.handle_all_urls", "delegate_permission/common.handle_all_urls",
"delegate_permission/common.get_login_creds" "delegate_permission/common.get_login_creds"
], ],
"target": { "target": {
"namespace": "android_app", "namespace": "android_app",
"package_name": "com.fluxer", "package_name": package_name,
"sha256_cert_fingerprints": [ "sha256_cert_fingerprints": [
"91:E4:98:E1:B8:A6:C8:BA:99:41:5E:DB:29:78:29:6B:6C:58:BA:A5:E2:D2:A6:49:CE:C6:2D:A7:A8:29:C7:BC" "91:E4:98:E1:B8:A6:C8:BA:99:41:5E:DB:29:78:29:6B:6C:58:BA:A5:E2:D2:A6:49:CE:C6:2D:A7:A8:29:C7:BC"
] ]
} }
}, })
{ });
"relation": [
"delegate_permission/common.handle_all_urls",
"delegate_permission/common.get_login_creds"
],
"target": {
"namespace": "android_app",
"package_name": "com.fluxer.canary",
"sha256_cert_fingerprints": [
"91:E4:98:E1:B8:A6:C8:BA:99:41:5E:DB:29:78:29:6B:6C:58:BA:A5:E2:D2:A6:49:CE:C6:2D:A7:A8:29:C7:BC"
]
}
}
]);
let mut response = Json(body).into_response(); (
response.headers_mut().insert( [(header::CACHE_CONTROL, "public, max-age=1800")],
header::CACHE_CONTROL, Json(body),
HeaderValue::from_static("public, max-age=1800"), )
); .into_response()
response
} }
#[cfg(test)] #[cfg(test)]
@@ -2,7 +2,7 @@
use axum::{ use axum::{
Json, Json,
http::{HeaderValue, header}, http::header,
response::{IntoResponse, Response}, response::{IntoResponse, Response},
}; };
use serde_json::json; use serde_json::json;
@@ -19,10 +19,9 @@ pub async fn apple_app_site_association() -> Response {
} }
}); });
let mut response = Json(body).into_response(); (
response.headers_mut().insert( [(header::CACHE_CONTROL, "public, max-age=1800")],
header::CACHE_CONTROL, Json(body),
HeaderValue::from_static("public, max-age=1800"), )
); .into_response()
response
} }
+43 -25
View File
@@ -9,9 +9,7 @@ use axum::{
}; };
use std::path::{Path as FsPath, PathBuf}; use std::path::{Path as FsPath, PathBuf};
use std::time::Duration; use std::time::Duration;
use tokio::io::{AsyncWriteExt, DuplexStream};
use tokio::sync::{OwnedSemaphorePermit, TryAcquireError}; use tokio::sync::{OwnedSemaphorePermit, TryAcquireError};
use tokio_util::io::ReaderStream;
use super::file_stream::stream_file; use super::file_stream::stream_file;
use super::spa_static::{CORS_ALLOW_ANY_VALUE, asset_cache_control, guess_mime, is_font_mime}; use super::spa_static::{CORS_ALLOW_ANY_VALUE, asset_cache_control, guess_mime, is_font_mime};
@@ -19,7 +17,6 @@ use super::spa_static::{CORS_ALLOW_ANY_VALUE, asset_cache_control, guess_mime, i
const ASSET_REQUEST_TIMEOUT: Duration = Duration::from_secs(15); const ASSET_REQUEST_TIMEOUT: Duration = Duration::from_secs(15);
const PRECOMPRESSED_VARIANTS: &[(&str, &str)] = &[("br", "br"), ("gzip", "gz")]; const PRECOMPRESSED_VARIANTS: &[(&str, &str)] = &[("br", "br"), ("gzip", "gz")];
const MAX_ASSET_SIZE_BYTES: u64 = 100 * 1024 * 1024; const MAX_ASSET_SIZE_BYTES: u64 = 100 * 1024 * 1024;
const UPSTREAM_ASSET_PUMP_BUFFER_BYTES: usize = 64 * 1024;
const UPSTREAM_FAILURE_CACHE_CONTROL: &str = "no-store"; const UPSTREAM_FAILURE_CACHE_CONTROL: &str = "no-store";
const UPSTREAM_FAILURE_STRIPPED_HEADERS: &[&str] = &[ const UPSTREAM_FAILURE_STRIPPED_HEADERS: &[&str] = &[
"cdn-cache-control", "cdn-cache-control",
@@ -148,38 +145,59 @@ pub async fn proxy_assets(
response_headers.insert(header::CONTENT_SECURITY_POLICY, state.csp.asset_header()); response_headers.insert(header::CONTENT_SECURITY_POLICY, state.csp.asset_header());
response_headers.remove("content-security-policy-report-only"); response_headers.remove("content-security-policy-report-only");
let body = Body::from_stream(upstream_asset_body(upstream_response, upstream_slot)); let body = upstream_asset_body(upstream_response, upstream_slot);
let mut response = Response::new(body); let mut response = Response::new(body);
*response.status_mut() = status; *response.status_mut() = status;
*response.headers_mut() = response_headers; *response.headers_mut() = response_headers;
response response
} }
struct UpstreamAssetReadState {
response: reqwest::Response,
_permit: OwnedSemaphorePermit,
remaining_bytes: u64,
}
fn upstream_asset_body( fn upstream_asset_body(
mut upstream_response: reqwest::Response, upstream_response: reqwest::Response,
upstream_slot: OwnedSemaphorePermit, upstream_slot: OwnedSemaphorePermit,
) -> ReaderStream<DuplexStream> { ) -> Body {
let (writer, reader) = tokio::io::duplex(UPSTREAM_ASSET_PUMP_BUFFER_BYTES); let state = UpstreamAssetReadState {
tokio::spawn(async move { response: upstream_response,
let _upstream_slot = upstream_slot; _permit: upstream_slot,
let mut writer = writer; remaining_bytes: MAX_ASSET_SIZE_BYTES,
loop { };
match upstream_response.chunk().await { Body::from_stream(futures_util::stream::try_unfold(
Ok(Some(chunk)) => { state,
if writer.write_all(&chunk).await.is_err() { |mut state| async move {
return; let Some(chunk) = state
} .response
} .chunk()
Ok(None) => break, .await
Err(err) => { .inspect_err(|err| {
tracing::warn!(%err, "upstream asset body ended early"); tracing::warn!(%err, "upstream asset body ended early");
return; })
} .map_err(axum::Error::new)?
else {
return Ok(None);
};
let chunk_bytes = chunk.len() as u64;
if chunk_bytes > state.remaining_bytes {
tracing::warn!(
maximum_bytes = MAX_ASSET_SIZE_BYTES,
remaining_bytes = state.remaining_bytes,
chunk_bytes,
"upstream asset body exceeds size cap"
);
return Err(axum::Error::new(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("upstream asset body exceeds {MAX_ASSET_SIZE_BYTES} bytes"),
)));
} }
} state.remaining_bytes -= chunk_bytes;
let _ = writer.shutdown().await; Ok(Some((chunk, state)))
}); },
ReaderStream::new(reader) ))
} }
pub(super) async fn serve_local_asset( pub(super) async fn serve_local_asset(
+7 -13
View File
@@ -13,7 +13,7 @@ use axum::{
Router, Router,
extract::Request, extract::Request,
http::{HeaderName, HeaderValue, header}, http::{HeaderName, HeaderValue, header},
middleware::{Next, from_fn, from_fn_with_state}, middleware::{Next, from_fn},
response::{IntoResponse, Response}, response::{IntoResponse, Response},
routing::get, routing::get,
}; };
@@ -56,10 +56,7 @@ pub fn build_router(state: AppState) -> Router {
.fallback(get(spa_index::spa_catch_all)) .fallback(get(spa_index::spa_catch_all))
.layer(from_fn(request_id_middleware)) .layer(from_fn(request_id_middleware))
.layer(from_fn(cache_headers_middleware)) .layer(from_fn(cache_headers_middleware))
.layer(from_fn_with_state( .layer(from_fn(security_headers_middleware))
state.clone(),
security_headers_middleware,
))
.layer( .layer(
CompressionLayer::new() CompressionLayer::new()
.compress_when(DefaultPredicate::new().and(NotForContentType::const_new("font/"))), .compress_when(DefaultPredicate::new().and(NotForContentType::const_new("font/"))),
@@ -68,14 +65,13 @@ pub fn build_router(state: AppState) -> Router {
.with_state(state) .with_state(state)
} }
async fn security_headers_middleware( async fn security_headers_middleware(request: Request, next: Next) -> Response {
axum::extract::State(_state): axum::extract::State<AppState>,
request: Request,
next: Next,
) -> Response {
let mut response = next.run(request).await; let mut response = next.run(request).await;
let headers = response.headers_mut(); set_security_headers(response.headers_mut());
response
}
fn set_security_headers(headers: &mut axum::http::HeaderMap) {
set_static_header( set_static_header(
headers, headers,
header::STRICT_TRANSPORT_SECURITY, header::STRICT_TRANSPORT_SECURITY,
@@ -89,8 +85,6 @@ async fn security_headers_middleware(
HeaderName::from_static("permissions-policy"), HeaderName::from_static("permissions-policy"),
PERMISSIONS_POLICY_VALUE, PERMISSIONS_POLICY_VALUE,
); );
response
} }
async fn cache_headers_middleware(request: Request, next: Next) -> Response { async fn cache_headers_middleware(request: Request, next: Next) -> Response {
+1 -18
View File
@@ -410,19 +410,7 @@ fn build_spa_response(
} else { } else {
headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-cache")); headers.insert(header::CACHE_CONTROL, HeaderValue::from_static("no-cache"));
} }
headers.insert( super::set_security_headers(headers);
header::STRICT_TRANSPORT_SECURITY,
HeaderValue::from_static("max-age=31536000; includeSubDomains; preload"),
);
headers.insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
);
headers.insert(header::X_FRAME_OPTIONS, HeaderValue::from_static("DENY"));
headers.insert(
header::REFERRER_POLICY,
HeaderValue::from_static("strict-origin-when-cross-origin"),
);
headers.insert( headers.insert(
axum::http::HeaderName::from_static("accept-ch"), axum::http::HeaderName::from_static("accept-ch"),
HeaderValue::from_static(ACCEPT_CH_VALUE), HeaderValue::from_static(ACCEPT_CH_VALUE),
@@ -431,11 +419,6 @@ fn build_spa_response(
axum::http::HeaderName::from_static("critical-ch"), axum::http::HeaderName::from_static("critical-ch"),
HeaderValue::from_static(CRITICAL_CH_VALUE), HeaderValue::from_static(CRITICAL_CH_VALUE),
); );
headers.insert(
axum::http::HeaderName::from_static("permissions-policy"),
HeaderValue::from_static(super::PERMISSIONS_POLICY_VALUE),
);
#[cfg(feature = "time-freeze")] #[cfg(feature = "time-freeze")]
{ {
if let Some(tf) = time_freeze_header if let Some(tf) = time_freeze_header
+1
View File
@@ -13,6 +13,7 @@ maxminddb = "0.28.1"
moka = { version = "0.12.15", features = ["sync"] } moka = { version = "0.12.15", features = ["sync"] }
reqwest = { version = "0.13.4", default-features = false, features = ["blocking", "rustls"] } reqwest = { version = "0.13.4", default-features = false, features = ["blocking", "rustls"] }
serde_json = "1.0.150" serde_json = "1.0.150"
tempfile = "3.27.0"
time = { version = "0.3.47", features = ["formatting", "macros", "parsing"] } time = { version = "0.3.47", features = ["formatting", "macros", "parsing"] }
tracing = "0.1.44" tracing = "0.1.44"
base64 = "0.22" base64 = "0.22"
+6 -6
View File
@@ -203,12 +203,12 @@ fn is_default_port(scheme: &str, port: u16) -> bool {
) )
} }
fn strip_trailing_dot(host: &str) -> &str {
host.strip_suffix('.').unwrap_or(host)
}
fn canonicalize_domain(value: &str) -> String { fn canonicalize_domain(value: &str) -> String {
strip_trailing_dot(value.trim().to_lowercase().as_str()).to_owned() let mut domain = value.trim().to_lowercase();
if domain.ends_with('.') {
domain.pop();
}
domain
} }
fn default_port(scheme: &str) -> u16 { fn default_port(scheme: &str) -> u16 {
@@ -285,10 +285,10 @@ where
} }
pub fn normalize_public_endpoint(url: &str, base_domain: &str, public_port: Option<u16>) -> String { pub fn normalize_public_endpoint(url: &str, base_domain: &str, public_port: Option<u16>) -> String {
let domain = canonicalize_domain(base_domain);
let Some(port) = public_port.filter(|port| *port != 0) else { let Some(port) = public_port.filter(|port| *port != 0) else {
return url.to_owned(); return url.to_owned();
}; };
let domain = canonicalize_domain(base_domain);
if domain.is_empty() { if domain.is_empty() {
return url.to_owned(); return url.to_owned();
} }
+3 -20
View File
@@ -19,7 +19,6 @@ use std::{
path::Path, path::Path,
time::{Duration, SystemTime}, time::{Duration, SystemTime},
}; };
use time::OffsetDateTime;
const DEFAULT_COUNTRY_CODE: &str = "US"; const DEFAULT_COUNTRY_CODE: &str = "US";
const GEOIP_CACHE_TTL: Duration = Duration::from_secs(10 * 60); const GEOIP_CACHE_TTL: Duration = Duration::from_secs(10 * 60);
@@ -242,7 +241,6 @@ fn download_s3_object(
); );
}; };
fs::create_dir_all(parent)?; fs::create_dir_all(parent)?;
let temp_path = temporary_download_path(destination);
let request = signed_s3_get_request(config, bucket, key)?; let request = signed_s3_get_request(config, bucket, key)?;
let mut response = reqwest::blocking::Client::new() let mut response = reqwest::blocking::Client::new()
.get(request.url) .get(request.url)
@@ -253,29 +251,14 @@ fn download_s3_object(
})?; })?;
if !response.status().is_success() { if !response.status().is_success() {
let status = response.status(); let status = response.status();
let _ = fs::remove_file(&temp_path);
anyhow::bail!("failed to download GeoIP database from s3://{bucket}/{key}: HTTP {status}"); anyhow::bail!("failed to download GeoIP database from s3://{bucket}/{key}: HTTP {status}");
} }
{ let mut file = tempfile::Builder::new().make_in(parent, |path| File::create_new(path))?;
let mut file = File::create(&temp_path)?; io::copy(&mut response, &mut file)?;
io::copy(&mut response, &mut file)?; file.persist(destination).map_err(|err| err.error)?;
}
fs::rename(&temp_path, destination).inspect_err(|_err| {
let _ = fs::remove_file(&temp_path);
})?;
Ok(()) Ok(())
} }
fn temporary_download_path(destination: &Path) -> std::path::PathBuf {
let pid = std::process::id();
let now = OffsetDateTime::now_utc().unix_timestamp_nanos();
let file_name = destination
.file_name()
.and_then(|value| value.to_str())
.unwrap_or("geoip.mmdb");
destination.with_file_name(format!("{file_name}.tmp-{pid}-{now}"))
}
struct SignedS3Request { struct SignedS3Request {
url: reqwest::Url, url: reqwest::Url,
headers: reqwest::header::HeaderMap, headers: reqwest::header::HeaderMap,
+3 -8
View File
@@ -66,20 +66,15 @@ where
pub async fn generate(&mut self) -> anyhow::Result<u64> { pub async fn generate(&mut self) -> anyhow::Result<u64> {
let mut timestamp = self.current_relative_timestamp()?; let mut timestamp = self.current_relative_timestamp()?;
if let Some(last_timestamp) = self.last_timestamp { match self.last_timestamp {
if timestamp < last_timestamp { Some(last_timestamp) if timestamp <= last_timestamp => {
timestamp = last_timestamp; timestamp = last_timestamp;
}
if timestamp == last_timestamp {
self.sequence = (self.sequence + 1) & MAX_SEQUENCE; self.sequence = (self.sequence + 1) & MAX_SEQUENCE;
if self.sequence == 0 { if self.sequence == 0 {
timestamp = self.wait_until_next_timestamp(last_timestamp).await?; timestamp = self.wait_until_next_timestamp(last_timestamp).await?;
} }
} else {
self.sequence = 0;
} }
} else { _ => self.sequence = 0,
self.sequence = 0;
} }
self.last_timestamp = Some(timestamp); self.last_timestamp = Some(timestamp);
Ok(create_snowflake_from_relative_timestamp( Ok(create_snowflake_from_relative_timestamp(
+140 -89
View File
@@ -4,10 +4,10 @@ use crate::types::{SERVICE_NAME, SnowflakeRequest, SnowflakeResponse};
use fluxer_svc::config::ServiceConfig; use fluxer_svc::config::ServiceConfig;
use fluxer_svc::hash_ring::HashRing; use fluxer_svc::hash_ring::HashRing;
use fluxer_svc::metrics::ServiceMetrics; use fluxer_svc::metrics::ServiceMetrics;
use fluxer_svc::transport::{NatsTransport, TransportMessage, TransportSubscriber}; use fluxer_svc::transport::{NatsMessage, NatsTransport, TransportMessage, TransportSubscriber};
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU32, Ordering}; use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use std::time::Duration; use std::time::{Duration, Instant};
use tokio::task::JoinSet; use tokio::task::JoinSet;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
@@ -61,8 +61,10 @@ pub async fn run_round_robin_router(
config: &ServiceConfig, config: &ServiceConfig,
transport: NatsTransport, transport: NatsTransport,
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
let request_subject = format!("svc.{SERVICE_NAME}"); anyhow::ensure!(
let queue_group = format!("{SERVICE_NAME}-router"); config.max_concurrent_requests > 0,
"snowflake router request concurrency must be positive"
);
let picker = Arc::new(SnowflakeShardPicker::new(config.shard_count)); let picker = Arc::new(SnowflakeShardPicker::new(config.shard_count));
let mut tasks = JoinSet::new(); let mut tasks = JoinSet::new();
let health_addr = config.listen_addr; let health_addr = config.listen_addr;
@@ -82,84 +84,12 @@ pub async fn run_round_robin_router(
.await .await
} }
}); });
let req_transport = transport.clone(); tasks.spawn(serve_requests(
tasks.spawn(async move { transport,
loop { picker,
let mut sub = req_transport http_metrics,
.subscribe_queue(&request_subject, &queue_group) config.max_concurrent_requests,
.await?; ));
info!(subject = request_subject, "snowflake router listening for requests");
loop {
let msg = tokio::select! {
msg_opt = sub.next() => {
let Some(msg) = msg_opt else {
warn!("snowflake router request subscription stream ended, will re-subscribe");
break;
};
msg
}
_ = req_transport.wait_for_reconnect() => {
info!("NATS reconnected, re-subscribing snowflake router request listener");
break;
}
};
let transport = req_transport.clone();
let picker = picker.clone();
if msg.payload().len() > MAX_ROUTER_REQUEST_BYTES {
warn!(
payload_bytes = msg.payload().len(),
max_payload_bytes = MAX_ROUTER_REQUEST_BYTES,
"rejecting oversized snowflake request"
);
reply_json_error(&msg, &transport, "request_too_large").await;
continue;
}
let request: SnowflakeRequest = match serde_json::from_slice(msg.payload()) {
Ok(request) => request,
Err(error) => {
warn!(error = %error, "failed to decode snowflake request");
reply_json_error(&msg, &transport, "decode_error").await;
continue;
}
};
tokio::spawn(async move {
let shard_id = picker.pick_shard(&request);
let shard_subject = format!("svc.{SERVICE_NAME}.shard.{shard_id}");
let payload = match rmp_serde::to_vec(&request) {
Ok(payload) => payload,
Err(error) => {
warn!(error = %error, "failed to encode snowflake shard request");
reply_json_error(&msg, &transport, "encode_error").await;
return;
}
};
let response_bytes = match transport
.request(&shard_subject, &payload, SHARD_REQUEST_TIMEOUT)
.await
{
Ok(response_bytes) => response_bytes,
Err(error) => {
debug!(error = %error, shard_id, "snowflake shard request failed");
reply_json_error(&msg, &transport, "shard_unavailable").await;
return;
}
};
let response: SnowflakeResponse = match rmp_serde::from_slice(&response_bytes) {
Ok(response) => response,
Err(error) => {
debug!(error = %error, shard_id, "failed to decode snowflake shard response");
reply_json_error(&msg, &transport, "shard_decode_error").await;
return;
}
};
if msg.has_reply() {
let response_json = serde_json::to_vec(&response).unwrap_or_default();
let _ = msg.reply(&transport, &response_json).await;
}
});
}
}
});
tokio::select! { tokio::select! {
result = tasks.join_next() => { result = tasks.join_next() => {
match result { match result {
@@ -177,15 +107,136 @@ pub async fn run_round_robin_router(
} }
} }
async fn reply_json_error( async fn serve_requests(
msg: &fluxer_svc::transport::NatsMessage, transport: NatsTransport,
transport: &NatsTransport, picker: Arc<SnowflakeShardPicker>,
error: &str, metrics: Arc<ServiceMetrics>,
max_concurrent_requests: usize,
) -> anyhow::Result<()> {
let request_subject = format!("svc.{SERVICE_NAME}");
let queue_group = format!("{SERVICE_NAME}-router");
let mut requests = JoinSet::new();
loop {
let mut sub = transport
.subscribe_queue(&request_subject, &queue_group)
.await?;
info!(
subject = request_subject,
max_concurrent_requests, "snowflake router listening for requests"
);
loop {
let msg = tokio::select! {
result = requests.join_next(), if !requests.is_empty() => {
if let Err(error) = result.expect("nonempty snowflake request set") {
warn!(error = %error, "snowflake request task failed");
}
continue;
}
msg = sub.next() => {
let Some(msg) = msg else {
warn!("snowflake router request subscription stream ended, will re-subscribe");
break;
};
msg
}
};
while let Some(result) = requests.try_join_next() {
if let Err(error) = result {
warn!(error = %error, "snowflake request task failed");
}
}
metrics.record_request();
if requests.len() >= max_concurrent_requests {
debug!(
max_concurrent_requests,
"shedding overloaded snowflake router request"
);
metrics.record_request_error();
reply_json_error(&msg, &transport, "overloaded").await;
continue;
}
requests.spawn(handle_request(
msg,
transport.clone(),
picker.clone(),
metrics.clone(),
));
}
}
}
async fn handle_request(
msg: NatsMessage,
transport: NatsTransport,
picker: Arc<SnowflakeShardPicker>,
metrics: Arc<ServiceMetrics>,
) { ) {
let started = Instant::now();
match forward_request(msg.payload(), &transport, &picker, &metrics).await {
Ok(response) => {
if msg.has_reply() {
let payload = serde_json::to_vec(&response)
.expect("snowflake responses contain only JSON-serializable strings");
reply(&msg, &transport, &payload).await;
}
}
Err(error) => {
metrics.record_request_error();
reply_json_error(&msg, &transport, error).await;
}
}
metrics.record_request_duration(started.elapsed().as_millis() as u64);
}
async fn forward_request(
payload: &[u8],
transport: &NatsTransport,
picker: &SnowflakeShardPicker,
metrics: &ServiceMetrics,
) -> Result<SnowflakeResponse, &'static str> {
if payload.len() > MAX_ROUTER_REQUEST_BYTES {
warn!(
payload_bytes = payload.len(),
max_payload_bytes = MAX_ROUTER_REQUEST_BYTES,
"rejecting oversized snowflake request"
);
return Err("request_too_large");
}
let request: SnowflakeRequest = serde_json::from_slice(payload).map_err(|error| {
warn!(error = %error, "failed to decode snowflake request");
"decode_error"
})?;
let shard_id = picker.pick_shard(&request);
let shard_subject = format!("svc.{SERVICE_NAME}.shard.{shard_id}");
let payload = rmp_serde::to_vec(&request).map_err(|error| {
warn!(error = %error, "failed to encode snowflake shard request");
"encode_error"
})?;
metrics.record_shard_forward();
let response_bytes = transport
.request(&shard_subject, &payload, SHARD_REQUEST_TIMEOUT)
.await
.map_err(|error| {
debug!(error = %error, shard_id, "snowflake shard request failed");
"shard_unavailable"
})?;
rmp_serde::from_slice(&response_bytes).map_err(|error| {
debug!(error = %error, shard_id, "failed to decode snowflake shard response");
"shard_decode_error"
})
}
async fn reply_json_error(msg: &NatsMessage, transport: &NatsTransport, error: &str) {
if msg.has_reply() { if msg.has_reply() {
let error_response = let payload = serde_json::to_vec(&serde_json::json!({ "error": error }))
serde_json::to_vec(&serde_json::json!({ "error": error })).unwrap_or_default(); .expect("snowflake error responses contain only a JSON-serializable string");
let _ = msg.reply(transport, &error_response).await; reply(msg, transport, &payload).await;
}
}
async fn reply(msg: &NatsMessage, transport: &NatsTransport, payload: &[u8]) {
if let Err(error) = msg.reply(transport, payload).await {
warn!(error = %error, "failed to reply to snowflake request");
} }
} }
+2 -1
View File
@@ -75,8 +75,9 @@ impl ServiceConfig {
}; };
let mode = match optional_from(&get, "FLUXER_SVC_MODE").as_deref() { let mode = match optional_from(&get, "FLUXER_SVC_MODE").as_deref() {
None | Some("router") => Mode::Router,
Some("shard") => Mode::Shard, Some("shard") => Mode::Shard,
_ => Mode::Router, Some(other) => anyhow::bail!("unsupported FLUXER_SVC_MODE: {other}"),
}; };
let shard_count = optional_from(&get, "FLUXER_SVC_SHARD_COUNT") let shard_count = optional_from(&get, "FLUXER_SVC_SHARD_COUNT")
+4
View File
@@ -9,6 +9,10 @@ impl HashRing {
Self { shard_count } Self { shard_count }
} }
pub fn shard_count(&self) -> u32 {
self.shard_count
}
pub fn owner(&self, route_key: &str) -> u32 { pub fn owner(&self, route_key: &str) -> u32 {
let mut best_shard = 0u32; let mut best_shard = 0u32;
let mut best_score = 0u64; let mut best_score = 0u64;
+154 -83
View File
@@ -2,11 +2,15 @@
use crate::config::ServiceConfig; use crate::config::ServiceConfig;
use crate::hash_ring::HashRing; use crate::hash_ring::HashRing;
use crate::metrics::{ServiceMetrics, now_ms}; use crate::metrics::ServiceMetrics;
use crate::transport::{Transport, TransportMessage, TransportSubscriber, reply_message}; use crate::transport::{
Transport, TransportMessage, TransportSubscriber, reply_bytes, reply_json_error,
};
use anyhow::Context;
use futures::stream::{FuturesUnordered, StreamExt};
use moka::future::Cache; use moka::future::Cache;
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::{Duration, Instant};
use tokio::sync::{Semaphore, TryAcquireError}; use tokio::sync::{Semaphore, TryAcquireError};
use tokio::task::JoinSet; use tokio::task::JoinSet;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
@@ -14,6 +18,7 @@ use tracing::{debug, info, warn};
pub(crate) const SHARD_REQUEST_TIMEOUT: Duration = Duration::from_secs(5); pub(crate) const SHARD_REQUEST_TIMEOUT: Duration = Duration::from_secs(5);
const INFLIGHT_TTL: Duration = Duration::from_millis(200); const INFLIGHT_TTL: Duration = Duration::from_millis(200);
const INFLIGHT_MAX_ENTRIES: u64 = 10_000; const INFLIGHT_MAX_ENTRIES: u64 = 10_000;
const MAX_BROADCAST_CONCURRENCY: usize = 32;
const MAX_ROUTER_REQUEST_BYTES: usize = 2 * 1024 * 1024; const MAX_ROUTER_REQUEST_BYTES: usize = 2 * 1024 * 1024;
const LEGACY_SHARD_DECODE_ERROR: &[u8] = br#"{"error":"shard_request_decode_error"}"#; const LEGACY_SHARD_DECODE_ERROR: &[u8] = br#"{"error":"shard_request_decode_error"}"#;
type InflightKey = (String, String); type InflightKey = (String, String);
@@ -30,6 +35,13 @@ pub trait RouterService: Send + Sync + 'static {
None None
} }
fn is_broadcast_request(_request: &Self::Request) -> bool {
false
}
fn is_broadcast_acknowledgement(_response: &Self::Response) -> bool {
false
}
fn l1_lookup(&self, _req: &Self::Request) -> Option<Self::Response> { fn l1_lookup(&self, _req: &Self::Request) -> Option<Self::Response> {
None None
} }
@@ -111,6 +123,88 @@ async fn dispatch_to_shard<S: RouterService>(
} }
} }
async fn forward_to_all_shards<S: RouterService>(
transport: &impl Transport,
service: &S,
ring: &HashRing,
request: &S::Request,
) -> anyhow::Result<Vec<u8>> {
let payload = rmp_serde::to_vec_named(request)
.context("failed to encode broadcast request as msgpack")?;
let payload = payload.as_slice();
let service_name = service.service_name();
let deadline = tokio::time::Instant::now() + SHARD_REQUEST_TIMEOUT;
let request_shard = |shard_id| async move {
let timeout = deadline.saturating_duration_since(tokio::time::Instant::now());
if timeout.is_zero() {
anyhow::bail!("broadcast deadline expired before shard {shard_id}");
}
let subject = format!("svc.{service_name}.shard.{shard_id}");
let response_bytes = transport
.request(&subject, payload, timeout)
.await
.with_context(|| format!("broadcast request to shard {shard_id} failed"))?;
let response = rmp_serde::from_slice::<S::Response>(&response_bytes)
.with_context(|| format!("invalid broadcast response from shard {shard_id}"))?;
if !S::is_broadcast_acknowledgement(&response) {
anyhow::bail!("unexpected broadcast acknowledgement from shard {shard_id}");
}
anyhow::Ok(response)
};
let mut shard_ids = 0..ring.shard_count();
let mut pending = FuturesUnordered::new();
for shard_id in shard_ids.by_ref().take(MAX_BROADCAST_CONCURRENCY) {
pending.push(request_shard(shard_id));
}
let mut acknowledgement = None;
let mut first_error = None;
while let Some(result) = pending.next().await {
match result {
Ok(response) => acknowledgement = Some(response),
Err(error) => {
first_error.get_or_insert(error);
}
}
if first_error.is_none()
&& let Some(shard_id) = shard_ids.next()
{
pending.push(request_shard(shard_id));
}
}
if let Some(error) = first_error {
return Err(error);
}
let acknowledgement = acknowledgement.context("broadcast request has no configured shards")?;
if S::CACHES_RESPONSES {
rmp_serde::to_vec_named(&acknowledgement)
.context("failed to encode broadcast acknowledgement as msgpack")
} else {
serde_json::to_vec(&acknowledgement)
.context("failed to encode broadcast acknowledgement as json")
}
}
async fn reply_json_response(
message: &impl TransportMessage,
transport: &impl Transport,
response: &impl serde::Serialize,
metrics: &ServiceMetrics,
) {
if !message.has_reply() {
return;
}
match serde_json::to_vec(response) {
Ok(payload) => reply_bytes(message, transport, &payload).await,
Err(error) => {
warn!(error = %error, subject = message.subject(), "failed to encode router response");
metrics.record_request_error();
reply_json_error(message, transport, "encode_error").await;
}
}
}
async fn handle_router_request<S, T>( async fn handle_router_request<S, T>(
msg: T::Message, msg: T::Message,
transport: T, transport: T,
@@ -122,7 +216,7 @@ async fn handle_router_request<S, T>(
S: RouterService, S: RouterService,
T: Transport, T: Transport,
{ {
let request_start = now_ms(); let request_start = Instant::now();
metrics.record_request(); metrics.record_request();
if msg.payload().len() > MAX_ROUTER_REQUEST_BYTES { if msg.payload().len() > MAX_ROUTER_REQUEST_BYTES {
warn!( warn!(
@@ -131,12 +225,7 @@ async fn handle_router_request<S, T>(
"rejecting oversized router request" "rejecting oversized router request"
); );
metrics.record_request_error(); metrics.record_request_error();
if msg.has_reply() { reply_json_error(&msg, &transport, "request_too_large").await;
let error_response =
serde_json::to_vec(&serde_json::json!({"error": "request_too_large"}))
.unwrap_or_default();
let _ = reply_message(&msg, &transport, &error_response).await;
}
return; return;
} }
let request: S::Request = match serde_json::from_slice(msg.payload()) { let request: S::Request = match serde_json::from_slice(msg.payload()) {
@@ -144,25 +233,20 @@ async fn handle_router_request<S, T>(
Err(err) => { Err(err) => {
warn!(error = %err, "failed to decode incoming request"); warn!(error = %err, "failed to decode incoming request");
metrics.record_request_error(); metrics.record_request_error();
if msg.has_reply() { reply_json_error(&msg, &transport, "decode_error").await;
let error_response =
serde_json::to_vec(&serde_json::json!({"error": "decode_error"}))
.unwrap_or_default();
let _ = reply_message(&msg, &transport, &error_response).await;
}
return; return;
} }
}; };
let request = Arc::new(request); let request = Arc::new(request);
let broadcast = S::is_broadcast_request(&request);
if let Some(cached) = service.l1_lookup(&request) { if S::CACHES_RESPONSES
&& !broadcast
&& let Some(cached) = service.l1_lookup(&request)
{
metrics.record_cache_hit(); metrics.record_cache_hit();
let elapsed = (now_ms() - request_start).max(0) as u64; metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
metrics.record_request_duration(elapsed); reply_json_response(&msg, &transport, &cached, &metrics).await;
if msg.has_reply() {
let response_bytes = serde_json::to_vec(&cached).unwrap_or_default();
let _ = reply_message(&msg, &transport, &response_bytes).await;
}
return; return;
} }
@@ -170,7 +254,9 @@ async fn handle_router_request<S, T>(
metrics.record_cache_miss(); metrics.record_cache_miss();
metrics.record_shard_forward(); metrics.record_shard_forward();
let coalesce_result = if let Some(coalesce_key) = S::coalesce_key(&request) { let coalesce_result = if broadcast {
forward_to_all_shards::<S>(&transport, service.as_ref(), ring.as_ref(), &request).await
} else if let Some(coalesce_key) = S::coalesce_key(&request) {
let forward_transport = transport.clone(); let forward_transport = transport.clone();
let forward_service = service.clone(); let forward_service = service.clone();
let forward_ring = ring.clone(); let forward_ring = ring.clone();
@@ -208,36 +294,31 @@ async fn handle_router_request<S, T>(
.await .await
}; };
let elapsed = (now_ms() - request_start).max(0) as u64; metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
metrics.record_request_duration(elapsed);
match coalesce_result { match coalesce_result {
Ok(response_bytes) => { Ok(response_bytes) => {
if !S::CACHES_RESPONSES { if !S::CACHES_RESPONSES {
if msg.has_reply() { if msg.has_reply() {
let _ = reply_message(&msg, &transport, &response_bytes).await; reply_bytes(&msg, &transport, &response_bytes).await;
} }
return; return;
} }
match rmp_serde::from_slice::<S::Response>(&response_bytes) { match rmp_serde::from_slice::<S::Response>(&response_bytes) {
Ok(response) => { Ok(response) => {
service.l1_insert(&request, &response); if !broadcast {
if msg.has_reply() { service.l1_insert(&request, &response);
let json = serde_json::to_vec(&response).unwrap_or_default();
let _ = reply_message(&msg, &transport, &json).await;
} }
reply_json_response(&msg, &transport, &response, &metrics).await;
} }
Err(err) => { Err(err) => {
debug!(error = %err, "failed to decode shard response"); debug!(error = %err, "failed to decode shard response");
if msg.has_reply() { if msg.has_reply() {
if serde_json::from_slice::<serde_json::Value>(&response_bytes).is_ok() { if serde_json::from_slice::<serde_json::Value>(&response_bytes).is_ok() {
let _ = reply_message(&msg, &transport, &response_bytes).await; reply_bytes(&msg, &transport, &response_bytes).await;
return; return;
} }
let error_response = reply_json_error(&msg, &transport, "shard_decode_error").await;
serde_json::to_vec(&serde_json::json!({"error": "shard_decode_error"}))
.unwrap_or_default();
let _ = reply_message(&msg, &transport, &error_response).await;
} }
} }
} }
@@ -245,12 +326,7 @@ async fn handle_router_request<S, T>(
Err(err) => { Err(err) => {
debug!(error = %err, "shard request failed (coalesced)"); debug!(error = %err, "shard request failed (coalesced)");
metrics.record_request_error(); metrics.record_request_error();
if msg.has_reply() { reply_json_error(&msg, &transport, "shard_unavailable").await;
let error_response =
serde_json::to_vec(&serde_json::json!({"error": "shard_unavailable"}))
.unwrap_or_default();
let _ = reply_message(&msg, &transport, &error_response).await;
}
} }
} }
} }
@@ -267,7 +343,6 @@ where
let ring = Arc::new(HashRing::new(config.shard_count)); let ring = Arc::new(HashRing::new(config.shard_count));
let name = service.service_name().to_owned(); let name = service.service_name().to_owned();
let request_subject = format!("svc.{name}"); let request_subject = format!("svc.{name}");
let invalidate_subject = format!("svc.{name}.invalidate.>");
let queue_group = format!("{name}-router"); let queue_group = format!("{name}-router");
let metrics = Arc::new(ServiceMetrics::default()); let metrics = Arc::new(ServiceMetrics::default());
@@ -295,6 +370,7 @@ where
let req_metrics = metrics.clone(); let req_metrics = metrics.clone();
let req_permits = Arc::new(Semaphore::new(config.max_concurrent_requests)); let req_permits = Arc::new(Semaphore::new(config.max_concurrent_requests));
tasks.spawn(async move { tasks.spawn(async move {
let mut requests = JoinSet::new();
loop { loop {
let mut sub = req_transport let mut sub = req_transport
.subscribe_queue(&request_subject, &req_queue) .subscribe_queue(&request_subject, &req_queue)
@@ -307,6 +383,12 @@ where
loop { loop {
let msg = tokio::select! { let msg = tokio::select! {
result = requests.join_next(), if !requests.is_empty() => {
if let Err(err) = result.expect("nonempty router request set") {
warn!(error = %err, "router request task failed");
}
continue;
}
msg_opt = sub.next() => { msg_opt = sub.next() => {
let Some(msg) = msg_opt else { let Some(msg) = msg_opt else {
warn!("router request subscription stream ended, will re-subscribe"); warn!("router request subscription stream ended, will re-subscribe");
@@ -314,24 +396,20 @@ where
}; };
msg msg
} }
_ = req_transport.wait_for_reconnect() => {
info!("NATS reconnected, re-subscribing router request listener");
break;
}
}; };
while let Some(result) = requests.try_join_next() {
if let Err(err) = result {
warn!(error = %err, "router request task failed");
}
}
let permit = match req_permits.clone().try_acquire_owned() { let permit = match req_permits.clone().try_acquire_owned() {
Ok(permit) => permit, Ok(permit) => permit,
Err(TryAcquireError::NoPermits) => { Err(TryAcquireError::NoPermits) => {
debug!("shedding router request, no permits available"); debug!("shedding router request, no permits available");
req_metrics.record_request(); req_metrics.record_request();
req_metrics.record_request_error(); req_metrics.record_request_error();
if msg.has_reply() { reply_json_error(&msg, &req_transport, "overloaded").await;
let error_response =
serde_json::to_vec(&serde_json::json!({"error": "overloaded"}))
.unwrap_or_default();
let _ = reply_message(&msg, &req_transport, &error_response).await;
}
continue; continue;
} }
Err(TryAcquireError::Closed) => return anyhow::Ok(()), Err(TryAcquireError::Closed) => return anyhow::Ok(()),
@@ -341,7 +419,7 @@ where
let ring = ring.clone(); let ring = ring.clone();
let inflight = inflight.clone(); let inflight = inflight.clone();
let metrics = req_metrics.clone(); let metrics = req_metrics.clone();
tokio::spawn(async move { requests.spawn(async move {
let _permit = permit; let _permit = permit;
handle_router_request::<S, _>(msg, transport, service, ring, inflight, metrics) handle_router_request::<S, _>(msg, transport, service, ring, inflight, metrics)
.await; .await;
@@ -350,39 +428,32 @@ where
} }
}); });
let inv_transport = transport.clone(); if S::CACHES_RESPONSES {
let inv_service = service.clone(); let inv_transport = transport.clone();
tasks.spawn(async move { let inv_service = service.clone();
loop { let invalidate_prefix = format!("svc.{name}.invalidate.");
let mut sub = inv_transport.subscribe(&invalidate_subject).await?; let invalidate_subject = format!("{invalidate_prefix}>");
info!( tasks.spawn(async move {
subject = invalidate_subject,
"router listening for cache invalidations"
);
loop { loop {
tokio::select! { let mut sub = inv_transport.subscribe(&invalidate_subject).await?;
msg_opt = sub.next() => { info!(
let Some(msg) = msg_opt else { subject = invalidate_subject,
warn!("router invalidation subscription stream ended, will re-subscribe"); "router listening for cache invalidations"
break; );
};
let subject = msg.subject().to_owned(); while let Some(msg) = sub.next().await {
let key = subject if let Some(key) = msg
.strip_prefix(&format!("svc.{name}.invalidate.")) .subject()
.unwrap_or(""); .strip_prefix(&invalidate_prefix)
if !key.is_empty() { .filter(|key| !key.is_empty())
inv_service.l1_invalidate(key); {
} inv_service.l1_invalidate(key);
}
_ = inv_transport.wait_for_reconnect() => {
info!("NATS reconnected, re-subscribing router invalidation listener");
break;
} }
} }
warn!("router invalidation subscription stream ended, will re-subscribe");
} }
} });
}); }
tokio::select! { tokio::select! {
result = tasks.join_next() => { result = tasks.join_next() => {
+70 -54
View File
@@ -1,11 +1,15 @@
// SPDX-License-Identifier: AGPL-3.0-or-later // SPDX-License-Identifier: AGPL-3.0-or-later
use crate::config::ServiceConfig; use crate::config::ServiceConfig;
use crate::metrics::{ServiceMetrics, now_ms}; use crate::metrics::ServiceMetrics;
use crate::transport::{Transport, TransportMessage, TransportSubscriber, reply_message}; use crate::transport::{
Transport, TransportMessage, TransportSubscriber, reply_bytes, reply_json_error,
};
use anyhow::Context;
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::atomic::{AtomicBool, Ordering};
use tokio::sync::{Semaphore, TryAcquireError}; use std::time::Instant;
use tokio::sync::{Semaphore, TryAcquireError, oneshot};
use tokio::task::JoinSet; use tokio::task::JoinSet;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
@@ -93,9 +97,14 @@ where
let shard_is_serving = is_serving.clone(); let shard_is_serving = is_serving.clone();
let shard_permits = request_permits.clone(); let shard_permits = request_permits.clone();
let shard_metrics = metrics.clone(); let shard_metrics = metrics.clone();
tasks.spawn(async move { let (stop_requests, mut stop_requests_rx) = oneshot::channel();
loop { let request_task = tasks.spawn(async move {
let mut sub = shard_transport.subscribe(&shard_subject).await?; let mut requests = JoinSet::new();
'listening: loop {
let mut sub = tokio::select! {
_ = &mut stop_requests_rx => break 'listening,
result = shard_transport.subscribe(&shard_subject) => result?,
};
info!( info!(
subject = shard_subject, subject = shard_subject,
shard_id, shard_id,
@@ -105,12 +114,23 @@ where
loop { loop {
tokio::select! { tokio::select! {
_ = &mut stop_requests_rx => break 'listening,
result = requests.join_next(), if !requests.is_empty() => {
if let Err(err) = result.expect("nonempty shard request set") {
warn!(error = %err, "shard request task failed");
}
}
msg_opt = sub.next() => { msg_opt = sub.next() => {
let Some(msg) = msg_opt else { let Some(msg) = msg_opt else {
warn!("shard subscription stream ended, will re-subscribe"); warn!("shard subscription stream ended, will re-subscribe");
break; break;
}; };
while let Some(result) = requests.try_join_next() {
if let Err(err) = result {
warn!(error = %err, "shard request task failed");
}
}
if !shard_is_serving.load(Ordering::SeqCst) { if !shard_is_serving.load(Ordering::SeqCst) {
continue; continue;
} }
@@ -133,27 +153,25 @@ where
debug!("shedding shard request, no permits available"); debug!("shedding shard request, no permits available");
shard_metrics.record_request(); shard_metrics.record_request();
shard_metrics.record_request_error(); shard_metrics.record_request_error();
reply_shard_error(&msg, &transport, "overloaded").await; reply_json_error(&msg, &transport, "overloaded").await;
continue; continue;
} }
Err(TryAcquireError::Closed) => return anyhow::Ok(()), Err(TryAcquireError::Closed) => return anyhow::Ok(()),
}; };
let raw_payload = msg.payload().to_vec(); requests.spawn(async move {
tokio::spawn(async move {
let _permit = permit; let _permit = permit;
if !is_serving.load(Ordering::SeqCst) { if !is_serving.load(Ordering::SeqCst) {
return; return;
} }
let request_start = now_ms(); let request_start = Instant::now();
metrics.record_request(); metrics.record_request();
let encoding = WireEncoding::detect(&raw_payload); let encoding = WireEncoding::detect(msg.payload());
let request: S::Request = match encoding.decode(&raw_payload) { let request: S::Request = match encoding.decode(msg.payload()) {
Ok(r) => r, Ok(r) => r,
Err(err) => { Err(err) => {
warn!(error = %err, ?encoding, "failed to decode shard request"); warn!(error = %err, ?encoding, "failed to decode shard request");
metrics.record_request_error(); metrics.record_request_error();
reply_shard_error(&msg, &transport, "shard_request_decode_error") reply_json_error(&msg, &transport, "shard_request_decode_error")
.await; .await;
return; return;
} }
@@ -161,25 +179,20 @@ where
match service.handle(request).await { match service.handle(request).await {
Ok(response) => { Ok(response) => {
let elapsed = (now_ms() - request_start).max(0) as u64; metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
metrics.record_request_duration(elapsed);
if msg.has_reply() { if msg.has_reply() {
match encoding.encode(&response) { match encoding.encode(&response) {
Ok(response_bytes) => { Ok(response_bytes) => {
if let Err(err) = reply_bytes(&msg, &transport, &response_bytes).await;
reply_message(&msg, &transport, &response_bytes).await
{
debug!(
error = %err,
"failed to send shard reply"
);
}
} }
Err(err) => { Err(err) => {
warn!( warn!(
error = %err, error = %err,
?encoding,
"failed to encode shard response" "failed to encode shard response"
); );
metrics.record_request_error();
reply_json_error(&msg, &transport, "encode_error").await;
} }
} }
} }
@@ -187,20 +200,19 @@ where
Err(err) => { Err(err) => {
warn!(error = %err, "shard handler returned error"); warn!(error = %err, "shard handler returned error");
metrics.record_request_error(); metrics.record_request_error();
let elapsed = (now_ms() - request_start).max(0) as u64; metrics.record_request_duration(request_start.elapsed().as_millis() as u64);
metrics.record_request_duration(elapsed); reply_json_error(&msg, &transport, "shard_handler_error").await;
reply_shard_error(&msg, &transport, "shard_handler_error").await;
} }
} }
}); });
} }
_ = shard_transport.wait_for_reconnect() => {
info!("NATS reconnected, re-subscribing shard listener");
break;
}
} }
} }
} }
while let Some(result) = requests.join_next().await {
result.context("shard request task failed while draining")?;
}
anyhow::Ok(())
}); });
tokio::select! { tokio::select! {
@@ -217,20 +229,34 @@ where
is_serving.store(false, Ordering::SeqCst); is_serving.store(false, Ordering::SeqCst);
let max_permits = config.max_concurrent_requests; let _ = stop_requests.send(());
let drain_permits = request_permits.clone(); let drain = async {
crate::shutdown::drain_with_timeout( loop {
async move { let (task_id, result) = tasks.join_next_with_id().await
if let Ok(_permit) = drain_permits.acquire_many(max_permits as u32).await { .expect("running shard listener while draining")
info!( .context("shard service task failed while draining")?;
max_concurrent_requests = max_permits, result?;
"all in-flight requests drained" if task_id == request_task.id() {
); return anyhow::Ok(());
} }
}, }
crate::shutdown::DEFAULT_DRAIN_TIMEOUT, };
) match tokio::time::timeout(crate::shutdown::DEFAULT_DRAIN_TIMEOUT, drain).await {
.await; Ok(result) => {
result?;
info!(
max_concurrent_requests = config.max_concurrent_requests,
"all in-flight requests drained"
);
info!("graceful drain completed");
}
Err(_) => {
warn!(
timeout_secs = crate::shutdown::DEFAULT_DRAIN_TIMEOUT.as_secs(),
"drain timeout exceeded, proceeding with shutdown"
);
}
}
info!("shard shutdown complete"); info!("shard shutdown complete");
Ok(()) Ok(())
@@ -238,16 +264,6 @@ where
} }
} }
async fn reply_shard_error(msg: &impl TransportMessage, transport: &impl Transport, code: &str) {
if !msg.has_reply() {
return;
}
let response = serde_json::to_vec(&serde_json::json!({ "error": code })).unwrap_or_default();
if let Err(err) = reply_message(msg, transport, &response).await {
debug!(error = %err, "failed to send shard error reply");
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
+24 -1
View File
@@ -6,7 +6,7 @@ use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration; use std::time::Duration;
use tokio::sync::Notify; use tokio::sync::Notify;
use tracing::{info, warn}; use tracing::{debug, info, warn};
const NATS_SUBSCRIPTION_CAPACITY: usize = 8_192; const NATS_SUBSCRIPTION_CAPACITY: usize = 8_192;
const SLOW_CONSUMER_LOG_INTERVAL_MS: u64 = 1_000; const SLOW_CONSUMER_LOG_INTERVAL_MS: u64 = 1_000;
@@ -62,6 +62,29 @@ where
Ok(()) Ok(())
} }
pub(crate) async fn reply_bytes(
message: &impl TransportMessage,
transport: &impl Transport,
payload: &[u8],
) {
if let Err(error) = reply_message(message, transport, payload).await {
debug!(error = %error, subject = message.subject(), "failed to send service reply");
}
}
pub(crate) async fn reply_json_error(
message: &impl TransportMessage,
transport: &impl Transport,
code: &str,
) {
if !message.has_reply() {
return;
}
let payload = serde_json::to_vec(&serde_json::json!({ "error": code }))
.expect("service error responses contain only a JSON-serializable string");
reply_bytes(message, transport, &payload).await;
}
#[derive(Clone)] #[derive(Clone)]
pub struct NatsTransport { pub struct NatsTransport {
client: async_nats::Client, client: async_nats::Client,
+1 -1
View File
@@ -11,7 +11,7 @@ chrono = { version = "0.4", default-features = false }
fluxer_common = { path = "../fluxer_common" } fluxer_common = { path = "../fluxer_common" }
fluxer-svc = { path = "../fluxer_svc" } fluxer-svc = { path = "../fluxer_svc" }
futures = "0.3.32" futures = "0.3.32"
moka = { version = "0.12.15", features = ["future", "sync"] } moka = { version = "0.12.15", features = ["future"] }
rmp-serde = "1.3" rmp-serde = "1.3"
scylla = { version = "1.6.0", features = ["chrono-04"], optional = true } scylla = { version = "1.6.0", features = ["chrono-04"], optional = true }
serde = { version = "1.0.228", features = ["derive"] } serde = { version = "1.0.228", features = ["derive"] }
+4 -17
View File
@@ -26,10 +26,7 @@ async fn main() -> anyhow::Result<()> {
); );
match config.mode { match config.mode {
Mode::Router => { Mode::Router => fluxer_svc::router::run_router(&config, UsersRouter, transport).await,
let router = UsersRouter::new(config.cache_max_entries, config.cache_ttl);
fluxer_svc::router::run_router(&config, router, transport).await
}
Mode::Shard => { Mode::Shard => {
let shard = match config.database_backend { let shard = match config.database_backend {
DatabaseBackend::Postgres => { DatabaseBackend::Postgres => {
@@ -37,12 +34,7 @@ async fn main() -> anyhow::Result<()> {
fluxer_svc::postgres::PostgresConfig::from_service_config(&config); fluxer_svc::postgres::PostgresConfig::from_service_config(&config);
let pool = fluxer_svc::postgres::connect(&postgres_config).await?; let pool = fluxer_svc::postgres::connect(&postgres_config).await?;
let kv = fluxer_svc::postgres::KvClient::new(pool, &postgres_config)?; let kv = fluxer_svc::postgres::KvClient::new(pool, &postgres_config)?;
UsersShard::new_postgres( UsersShard::new_postgres(kv, config.cache_max_entries, config.cache_ttl)
kv,
transport.clone(),
config.cache_max_entries,
config.cache_ttl,
)?
} }
DatabaseBackend::Cassandra => { DatabaseBackend::Cassandra => {
#[cfg(feature = "scylla")] #[cfg(feature = "scylla")]
@@ -50,13 +42,8 @@ async fn main() -> anyhow::Result<()> {
let scylla_config = let scylla_config =
fluxer_svc::scylla::ScyllaConfig::from_service_config(&config); fluxer_svc::scylla::ScyllaConfig::from_service_config(&config);
let db = fluxer_svc::scylla::connect(&scylla_config).await?; let db = fluxer_svc::scylla::connect(&scylla_config).await?;
UsersShard::new_scylla( UsersShard::new_scylla(db, config.cache_max_entries, config.cache_ttl)
db, .await?
transport.clone(),
config.cache_max_entries,
config.cache_ttl,
)
.await?
} }
#[cfg(not(feature = "scylla"))] #[cfg(not(feature = "scylla"))]
{ {
+36 -187
View File
@@ -2,36 +2,24 @@
use crate::types::{UserRequest, UserResponse}; use crate::types::{UserRequest, UserResponse};
use fluxer_svc::router::RouterService; use fluxer_svc::router::RouterService;
use moka::sync::Cache;
use std::time::Duration;
pub struct UsersRouter { pub struct UsersRouter;
l1: Cache<String, UserResponse>,
}
impl UsersRouter {
pub fn new(max_entries: u64, ttl: Duration) -> Self {
Self {
l1: Cache::builder()
.max_capacity(max_entries)
.time_to_live(ttl)
.build(),
}
}
}
impl RouterService for UsersRouter { impl RouterService for UsersRouter {
type Request = UserRequest; type Request = UserRequest;
type Response = UserResponse; type Response = UserResponse;
const CACHES_RESPONSES: bool = false;
fn service_name(&self) -> &str { fn service_name(&self) -> &str {
"users" "users"
} }
fn route_key(req: &UserRequest) -> String { fn route_key(req: &UserRequest) -> String {
match req { match req {
UserRequest::GetById { user_id } => user_id.to_string(), UserRequest::GetById { user_id }
UserRequest::GetPartialById { user_id } => user_id.to_string(), | UserRequest::GetPartialById { user_id }
| UserRequest::Invalidate { user_id } => user_id.to_string(),
UserRequest::GetPartialsByIds { user_ids } => user_ids UserRequest::GetPartialsByIds { user_ids } => user_ids
.iter() .iter()
.min() .min()
@@ -43,7 +31,6 @@ impl RouterService for UsersRouter {
.min() .min()
.cloned() .cloned()
.unwrap_or_else(|| "0".to_owned()), .unwrap_or_else(|| "0".to_owned()),
UserRequest::Invalidate { user_id } => user_id.to_string(),
} }
} }
@@ -74,117 +61,12 @@ impl RouterService for UsersRouter {
} }
} }
fn l1_lookup(&self, req: &UserRequest) -> Option<UserResponse> { fn is_broadcast_request(req: &UserRequest) -> bool {
match req { matches!(req, UserRequest::Invalidate { .. })
UserRequest::GetById { user_id } => self.l1.get(&user_id.to_string()),
UserRequest::GetPartialById { user_id } => {
let cached = self.l1.get(&user_id.to_string())?;
match cached {
UserResponse::Found(ref user) => {
Some(UserResponse::FoundPartial(user.to_partial()))
}
UserResponse::FoundPartial(_) => Some(cached),
UserResponse::NotFound => Some(UserResponse::NotFound),
_ => None,
}
}
UserRequest::GetPartialsByIds { user_ids } => {
let mut partials = Vec::with_capacity(user_ids.len());
for user_id in user_ids {
let cached = self.l1.get(&user_id.to_string())?;
match cached {
UserResponse::Found(ref user) => partials.push(user.to_partial()),
UserResponse::FoundPartial(partial) => partials.push(partial),
UserResponse::NotFound => {}
_ => return None,
}
}
Some(UserResponse::FoundPartials(partials))
}
UserRequest::GetApiPartialById { user_id } => {
let cached = self.l1.get(user_id)?;
match cached {
UserResponse::Found(ref user) => {
Some(UserResponse::FoundApiPartial(user.to_api_partial()))
}
UserResponse::FoundPartial(ref partial) => {
Some(UserResponse::FoundApiPartial(partial.to_api_partial()))
}
UserResponse::FoundApiPartial(partial) => {
Some(UserResponse::FoundApiPartial(partial))
}
UserResponse::NotFound => Some(UserResponse::NotFound),
_ => None,
}
}
UserRequest::GetApiPartialsByIds { user_ids } => {
let mut partials = Vec::with_capacity(user_ids.len());
for user_id in user_ids {
let cached = self.l1.get(user_id)?;
match cached {
UserResponse::Found(ref user) => partials.push(user.to_api_partial()),
UserResponse::FoundPartial(ref partial) => {
partials.push(partial.to_api_partial())
}
UserResponse::FoundApiPartial(partial) => partials.push(partial),
UserResponse::NotFound => {}
_ => return None,
}
}
Some(UserResponse::FoundApiPartials(partials))
}
UserRequest::Invalidate { .. } => None,
}
} }
fn l1_insert(&self, req: &UserRequest, resp: &UserResponse) { fn is_broadcast_acknowledgement(response: &UserResponse) -> bool {
match req { matches!(response, UserResponse::Invalidated)
UserRequest::GetById { user_id } => {
self.l1.insert(user_id.to_string(), resp.clone());
}
UserRequest::GetPartialById { user_id } => {
if !matches!(
self.l1.get(&user_id.to_string()),
Some(UserResponse::Found(_))
) {
self.l1.insert(user_id.to_string(), resp.clone());
}
}
UserRequest::GetPartialsByIds { .. } => {
if let UserResponse::FoundPartials(partials) = resp {
for partial in partials {
let key = partial.user_id.to_string();
if !matches!(self.l1.get(&key), Some(UserResponse::Found(_))) {
self.l1
.insert(key, UserResponse::FoundPartial(partial.clone()));
}
}
}
}
UserRequest::GetApiPartialById { .. } => {
if let UserResponse::FoundApiPartial(partial) = resp {
self.l1.insert(
partial.id.clone(),
UserResponse::FoundApiPartial(partial.clone()),
);
}
}
UserRequest::GetApiPartialsByIds { .. } => {
if let UserResponse::FoundApiPartials(partials) = resp {
for partial in partials {
self.l1.insert(
partial.id.clone(),
UserResponse::FoundApiPartial(partial.clone()),
);
}
}
}
UserRequest::Invalidate { .. } => {}
}
}
fn l1_invalidate(&self, key: &str) {
self.l1.invalidate(key);
} }
} }
@@ -192,72 +74,39 @@ impl RouterService for UsersRouter {
mod tests { mod tests {
use super::*; use super::*;
fn api_partial(id: &str) -> crate::types::ApiUserPartial { #[test]
crate::types::ApiUserPartial { fn coalesce_key_ignores_batch_id_order() {
id: id.to_owned(), let forward = UserRequest::GetPartialsByIds {
username: "Ada".to_owned(), user_ids: vec![7, 42, 7],
discriminator: "0007".to_owned(), };
global_name: Some("Ada Lovelace".to_owned()), let reversed = UserRequest::GetPartialsByIds {
avatar: Some("avatar_hash".to_owned()), user_ids: vec![42, 7],
avatar_color: Some(0x336699), };
bot: None, assert_eq!(
system: None, UsersRouter::coalesce_key(&forward),
flags: 1, UsersRouter::coalesce_key(&reversed)
mention_flags: None, );
}
} }
#[test] #[test]
fn l1_caches_api_partial_batch_responses() { fn coalesce_key_separates_batches_with_different_ids() {
let router = UsersRouter::new(100, Duration::from_secs(30)); let left = UserRequest::GetApiPartialsByIds {
let partial = api_partial("9223372036854775807"); user_ids: vec!["7".to_owned(), "42".to_owned()],
router.l1_insert( };
&UserRequest::GetApiPartialsByIds { let right = UserRequest::GetApiPartialsByIds {
user_ids: vec![partial.id.clone()], user_ids: vec!["7".to_owned()],
}, };
&UserResponse::FoundApiPartials(vec![partial.clone()]), assert_ne!(
UsersRouter::coalesce_key(&left),
UsersRouter::coalesce_key(&right)
); );
let cached = router.l1_lookup(&UserRequest::GetApiPartialsByIds {
user_ids: vec![partial.id.clone()],
});
match cached {
Some(UserResponse::FoundApiPartials(partials)) => {
assert_eq!(partials.len(), 1);
assert_eq!(partials[0].id, partial.id);
assert_eq!(partials[0].username, partial.username);
}
other => panic!("unexpected cached response: {other:?}"),
}
} }
#[test] #[test]
fn l1_invalidates_cached_api_partials_by_user_id() { fn invalidations_are_never_coalesced() {
let router = UsersRouter::new(100, Duration::from_secs(30)); assert_eq!(
let partial = api_partial("42"); UsersRouter::coalesce_key(&UserRequest::Invalidate { user_id: 42 }),
router.l1_insert( None
&UserRequest::GetApiPartialById {
user_id: partial.id.clone(),
},
&UserResponse::FoundApiPartial(partial.clone()),
);
assert!(
router
.l1_lookup(&UserRequest::GetApiPartialById {
user_id: partial.id.clone(),
})
.is_some()
);
router.l1_invalidate(&partial.id);
assert!(
router
.l1_lookup(&UserRequest::GetApiPartialById {
user_id: partial.id,
})
.is_none()
); );
} }
} }
+167 -131
View File
@@ -4,7 +4,6 @@ use crate::types::{ApiUserPartial, User, UserPartial, UserRequest, UserResponse}
#[cfg(feature = "scylla")] #[cfg(feature = "scylla")]
use chrono::{DateTime, NaiveDate, Utc}; use chrono::{DateTime, NaiveDate, Utc};
use fluxer_svc::shard::ShardService; use fluxer_svc::shard::ShardService;
use fluxer_svc::transport::NatsTransport;
use fluxer_svc::{postgres, postgres::KeyPart}; use fluxer_svc::{postgres, postgres::KeyPart};
use futures::stream::{self, StreamExt}; use futures::stream::{self, StreamExt};
use moka::future::Cache; use moka::future::Cache;
@@ -17,8 +16,12 @@ use scylla::statement::prepared::PreparedStatement;
#[cfg(feature = "scylla")] #[cfg(feature = "scylla")]
use scylla::value::MaybeEmpty; use scylla::value::MaybeEmpty;
use serde::Deserialize; use serde::Deserialize;
use std::collections::{HashMap, hash_map::DefaultHasher};
use std::fmt::Write;
use std::future::Future; use std::future::Future;
use std::hash::{Hash, Hasher};
use std::sync::Arc; use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration; use std::time::Duration;
#[cfg(feature = "scylla")] #[cfg(feature = "scylla")]
@@ -60,6 +63,8 @@ const PARTIAL_USER_COLUMNS: &str = "\
mention_flags"; mention_flags";
const USER_BATCH_SIZE: usize = 128; const USER_BATCH_SIZE: usize = 128;
const USER_BATCH_CONCURRENCY: usize = 8; const USER_BATCH_CONCURRENCY: usize = 8;
const USER_CACHE_MIN_GENERATION_STRIPES: usize = 4096;
const USER_CACHE_MAX_GENERATION_STRIPES: usize = 1 << 20;
const FLUXER_SYSTEM_USER_ID: i64 = 0; const FLUXER_SYSTEM_USER_ID: i64 = 0;
const FLUXER_SYSTEM_USERNAME: &str = "Fluxer"; const FLUXER_SYSTEM_USERNAME: &str = "Fluxer";
const FLUXER_SYSTEM_DISCRIMINATOR: i32 = 0; const FLUXER_SYSTEM_DISCRIMINATOR: i32 = 0;
@@ -68,12 +73,19 @@ const USER_FLAG_STAFF: i64 = 1;
pub struct UsersShard { pub struct UsersShard {
storage: UsersStorage, storage: UsersStorage,
caches: UserCaches, caches: UserCaches,
transport: NatsTransport,
} }
struct UserCaches { struct UserCaches {
full: Cache<i64, Option<User>>, full: Cache<UserCacheKey, Option<User>>,
partial: Cache<i64, Option<UserPartial>>, partial: Cache<UserCacheKey, Option<UserPartial>>,
generations: Box<[AtomicU64]>,
generation_bumps: AtomicU64,
}
#[derive(Clone, Copy, Eq, Hash, PartialEq)]
struct UserCacheKey {
user_id: i64,
generation: u64,
} }
#[derive(Clone)] #[derive(Clone)]
@@ -159,26 +171,9 @@ struct FullUserDbRow {
timezone_privacy_flags: Option<i32>, timezone_privacy_flags: Option<i32>,
} }
#[cfg(feature = "scylla")]
#[derive(Debug, DeserializeRow)]
struct PartialUserDbRow {
user_id: i64,
username: String,
discriminator: i32,
global_name: Option<String>,
avatar_hash: Option<String>,
bot: Option<bool>,
system: Option<bool>,
flags: Option<i64>,
banner_hash: Option<String>,
banner_color: Option<i32>,
accent_color: Option<i32>,
avatar_color: Option<i32>,
mention_flags: Option<i32>,
}
#[derive(Debug, Deserialize)] #[derive(Debug, Deserialize)]
struct PartialUserKvRow { #[cfg_attr(feature = "scylla", derive(DeserializeRow))]
struct PartialUserDbRow {
user_id: i64, user_id: i64,
username: String, username: String,
discriminator: i32, discriminator: i32,
@@ -256,6 +251,16 @@ struct FullUserKvRow {
timezone_privacy_flags: Option<i32>, timezone_privacy_flags: Option<i32>,
} }
fn generation_stripes(max_entries: u64) -> usize {
usize::try_from(max_entries)
.unwrap_or(USER_CACHE_MAX_GENERATION_STRIPES)
.clamp(
USER_CACHE_MIN_GENERATION_STRIPES,
USER_CACHE_MAX_GENERATION_STRIPES,
)
.next_power_of_two()
}
impl UserCaches { impl UserCaches {
fn new(max_entries: u64, ttl: Duration) -> Self { fn new(max_entries: u64, ttl: Duration) -> Self {
Self { Self {
@@ -267,6 +272,23 @@ impl UserCaches {
.max_capacity(max_entries) .max_capacity(max_entries)
.time_to_live(ttl) .time_to_live(ttl)
.build(), .build(),
generations: (0..generation_stripes(max_entries))
.map(|_| AtomicU64::new(0))
.collect(),
generation_bumps: AtomicU64::new(0),
}
}
fn generation(&self, user_id: i64) -> &AtomicU64 {
let mut hasher = DefaultHasher::new();
user_id.hash(&mut hasher);
&self.generations[hasher.finish() as usize % self.generations.len()]
}
fn key(&self, user_id: i64) -> UserCacheKey {
UserCacheKey {
user_id,
generation: self.generation(user_id).load(Ordering::SeqCst),
} }
} }
@@ -274,13 +296,14 @@ impl UserCaches {
where where
F: Future<Output = anyhow::Result<Option<User>>>, F: Future<Output = anyhow::Result<Option<User>>>,
{ {
let key = self.key(user_id);
let user = self let user = self
.full .full
.try_get_with(user_id, fetch) .try_get_with(key, fetch)
.await .await
.map_err(|e: Arc<anyhow::Error>| anyhow::anyhow!("{e}"))?; .map_err(|e: Arc<anyhow::Error>| anyhow::anyhow!("{e}"))?;
self.partial self.partial
.insert(user_id, user.as_ref().map(User::to_partial)) .insert(key, user.as_ref().map(User::to_partial))
.await; .await;
Ok(user) Ok(user)
} }
@@ -294,43 +317,47 @@ impl UserCaches {
F: Future<Output = anyhow::Result<Option<UserPartial>>>, F: Future<Output = anyhow::Result<Option<UserPartial>>>,
{ {
self.partial self.partial
.try_get_with(user_id, fetch) .try_get_with(self.key(user_id), fetch)
.await .await
.map_err(|e: Arc<anyhow::Error>| anyhow::anyhow!("{e}")) .map_err(|e: Arc<anyhow::Error>| anyhow::anyhow!("{e}"))
} }
async fn get_partial(&self, user_id: i64) -> Option<Option<UserPartial>> { async fn get_partial(&self, user_id: i64) -> Option<Option<UserPartial>> {
self.partial.get(&user_id).await self.partial.get(&self.key(user_id)).await
} }
async fn insert_partial(&self, user_id: i64, partial: Option<UserPartial>) { async fn insert_partial(&self, key: UserCacheKey, partial: Option<UserPartial>) {
self.partial.insert(user_id, partial).await; self.partial.insert(key, partial).await;
} }
async fn invalidate(&self, user_id: i64) { async fn invalidate(&self, user_id: i64) {
self.full.invalidate(&user_id).await; let generation = self
self.partial.invalidate(&user_id).await; .generation(user_id)
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| {
current.checked_add(1)
})
.expect("user cache generation exhausted");
self.generation_bumps.fetch_add(1, Ordering::Relaxed);
let key = UserCacheKey {
user_id,
generation,
};
self.full.invalidate(&key).await;
self.partial.invalidate(&key).await;
} }
} }
impl UsersShard { impl UsersShard {
pub fn new_postgres( pub fn new_postgres(kv: postgres::KvClient, max_entries: u64, ttl: Duration) -> Self {
kv: postgres::KvClient, Self {
transport: NatsTransport,
max_entries: u64,
ttl: Duration,
) -> anyhow::Result<Self> {
Ok(Self {
storage: UsersStorage::Postgres(PostgresUsersStorage { kv }), storage: UsersStorage::Postgres(PostgresUsersStorage { kv }),
caches: UserCaches::new(max_entries, ttl), caches: UserCaches::new(max_entries, ttl),
transport, }
})
} }
#[cfg(feature = "scylla")] #[cfg(feature = "scylla")]
pub async fn new_scylla( pub async fn new_scylla(
db: Arc<Session>, db: Arc<Session>,
transport: NatsTransport,
max_entries: u64, max_entries: u64,
ttl: Duration, ttl: Duration,
) -> anyhow::Result<Self> { ) -> anyhow::Result<Self> {
@@ -358,7 +385,6 @@ impl UsersShard {
stmt_partial_batch, stmt_partial_batch,
})), })),
caches: UserCaches::new(max_entries, ttl), caches: UserCaches::new(max_entries, ttl),
transport,
}) })
} }
@@ -366,12 +392,8 @@ impl UsersShard {
if user_id == FLUXER_SYSTEM_USER_ID { if user_id == FLUXER_SYSTEM_USER_ID {
return Ok(Some(fluxer_system_user())); return Ok(Some(fluxer_system_user()));
} }
let storage = self.storage.clone();
self.caches self.caches
.get_or_fetch_full( .get_or_fetch_full(user_id, self.storage.fetch_full_user(user_id))
user_id,
async move { storage.fetch_full_user(user_id).await },
)
.await .await
} }
@@ -379,17 +401,12 @@ impl UsersShard {
if user_id == FLUXER_SYSTEM_USER_ID { if user_id == FLUXER_SYSTEM_USER_ID {
return Ok(Some(fluxer_system_user().to_partial())); return Ok(Some(fluxer_system_user().to_partial()));
} }
let storage = self.storage.clone();
self.caches self.caches
.get_or_fetch_partial( .get_or_fetch_partial(user_id, self.storage.fetch_partial_user(user_id))
user_id,
async move { storage.fetch_partial_user(user_id).await },
)
.await .await
} }
async fn get_partial_users(&self, user_ids: Vec<i64>) -> anyhow::Result<Vec<UserPartial>> { async fn get_partial_users(&self, mut user_ids: Vec<i64>) -> anyhow::Result<Vec<UserPartial>> {
let mut user_ids = user_ids;
user_ids.sort_unstable(); user_ids.sort_unstable();
user_ids.dedup(); user_ids.dedup();
let mut partials = Vec::new(); let mut partials = Vec::new();
@@ -451,43 +468,40 @@ impl UsersShard {
} }
async fn fetch_partial_batch(&self, user_ids: Vec<i64>) -> anyhow::Result<Vec<UserPartial>> { async fn fetch_partial_batch(&self, user_ids: Vec<i64>) -> anyhow::Result<Vec<UserPartial>> {
let mut partials = Vec::new(); let mut cache_keys = user_ids
let user_ids = user_ids .iter()
.into_iter() .map(|&user_id| (user_id, self.caches.key(user_id)))
.filter(|user_id| { .collect::<HashMap<_, _>>();
if *user_id == FLUXER_SYSTEM_USER_ID {
partials.push(fluxer_system_user().to_partial());
false
} else {
true
}
})
.collect::<Vec<_>>();
if user_ids.is_empty() {
return Ok(partials);
}
let fetched_partials = match self.storage.fetch_partial_batch(user_ids.clone()).await { let fetched_partials = match self.storage.fetch_partial_batch(user_ids.clone()).await {
Ok(partials) => partials, Ok(fetched) => fetched,
Err(_) => { Err(err) => {
partials.extend(self.fetch_partial_batch_individually(user_ids).await?); tracing::warn!(
return Ok(partials); error = %err,
user_count = user_ids.len(),
"user partial batch read failed, retrying per user"
);
return self.fetch_partial_batch_individually(user_ids).await;
} }
}; };
let found_ids = fetched_partials let mut unexpected = 0usize;
.iter() let mut partials = Vec::with_capacity(fetched_partials.len());
.map(|partial| partial.user_id) for partial in fetched_partials {
.collect::<std::collections::HashSet<_>>(); let Some(key) = cache_keys.remove(&partial.user_id) else {
for partial in &fetched_partials { unexpected += 1;
self.caches continue;
.insert_partial(partial.user_id, Some(partial.clone())) };
.await; self.caches.insert_partial(key, Some(partial.clone())).await;
partials.push(partial);
} }
for user_id in user_ids { if unexpected > 0 {
if !found_ids.contains(&user_id) { tracing::warn!(
self.caches.insert_partial(user_id, None).await; unexpected,
} "user batch returned duplicate or unrequested users"
);
}
for key in cache_keys.into_values() {
self.caches.insert_partial(key, None).await;
} }
partials.extend(fetched_partials);
Ok(partials) Ok(partials)
} }
@@ -603,7 +617,7 @@ fn decode_postgres_user(row: serde_json::Value) -> anyhow::Result<User> {
fn decode_postgres_user_partial(row: serde_json::Value) -> anyhow::Result<UserPartial> { fn decode_postgres_user_partial(row: serde_json::Value) -> anyhow::Result<UserPartial> {
let row = postgres::decode_row_dates_as_millis(row)?; let row = postgres::decode_row_dates_as_millis(row)?;
let row: PartialUserKvRow = serde_json::from_value(row)?; let row: PartialUserDbRow = serde_json::from_value(row)?;
Ok(row.into()) Ok(row.into())
} }
@@ -678,34 +692,49 @@ impl ShardService for UsersShard {
"users" "users"
} }
fn render_prometheus_metrics(&self, output: &mut String) {
let _ = writeln!(
output,
"# TYPE fluxer_users_shard_cache_generation_bumps_total counter"
);
let _ = writeln!(
output,
"fluxer_users_shard_cache_generation_bumps_total {}",
self.caches.generation_bumps.load(Ordering::Relaxed)
);
let _ = writeln!(
output,
"# TYPE fluxer_users_shard_cache_generation_stripes gauge"
);
let _ = writeln!(
output,
"fluxer_users_shard_cache_generation_stripes {}",
self.caches.generations.len()
);
}
async fn handle(&self, request: UserRequest) -> anyhow::Result<UserResponse> { async fn handle(&self, request: UserRequest) -> anyhow::Result<UserResponse> {
match request { match request {
UserRequest::GetById { user_id } => match self.get_full_user(user_id).await? { UserRequest::GetById { user_id } => Ok(self
Some(user) => Ok(UserResponse::Found(user)), .get_full_user(user_id)
None => Ok(UserResponse::NotFound), .await?
}, .map_or(UserResponse::NotFound, UserResponse::Found)),
UserRequest::GetPartialById { user_id } => { UserRequest::GetPartialById { user_id } => Ok(self
match self.get_partial_user(user_id).await? { .get_partial_user(user_id)
Some(partial) => Ok(UserResponse::FoundPartial(partial)), .await?
None => Ok(UserResponse::NotFound), .map_or(UserResponse::NotFound, UserResponse::FoundPartial)),
}
}
UserRequest::GetPartialsByIds { user_ids } => Ok(UserResponse::FoundPartials( UserRequest::GetPartialsByIds { user_ids } => Ok(UserResponse::FoundPartials(
self.get_partial_users(user_ids).await?, self.get_partial_users(user_ids).await?,
)), )),
UserRequest::GetApiPartialById { user_id } => { UserRequest::GetApiPartialById { user_id } => Ok(self
match self.get_api_partial_user(user_id).await? { .get_api_partial_user(user_id)
Some(partial) => Ok(UserResponse::FoundApiPartial(partial)), .await?
None => Ok(UserResponse::NotFound), .map_or(UserResponse::NotFound, UserResponse::FoundApiPartial)),
}
}
UserRequest::GetApiPartialsByIds { user_ids } => Ok(UserResponse::FoundApiPartials( UserRequest::GetApiPartialsByIds { user_ids } => Ok(UserResponse::FoundApiPartials(
self.get_api_partial_users(user_ids).await?, self.get_api_partial_users(user_ids).await?,
)), )),
UserRequest::Invalidate { user_id } => { UserRequest::Invalidate { user_id } => {
self.caches.invalidate(user_id).await; self.caches.invalidate(user_id).await;
let subject = format!("svc.users.invalidate.{user_id}");
self.transport.publish(&subject, &[]).await?;
Ok(UserResponse::Invalidated) Ok(UserResponse::Invalidated)
} }
} }
@@ -734,7 +763,6 @@ fn optional_date_string(value: OptionalDate) -> Option<String> {
}) })
} }
#[cfg(feature = "scylla")]
impl From<PartialUserDbRow> for UserPartial { impl From<PartialUserDbRow> for UserPartial {
fn from(row: PartialUserDbRow) -> Self { fn from(row: PartialUserDbRow) -> Self {
Self { Self {
@@ -755,26 +783,6 @@ impl From<PartialUserDbRow> for UserPartial {
} }
} }
impl From<PartialUserKvRow> for UserPartial {
fn from(row: PartialUserKvRow) -> Self {
Self {
user_id: row.user_id,
username: row.username,
discriminator: row.discriminator,
global_name: row.global_name,
avatar_hash: row.avatar_hash,
bot: row.bot,
system: row.system,
flags: row.flags,
banner_hash: row.banner_hash,
banner_color: row.banner_color,
accent_color: row.accent_color,
avatar_color: row.avatar_color,
mention_flags: row.mention_flags,
}
}
}
#[cfg(feature = "scylla")] #[cfg(feature = "scylla")]
impl From<FullUserDbRow> for User { impl From<FullUserDbRow> for User {
fn from(row: FullUserDbRow) -> Self { fn from(row: FullUserDbRow) -> Self {
@@ -967,7 +975,7 @@ mod tests {
caches.get_partial(42).await.unwrap().unwrap().username, caches.get_partial(42).await.unwrap().unwrap().username,
"Ada" "Ada"
); );
assert!(caches.full.get(&42).await.is_none()); assert!(caches.full.get(&caches.key(42)).await.is_none());
} }
#[tokio::test] #[tokio::test]
@@ -1001,7 +1009,14 @@ mod tests {
assert_eq!(fetched.email.as_deref(), Some("[email protected]")); assert_eq!(fetched.email.as_deref(), Some("[email protected]"));
assert_eq!( assert_eq!(
caches.full.get(&42).await.unwrap().unwrap().bio.as_deref(), caches
.full
.get(&caches.key(42))
.await
.unwrap()
.unwrap()
.bio
.as_deref(),
Some("analytical engine enjoyer") Some("analytical engine enjoyer")
); );
let cached_partial = caches.get_partial(42).await.unwrap().unwrap(); let cached_partial = caches.get_partial(42).await.unwrap().unwrap();
@@ -1019,7 +1034,7 @@ mod tests {
.unwrap(); .unwrap();
assert!(fetched.is_none()); assert!(fetched.is_none());
assert!(matches!(caches.full.get(&42).await, Some(None))); assert!(matches!(caches.full.get(&caches.key(42)).await, Some(None)));
assert!(matches!(caches.get_partial(42).await, Some(None))); assert!(matches!(caches.get_partial(42).await, Some(None)));
} }
@@ -1030,12 +1045,12 @@ mod tests {
.get_or_fetch_full(42, async { Ok(Some(test_user(42))) }) .get_or_fetch_full(42, async { Ok(Some(test_user(42))) })
.await .await
.unwrap(); .unwrap();
assert!(caches.full.get(&42).await.is_some()); assert!(caches.full.get(&caches.key(42)).await.is_some());
assert!(caches.get_partial(42).await.is_some()); assert!(caches.get_partial(42).await.is_some());
caches.invalidate(42).await; caches.invalidate(42).await;
assert!(caches.full.get(&42).await.is_none()); assert!(caches.full.get(&caches.key(42)).await.is_none());
assert!(caches.get_partial(42).await.is_none()); assert!(caches.get_partial(42).await.is_none());
} }
@@ -1052,6 +1067,27 @@ mod tests {
assert!(caches.get_partial(42).await.is_none()); assert!(caches.get_partial(42).await.is_none());
} }
#[tokio::test]
async fn invalidation_counts_generation_bumps() {
let caches = caches();
assert_eq!(caches.generation_bumps.load(Ordering::Relaxed), 0);
caches.invalidate(42).await;
caches.invalidate(43).await;
assert_eq!(caches.generation_bumps.load(Ordering::Relaxed), 2);
}
#[test]
fn generation_stripes_follow_the_configured_capacity() {
assert_eq!(generation_stripes(16), USER_CACHE_MIN_GENERATION_STRIPES);
assert_eq!(generation_stripes(100_000), 131_072);
assert_eq!(
generation_stripes(u64::MAX),
USER_CACHE_MAX_GENERATION_STRIPES
);
}
#[cfg(feature = "scylla")] #[cfg(feature = "scylla")]
#[test] #[test]
fn partial_columns_match_the_user_partial_fields_exactly() { fn partial_columns_match_the_user_partial_fields_exactly() {
+2 -40
View File
@@ -159,24 +159,6 @@ impl User {
mention_flags: self.mention_flags, mention_flags: self.mention_flags,
} }
} }
pub fn to_api_partial(&self) -> ApiUserPartial {
if self.user_id == FLUXER_SYSTEM_USER_ID {
return fluxer_system_user();
}
ApiUserPartial {
id: self.user_id.to_string(),
username: self.username.clone(),
discriminator: format!("{:04}", self.discriminator),
global_name: self.global_name.clone(),
avatar: self.avatar_hash.clone(),
avatar_color: self.avatar_color,
bot: self.bot.filter(|bot| *bot),
system: self.system.filter(|system| *system),
flags: visible_user_flags(self.flags.unwrap_or_default()),
mention_flags: self.mention_flags.filter(|flags| *flags != 0),
}
}
} }
impl UserPartial { impl UserPartial {
@@ -320,30 +302,10 @@ mod tests {
} }
#[test] #[test]
fn direct_api_partial_matches_the_two_step_conversion() { fn api_partial_ignores_fields_outside_the_partial() {
for user_id in [123, FLUXER_SYSTEM_USER_ID] {
for flags in [
0,
USER_FLAG_STAFF,
USER_FLAG_STAFF | USER_FLAG_STAFF_HIDDEN | USER_FLAG_PARTNER,
USER_FLAG_DELETED,
] {
let user = user_with_flags(user_id, flags);
assert_eq!(
serde_json::to_value(user.to_api_partial()).unwrap(),
serde_json::to_value(user.to_partial().to_api_partial()).unwrap(),
"user_id {user_id} flags {flags}"
);
}
}
}
#[test]
fn direct_api_partial_ignores_fields_outside_the_partial() {
let mut user = user_with_flags(123, USER_FLAG_STAFF); let mut user = user_with_flags(123, USER_FLAG_STAFF);
user.mention_flags = Some(0); user.mention_flags = Some(0);
let api_partial = user.to_api_partial(); let api_partial = user.to_partial().to_api_partial();
assert_eq!(api_partial.id, "123"); assert_eq!(api_partial.id, "123");
assert_eq!(api_partial.username, "Ada"); assert_eq!(api_partial.username, "Ada");
+7 -3
View File
@@ -12,7 +12,12 @@
"tailwindcss", "tailwindcss",
"upng-js" "upng-js"
], ],
"ignoreFiles": ["fluxer_admin/src/styles/app.css", "fluxer_admin/static/htmx.min.js"], "ignoreFiles": [
"fluxer_admin/src/styles/app.css",
"fluxer_admin/static/htmx.min.js",
"tools/ci/templates/libfluxcore_wrapper.d.ts",
"tools/ci/templates/libfluxcore_wrapper.js"
],
"ignoreIssues": { "ignoreIssues": {
"fluxer_app/scripts/build/Config.tsx": ["exports"], "fluxer_app/scripts/build/Config.tsx": ["exports"],
"fluxer_app/scripts/build/rspack/static-files.mjs": ["exports"], "fluxer_app/scripts/build/rspack/static-files.mjs": ["exports"],
@@ -30,8 +35,7 @@
}, },
"workspaces": { "workspaces": {
".": { ".": {
"ignoreDependencies": ["vite"], "ignoreDependencies": ["vite"]
"project": ["scripts/**/*.mjs"]
}, },
"fluxer_api": { "fluxer_api": {
"ignoreDependencies": [ "ignoreDependencies": [
-1
View File
@@ -37,7 +37,6 @@
"lint": "cargo run -p fluxer-dev -- lint", "lint": "cargo run -p fluxer-dev -- lint",
"openapi:generate": "pnpm --filter @fluxer/openapi generate", "openapi:generate": "pnpm --filter @fluxer/openapi generate",
"openapi:validate": "pnpm --filter @fluxer/openapi validate", "openapi:validate": "pnpm --filter @fluxer/openapi validate",
"remote": "node scripts/remote.mjs",
"test": "cargo run -p fluxer-dev -- test", "test": "cargo run -p fluxer-dev -- test",
"typecheck": "cargo run -p fluxer-dev -- typecheck" "typecheck": "cargo run -p fluxer-dev -- typecheck"
}, },
+30 -12
View File
@@ -87,6 +87,9 @@ catalogs:
'@marsidev/react-turnstile': '@marsidev/react-turnstile':
specifier: 1.4.2 specifier: 1.4.2
version: 1.4.2 version: 1.4.2
'@mdx-js/mdx':
specifier: 3.1.1
version: 3.1.1
'@messageformat/core': '@messageformat/core':
specifier: 3.4.0 specifier: 3.4.0
version: 3.4.0 version: 3.4.0
@@ -138,6 +141,9 @@ catalogs:
'@types/luxon': '@types/luxon':
specifier: 3.7.1 specifier: 3.7.1
version: 3.7.1 version: 3.7.1
'@types/mdast':
specifier: 4.0.4
version: 4.0.4
'@types/node': '@types/node':
specifier: 25.3.0 specifier: 25.3.0
version: 25.3.0 version: 25.3.0
@@ -279,6 +285,9 @@ catalogs:
maxmind: maxmind:
specifier: 5.0.5 specifier: 5.0.5
version: 5.0.5 version: 5.0.5
mdast-util-mdx-jsx:
specifier: 3.2.0
version: 3.2.0
mdast-util-to-hast: mdast-util-to-hast:
specifier: 13.2.1 specifier: 13.2.1
version: 13.2.1 version: 13.2.1
@@ -363,6 +372,9 @@ catalogs:
react-hotkeys-hook: react-hotkeys-hook:
specifier: 5.2.4 specifier: 5.2.4
version: 5.2.4 version: 5.2.4
remark-gfm:
specifier: 4.0.1
version: 4.0.1
rxjs: rxjs:
specifier: 7.8.2 specifier: 7.8.2
version: 7.8.2 version: 7.8.2
@@ -411,6 +423,9 @@ catalogs:
undici: undici:
specifier: 7.29.0 specifier: 7.29.0
version: 7.29.0 version: 7.29.0
undici-types:
specifier: 7.18.2
version: 7.18.2
unique-names-generator: unique-names-generator:
specifier: 4.7.1 specifier: 4.7.1
version: 4.7.1 version: 4.7.1
@@ -597,9 +612,6 @@ importers:
'@pkgs/http_client': '@pkgs/http_client':
specifier: workspace:* specifier: workspace:*
version: link:pkgs/http_client version: link:pkgs/http_client
'@pkgs/initialization':
specifier: workspace:*
version: link:pkgs/initialization
'@pkgs/kv_client': '@pkgs/kv_client':
specifier: workspace:* specifier: workspace:*
version: link:pkgs/kv_client version: link:pkgs/kv_client
@@ -887,19 +899,13 @@ importers:
'@typescript/native-preview': '@typescript/native-preview':
specifier: 'catalog:' specifier: 'catalog:'
version: 7.0.0-dev.20260224.1 version: 7.0.0-dev.20260224.1
undici-types:
specifier: 'catalog:'
version: 7.18.2
vitest: vitest:
specifier: 'catalog:' specifier: 'catalog:'
version: 4.1.11(@opentelemetry/[email protected])(@types/[email protected])(@vitest/[email protected])([email protected])([email protected](@noble/[email protected]))([email protected](@types/[email protected])([email protected]))([email protected](@types/[email protected])([email protected])([email protected])([email protected])([email protected])) version: 4.1.11(@opentelemetry/[email protected])(@types/[email protected])(@vitest/[email protected])([email protected])([email protected](@noble/[email protected]))([email protected](@types/[email protected])([email protected]))([email protected](@types/[email protected])([email protected])([email protected])([email protected])([email protected]))
fluxer_api/pkgs/initialization:
devDependencies:
'@types/node':
specifier: 'catalog:'
version: 25.3.0
'@typescript/native-preview':
specifier: 'catalog:'
version: 7.0.0-dev.20260224.1
fluxer_api/pkgs/kv_client: fluxer_api/pkgs/kv_client:
dependencies: dependencies:
'@fluxer/constants': '@fluxer/constants':
@@ -1919,18 +1925,30 @@ importers:
'@fluxer/openapi': '@fluxer/openapi':
specifier: workspace:* specifier: workspace:*
version: link:../packages/openapi version: link:../packages/openapi
'@mdx-js/mdx':
specifier: 'catalog:'
version: 3.1.1
'@types/hast': '@types/hast':
specifier: 'catalog:' specifier: 'catalog:'
version: 3.0.5 version: 3.0.5
'@types/mdast':
specifier: 'catalog:'
version: 4.0.4
'@types/node': '@types/node':
specifier: 'catalog:' specifier: 'catalog:'
version: 25.3.0 version: 25.3.0
'@types/send': '@types/send':
specifier: 'catalog:' specifier: 'catalog:'
version: 1.2.1 version: 1.2.1
mdast-util-mdx-jsx:
specifier: 'catalog:'
version: 3.2.0
mdast-util-to-hast: mdast-util-to-hast:
specifier: 'catalog:' specifier: 'catalog:'
version: 13.2.1 version: 13.2.1
remark-gfm:
specifier: 'catalog:'
version: 4.0.1
typescript: typescript:
specifier: 6.0.2 specifier: 6.0.2
version: 6.0.2 version: 6.0.2
+5
View File
@@ -60,6 +60,7 @@ catalog:
'@lingui/swc-plugin': 5.11.0 '@lingui/swc-plugin': 5.11.0
'@livekit/components-react': 2.9.20 '@livekit/components-react': 2.9.20
'@marsidev/react-turnstile': 1.4.2 '@marsidev/react-turnstile': 1.4.2
'@mdx-js/mdx': 3.1.1
'@messageformat/core': 3.4.0 '@messageformat/core': 3.4.0
'@messageformat/parser': 5.1.1 '@messageformat/parser': 5.1.1
'@phosphor-icons/react': 2.1.10 '@phosphor-icons/react': 2.1.10
@@ -79,6 +80,7 @@ catalog:
'@types/hast': 3.0.5 '@types/hast': 3.0.5
'@types/lodash': 4.17.24 '@types/lodash': 4.17.24
'@types/luxon': 3.7.1 '@types/luxon': 3.7.1
'@types/mdast': 4.0.4
'@types/node': 25.3.0 '@types/node': 25.3.0
'@types/nodemailer': 7.0.11 '@types/nodemailer': 7.0.11
'@types/pg': 8.16.0 '@types/pg': 8.16.0
@@ -127,6 +129,7 @@ catalog:
luxon: 3.7.2 luxon: 3.7.2
match-sorter: 8.2.0 match-sorter: 8.2.0
maxmind: 5.0.5 maxmind: 5.0.5
mdast-util-mdx-jsx: 3.2.0
mdast-util-to-hast: 13.2.1 mdast-util-to-hast: 13.2.1
mime: 4.1.0 mime: 4.1.0
mobx: 6.15.0 mobx: 6.15.0
@@ -155,6 +158,7 @@ catalog:
react-dom: 19.2.4 react-dom: 19.2.4
react-hook-form: 7.71.2 react-hook-form: 7.71.2
react-hotkeys-hook: 5.2.4 react-hotkeys-hook: 5.2.4
remark-gfm: 4.0.1
rxjs: 7.8.2 rxjs: 7.8.2
send: 1.2.1 send: 1.2.1
sharp: 0.35.4 sharp: 0.35.4
@@ -171,6 +175,7 @@ catalog:
typescript-eslint: 8.70.0 typescript-eslint: 8.70.0
uint8array-extras: 1.5.0 uint8array-extras: 1.5.0
undici: 7.29.0 undici: 7.29.0
undici-types: 7.18.2
unique-names-generator: 4.7.1 unique-names-generator: 4.7.1
urlpattern-polyfill: 10.1.0 urlpattern-polyfill: 10.1.0
valibot: 1.4.2 valibot: 1.4.2
-370
View File
@@ -1,370 +0,0 @@
#!/usr/bin/env node
// SPDX-License-Identifier: AGPL-3.0-or-later
import {spawnSync} from 'node:child_process';
import {existsSync, mkdirSync, readFileSync} from 'node:fs';
import {homedir} from 'node:os';
import path from 'node:path';
import process from 'node:process';
const repoRoot = path.resolve(new URL('..', import.meta.url).pathname);
const defaultConfigPath = path.join(repoRoot, '.fluxer', 'remote-hosts.json');
const defaultControlDir = path.join(repoRoot, '.fluxer', 'remote-ssh');
function usage(exitCode = 0) {
console.log(`Usage:
pnpm remote list
pnpm remote doctor <host>
pnpm remote bootstrap <host> [--branch <branch>]
pnpm remote pull <host> [--branch <branch>]
pnpm remote run <host> -- <command>
pnpm remote apply-diff <host> [-- <path>...]
pnpm remote tunnel start|status|stop <host>
pnpm remote macos-setup [--pubkey ~/.ssh/id_ed25519.pub]
Global options:
--config <path> Host config path. Defaults to .fluxer/remote-hosts.json.
--verbose Print spawned ssh commands.
`);
process.exit(exitCode);
}
function fail(message, exitCode = 1) {
console.error(message);
process.exit(exitCode);
}
function parseGlobalArgs(argv) {
const options = {configPath: defaultConfigPath, verbose: false};
const rest = [];
for (let i = 0; i < argv.length; i++) {
const arg = argv[i];
if (arg === '--config') {
options.configPath = path.resolve(argv[++i] ?? fail('--config needs a path'));
} else if (arg === '--verbose') {
options.verbose = true;
} else if (arg === '--help' || arg === '-h') {
usage(0);
} else {
rest.push(arg);
}
}
return {options, args: rest};
}
function readJson(filePath) {
try {
return JSON.parse(readFileSync(filePath, 'utf8'));
} catch (error) {
fail(`Failed to read ${filePath}: ${error instanceof Error ? error.message : String(error)}`);
}
}
function loadConfig(configPath) {
if (!existsSync(configPath)) {
fail(
`Remote host config not found: ${configPath}\n` +
`Copy scripts/remote/hosts.example.json to .fluxer/remote-hosts.json and edit it.`,
);
}
const config = readJson(configPath);
if (!config || typeof config !== 'object' || !config.hosts || typeof config.hosts !== 'object') {
fail(`Invalid remote host config: ${configPath}`);
}
return config;
}
function getHost(config, name) {
const host = config.hosts[name];
if (!host || typeof host !== 'object') {
const known = Object.keys(config.hosts).sort().join(', ') || '(none)';
fail(`Unknown remote host "${name}". Known hosts: ${known}`);
}
if (typeof host.host !== 'string' || host.host.length === 0) {
fail(`Host "${name}" is missing "host"`);
}
if (host.platform !== 'windows' && host.platform !== 'macos' && host.platform !== 'linux') {
fail(`Host "${name}" needs platform "windows", "macos", or "linux"`);
}
return host;
}
function expandHome(value) {
if (typeof value !== 'string') return value;
if (value === '~') return homedir();
if (value.startsWith('~/')) return path.join(homedir(), value.slice(2));
return value;
}
function targetFor(host) {
return host.user ? `${host.user}@${host.host}` : host.host;
}
function controlPathFor(name, host) {
if (host.controlPath) return path.resolve(expandHome(host.controlPath));
mkdirSync(defaultControlDir, {recursive: true, mode: 0o700});
return path.join(defaultControlDir, `${name.replace(/[^A-Za-z0-9_.-]/g, '_')}.ctl`);
}
function sshBaseArgs(name, host, {control = true} = {}) {
const args = [];
if (host.port) args.push('-p', String(host.port));
if (host.identityFile) args.push('-i', path.resolve(expandHome(host.identityFile)));
args.push('-o', 'ServerAliveInterval=30', '-o', 'ServerAliveCountMax=3', '-o', 'StrictHostKeyChecking=accept-new');
if (control) {
args.push('-S', controlPathFor(name, host), '-o', 'ControlMaster=auto', '-o', 'ControlPersist=30m');
}
return args;
}
function forwardArgs(host) {
const forwards = Array.isArray(host.forwards) ? host.forwards : [];
const args = [];
for (const forward of forwards) {
if (typeof forward === 'string') {
args.push('-L', forward);
} else if (forward && typeof forward === 'object' && forward.local && forward.remote) {
args.push('-L', `${forward.local}:${forward.remote}`);
} else {
fail(`Invalid forward entry for ${host.host}`);
}
}
return args;
}
function shQuote(value) {
return `'${String(value).replace(/'/g, `'\\''`)}'`;
}
function psQuote(value) {
return `'${String(value).replace(/'/g, "''")}'`;
}
function remoteCommand(host, command) {
const repo = host.repo;
if (host.platform === 'windows') {
const body = `${repo ? `Set-Location -LiteralPath ${psQuote(repo)}; ` : ''}${command}`;
return `powershell.exe -NoProfile -ExecutionPolicy Bypass -EncodedCommand ${Buffer.from(body, 'utf16le').toString('base64')}`;
}
const body = `${repo ? `cd ${shQuote(repo)} && ` : ''}${command}`;
return `bash -lc ${shQuote(body)}`;
}
function spawnChecked(command, args, {verbose = false, input} = {}) {
if (verbose) console.error(`$ ${command} ${args.join(' ')}`);
const result = spawnSync(command, args, {
cwd: repoRoot,
input,
stdio: input === undefined ? 'inherit' : ['pipe', 'inherit', 'inherit'],
});
if (result.error) fail(`${command} failed to start: ${result.error.message}`);
if (result.status !== 0) process.exit(result.status ?? 1);
}
function spawnCapture(command, args) {
const result = spawnSync(command, args, {cwd: repoRoot, encoding: 'utf8'});
if (result.error || result.status !== 0) return null;
return result.stdout.trim();
}
function spawnCaptureRaw(command, args) {
const result = spawnSync(command, args, {cwd: repoRoot, encoding: 'utf8'});
if (result.error || result.status !== 0) return null;
return result.stdout;
}
function ssh(name, host, command, options = {}) {
spawnChecked('ssh', [...sshBaseArgs(name, host), targetFor(host), remoteCommand(host, command)], options);
}
function sshWithInput(name, host, command, input, options = {}) {
spawnChecked('ssh', [...sshBaseArgs(name, host), targetFor(host), remoteCommand(host, command)], {
...options,
input,
});
}
function currentBranch() {
return spawnCapture('git', ['rev-parse', '--abbrev-ref', 'HEAD']) || 'main';
}
function originUrl() {
return spawnCapture('git', ['config', '--get', 'remote.origin.url']) || '';
}
function splitAfterDoubleDash(args) {
const index = args.indexOf('--');
if (index === -1) return {head: args, tail: []};
return {head: args.slice(0, index), tail: args.slice(index + 1)};
}
function optionValue(args, name, fallback) {
const index = args.indexOf(name);
if (index === -1) return fallback;
const value = args[index + 1];
if (!value) fail(`${name} needs a value`);
return value;
}
function cloneOrUpdateCommand(host, branch) {
const repo = host.repo;
const repoUrl = host.repoUrl || originUrl();
if (!repo) fail('bootstrap/pull requires host.repo in the config');
if (!repoUrl) fail('bootstrap requires host.repoUrl or a local git remote.origin.url');
if (host.platform === 'windows') {
return [
`$repo = ${psQuote(repo)}`,
`$repoUrl = ${psQuote(repoUrl)}`,
`$branch = ${psQuote(branch)}`,
`if (!(Test-Path -LiteralPath (Join-Path $repo '.git'))) {`,
` New-Item -ItemType Directory -Force -Path (Split-Path -Parent $repo) | Out-Null`,
` git clone $repoUrl $repo`,
`}`,
`Set-Location -LiteralPath $repo`,
`git fetch --all --prune`,
`git checkout $branch`,
`git pull --ff-only`,
`corepack enable`,
`pnpm install`,
].join('; ');
}
return [
`repo=${shQuote(repo)}`,
`repo_url=${shQuote(repoUrl)}`,
`branch=${shQuote(branch)}`,
`if [ ! -d "$repo/.git" ]; then mkdir -p "$(dirname "$repo")"; git clone "$repo_url" "$repo"; fi`,
`cd "$repo"`,
`git fetch --all --prune`,
`git checkout "$branch"`,
`git pull --ff-only`,
`corepack enable`,
`pnpm install`,
].join(' && ');
}
function pullCommand(host, branch) {
const repo = host.repo;
if (!repo) fail('pull requires host.repo in the config');
if (host.platform === 'windows') {
return [
`Set-Location -LiteralPath ${psQuote(repo)}`,
`git fetch --all --prune`,
`git checkout ${psQuote(branch)}`,
`git pull --ff-only`,
].join('; ');
}
return [
`cd ${shQuote(repo)}`,
`git fetch --all --prune`,
`git checkout ${shQuote(branch)}`,
`git pull --ff-only`,
].join(' && ');
}
function macosSetup(args) {
const pubkeyPath = path.resolve(expandHome(optionValue(args, '--pubkey', '~/.ssh/id_ed25519.pub')));
const pubkey = existsSync(pubkeyPath) ? readFileSync(pubkeyPath, 'utf8').trim() : '<paste-your-public-ssh-key-here>';
console.log(`# Run this on the macOS machine while Tailscale is connected.
sudo systemsetup -setremotelogin on
mkdir -p ~/.ssh
chmod 700 ~/.ssh
printf '%s\\n' ${shQuote(pubkey)} >> ~/.ssh/authorized_keys
chmod 600 ~/.ssh/authorized_keys
tailscale status
tailscale ip -4
sudo lsof -nP -iTCP:22 -sTCP:LISTEN
# Optional, only if you intentionally use Tailscale SSH ACLs instead of OpenSSH keys:
# sudo tailscale up --ssh
`);
}
const {options, args} = parseGlobalArgs(process.argv.slice(2));
const command = args[0];
if (!command) usage(1);
if (command === 'macos-setup') {
macosSetup(args.slice(1));
process.exit(0);
}
const config = loadConfig(options.configPath);
switch (command) {
case 'list': {
for (const [name, host] of Object.entries(config.hosts)) {
console.log(`${name}\t${host.platform}\t${targetFor(host)}\t${host.repo ?? ''}`);
}
break;
}
case 'doctor': {
const name = args[1] ?? fail('doctor needs a host');
const host = getHost(config, name);
const remote =
host.platform === 'windows'
? '$PSVersionTable.PSVersion.ToString(); git --version; node --version; pnpm --version; rustc --version; cargo --version; Get-ComputerInfo | Select-Object -ExpandProperty OsName'
: 'uname -a; git --version; node --version; pnpm --version; rustc --version; cargo --version';
ssh(name, host, remote, options);
break;
}
case 'bootstrap': {
const name = args[1] ?? fail('bootstrap needs a host');
const host = getHost(config, name);
ssh(name, host, cloneOrUpdateCommand(host, optionValue(args, '--branch', currentBranch())), options);
break;
}
case 'pull': {
const name = args[1] ?? fail('pull needs a host');
const host = getHost(config, name);
ssh(name, host, pullCommand(host, optionValue(args, '--branch', currentBranch())), options);
break;
}
case 'run': {
const {head, tail} = splitAfterDoubleDash(args.slice(1));
const name = head[0] ?? fail('run needs a host');
if (tail.length === 0) fail('run needs -- <command>');
ssh(name, getHost(config, name), tail.join(' '), options);
break;
}
case 'apply-diff': {
const {head, tail} = splitAfterDoubleDash(args.slice(1));
const name = head[0] ?? fail('apply-diff needs a host');
const host = getHost(config, name);
const diff = spawnCaptureRaw('git', ['diff', '--binary', '--', ...tail]);
if (!diff) {
console.log('No local diff to apply.');
break;
}
sshWithInput(name, host, 'git apply --whitespace=nowarn -', diff, options);
break;
}
case 'tunnel': {
const action = args[1] ?? fail('tunnel needs start|status|stop');
const name = args[2] ?? fail('tunnel needs a host');
const host = getHost(config, name);
if (action === 'start') {
spawnChecked(
'ssh',
[
...sshBaseArgs(name, host),
...forwardArgs(host),
'-fN',
'-M',
'-o',
'ExitOnForwardFailure=yes',
targetFor(host),
],
options,
);
} else if (action === 'status') {
spawnChecked('ssh', [...sshBaseArgs(name, host), '-O', 'check', targetFor(host)], options);
} else if (action === 'stop') {
spawnChecked('ssh', [...sshBaseArgs(name, host), '-O', 'exit', targetFor(host)], options);
} else {
fail('tunnel action must be start, status, or stop');
}
break;
}
default:
usage(1);
}
-22
View File
@@ -1,22 +0,0 @@
{
"hosts": {
"windows-iva": {
"platform": "windows",
"host": "iva",
"user": "Hampus",
"repo": "C:\\Users\\Hampus\\src\\fluxer-private",
"repoUrl": "[email protected]:REPLACE_ME/fluxer-private.git",
"identityFile": "~/.ssh/id_ed25519",
"forwards": []
},
"macos": {
"platform": "macos",
"host": "replace-with-mac-tailnet-name",
"user": "replace-with-mac-username",
"repo": "/Users/replace-with-mac-username/src/fluxer-private",
"repoUrl": "[email protected]:REPLACE_ME/fluxer-private.git",
"identityFile": "~/.ssh/id_ed25519",
"forwards": []
}
}
}
+3 -16
View File
@@ -6,6 +6,7 @@ use crate::common::{
remove_dir_if_exists, require_env, resolve_calver, run_command, runner_temp, s3_client, remove_dir_if_exists, require_env, resolve_calver, run_command, runner_temp, s3_client,
trim_option, upload_s3_plan_append_only, trim_option, upload_s3_plan_append_only,
}; };
use crate::functions::sha256_reader;
use anyhow::{Context, Result, anyhow, ensure}; use anyhow::{Context, Result, anyhow, ensure};
use base64::Engine; use base64::Engine;
use base64::engine::general_purpose::STANDARD as BASE64; use base64::engine::general_purpose::STANDARD as BASE64;
@@ -13,11 +14,9 @@ use chrono::Utc;
use clap::{Args, ValueEnum}; use clap::{Args, ValueEnum};
use reqwest::Client; use reqwest::Client;
use serde_json::{Map, Value, json}; use serde_json::{Map, Value, json};
use sha2::{Digest, Sha256};
use std::collections::{BTreeMap, BTreeSet}; use std::collections::{BTreeMap, BTreeSet};
use std::env; use std::env;
use std::fs::{self, File}; use std::fs::{self, File};
use std::io::Read;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::sync::Arc; use std::sync::Arc;
use std::time::Duration; use std::time::Duration;
@@ -586,20 +585,8 @@ fn asset_tree_digests(root: &Path) -> Result<BTreeMap<String, String>> {
} }
fn sha256_file(path: &Path) -> Result<String> { fn sha256_file(path: &Path) -> Result<String> {
let mut file = let file = File::open(path).with_context(|| format!("Failed to open {}", path.display()))?;
File::open(path).with_context(|| format!("Failed to open {}", path.display()))?; sha256_reader(file).with_context(|| format!("Failed to read {}", path.display()))
let mut hasher = Sha256::new();
let mut buffer = [0u8; 64 * 1024];
loop {
let read = file
.read(&mut buffer)
.with_context(|| format!("Failed to read {}", path.display()))?;
if read == 0 {
break;
}
hasher.update(&buffer[..read]);
}
Ok(hex::encode(hasher.finalize()))
} }
fn tree_differences( fn tree_differences(
+4 -8
View File
@@ -1,8 +1,8 @@
// SPDX-License-Identifier: AGPL-3.0-or-later // SPDX-License-Identifier: AGPL-3.0-or-later
use crate::common::{ use crate::common::{
CALVER_SCHEME, CalverEnv, append_github_env, append_github_output, parse_version_instant, CALVER_SCHEME, CalverEnv, append_github_env, append_github_output, micro_segment,
resolve_calver, trim_option, month_day_segment, parse_version_instant, resolve_calver, trim_option,
}; };
use anyhow::Result; use anyhow::Result;
use chrono::{Datelike, Timelike, Utc}; use chrono::{Datelike, Timelike, Utc};
@@ -75,14 +75,10 @@ fn calver_outputs(version: &str) -> Result<CalverOutputs> {
instant.minute(), instant.minute(),
instant.second() instant.second()
); );
let micro = time
.parse::<u32>()
.expect("HHMMSS time segment should parse")
.to_string();
Ok(CalverOutputs { Ok(CalverOutputs {
version: version.to_string(), version: version.to_string(),
time, time,
micro, micro: micro_segment(instant),
date: format!( date: format!(
"{:04}{:02}{:02}", "{:04}{:02}{:02}",
instant.year(), instant.year(),
@@ -92,7 +88,7 @@ fn calver_outputs(version: &str) -> Result<CalverOutputs> {
year: instant.year().to_string(), year: instant.year().to_string(),
month: instant.month().to_string(), month: instant.month().to_string(),
day: format!("{:02}", instant.day()), day: format!("{:02}", instant.day()),
month_day: format!("{}{:02}", instant.month(), instant.day()), month_day: month_day_segment(instant),
}) })
} }
+17 -25
View File
@@ -365,20 +365,12 @@ fn format_calver(instant: DateTime<Utc>) -> String {
) )
} }
fn month_day_segment(instant: DateTime<Utc>) -> String { pub(crate) fn month_day_segment(instant: DateTime<Utc>) -> String {
format!("{}{:02}", instant.month(), instant.day()) format!("{}{:02}", instant.month(), instant.day())
} }
fn micro_segment(instant: DateTime<Utc>) -> String { pub(crate) fn micro_segment(instant: DateTime<Utc>) -> String {
format!( (instant.hour() * 10_000 + instant.minute() * 100 + instant.second()).to_string()
"{:02}{:02}{:02}",
instant.hour(),
instant.minute(),
instant.second()
)
.parse::<u32>()
.expect("HHMMSS time segment should parse")
.to_string()
} }
pub(crate) fn parse_version_instant(version: &str) -> Result<DateTime<Utc>> { pub(crate) fn parse_version_instant(version: &str) -> Result<DateTime<Utc>> {
@@ -1312,13 +1304,13 @@ pub(crate) fn collect_files(root: &Path) -> Result<Vec<PathBuf>> {
if !root.exists() { if !root.exists() {
return Ok(Vec::new()); return Ok(Vec::new());
} }
let mut files = WalkDir::new(root) let mut files = Vec::new();
.into_iter() for entry in WalkDir::new(root) {
.collect::<std::result::Result<Vec<_>, _>>()? let entry = entry?;
.into_iter() if entry.file_type().is_file() {
.filter(|entry| entry.file_type().is_file()) files.push(entry.into_path());
.map(|entry| entry.into_path()) }
.collect::<Vec<_>>(); }
files.sort(); files.sort();
Ok(files) Ok(files)
} }
@@ -1328,13 +1320,13 @@ pub(crate) fn count_files(root: &Path) -> Result<usize> {
} }
pub(crate) fn count_files_min_depth(root: &Path, min_depth: usize) -> Result<usize> { pub(crate) fn count_files_min_depth(root: &Path, min_depth: usize) -> Result<usize> {
Ok(WalkDir::new(root) let mut count = 0;
.min_depth(min_depth) for entry in WalkDir::new(root).min_depth(min_depth) {
.into_iter() if entry?.file_type().is_file() {
.collect::<std::result::Result<Vec<_>, _>>()? count += 1;
.into_iter() }
.filter(|entry| entry.file_type().is_file()) }
.count()) Ok(count)
} }
pub(crate) fn title_case(value: &str) -> String { pub(crate) fn title_case(value: &str) -> String {
+15 -1
View File
@@ -2,8 +2,9 @@
use anyhow::{Context, Result}; use anyhow::{Context, Result};
use serde::Serialize; use serde::Serialize;
use sha2::{Digest, Sha256};
use std::fs; use std::fs;
use std::io; use std::io::{self, Read};
use std::path::Path; use std::path::Path;
pub(crate) fn remove_file_if_exists(path: &Path) -> Result<()> { pub(crate) fn remove_file_if_exists(path: &Path) -> Result<()> {
@@ -32,6 +33,19 @@ pub(crate) fn write_json_pretty<T: Serialize + ?Sized>(path: &Path, value: &T) -
fs::write(path, bytes).with_context(|| format!("Failed to write {}", path.display())) fs::write(path, bytes).with_context(|| format!("Failed to write {}", path.display()))
} }
pub(crate) fn sha256_reader(mut reader: impl Read) -> io::Result<String> {
let mut hasher = Sha256::new();
let mut buffer = [0u8; 64 * 1024];
loop {
let read = reader.read(&mut buffer)?;
if read == 0 {
break;
}
hasher.update(&buffer[..read]);
}
Ok(hex::encode(hasher.finalize()))
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
+3 -15
View File
@@ -1,14 +1,13 @@
// SPDX-License-Identifier: AGPL-3.0-or-later // SPDX-License-Identifier: AGPL-3.0-or-later
use crate::common::{CommandSpec, output_text, parse_version_instant, run_command}; use crate::common::{CommandSpec, output_text, parse_version_instant, run_command};
use crate::functions::sha256_reader;
use anyhow::{Context, Result, bail, ensure}; use anyhow::{Context, Result, bail, ensure};
use chrono::{DateTime, Utc}; use chrono::{DateTime, Utc};
use clap::{Args, Subcommand}; use clap::{Args, Subcommand};
use serde::{Deserialize, Serialize}; use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::collections::{BTreeMap, BTreeSet}; use std::collections::{BTreeMap, BTreeSet};
use std::fs::{self, File}; use std::fs::{self, File};
use std::io::Read;
use std::path::{Path, PathBuf}; use std::path::{Path, PathBuf};
use std::thread; use std::thread;
use std::time::Duration; use std::time::Duration;
@@ -753,20 +752,9 @@ fn local_release_assets(
} }
fn sha256_file(path: &Path) -> Result<String> { fn sha256_file(path: &Path) -> Result<String> {
let mut file = File::open(path) let file = File::open(path)
.with_context(|| format!("Failed to open release asset {}", path.display()))?; .with_context(|| format!("Failed to open release asset {}", path.display()))?;
let mut hasher = Sha256::new(); sha256_reader(file).with_context(|| format!("Failed to read release asset {}", path.display()))
let mut buffer = [0u8; 64 * 1024];
loop {
let read = file
.read(&mut buffer)
.with_context(|| format!("Failed to read release asset {}", path.display()))?;
if read == 0 {
break;
}
hasher.update(&buffer[..read]);
}
Ok(hex::encode(hasher.finalize()))
} }
fn create_draft_release( fn create_draft_release(
-3
View File
@@ -46,9 +46,6 @@ fn find_import_start(source: &str) -> Option<usize> {
} }
offset += line.len(); offset += line.len();
} }
if source[offset..].starts_with("import ") {
return Some(offset);
}
None None
} }
+1 -3
View File
@@ -22,8 +22,6 @@ scylla = { version = "1.6.0", features = ["chrono-04"] }
serde = { version = "1.0.228", features = ["derive"] } serde = { version = "1.0.228", features = ["derive"] }
serde_json = "1.0.150" serde_json = "1.0.150"
sha2 = "0.11.0" sha2 = "0.11.0"
tempfile = "3.27.0"
tokio = { version = "1.52.3", features = ["fs", "io-util", "macros", "net", "process", "rt-multi-thread", "signal", "time"] } tokio = { version = "1.52.3", features = ["fs", "io-util", "macros", "net", "process", "rt-multi-thread", "signal", "time"] }
url = "2.5.8" url = "2.5.8"
[dev-dependencies]
tempfile = "3.27.0"
+40
View File
@@ -788,6 +788,10 @@
"name": "requested_at", "name": "requested_at",
"type": "timestamp" "type": "timestamp"
}, },
{
"name": "attempt_id",
"type": "text"
},
{ {
"name": "started_at", "name": "started_at",
"type": "timestamp" "type": "timestamp"
@@ -800,6 +804,10 @@
"name": "failed_at", "name": "failed_at",
"type": "timestamp" "type": "timestamp"
}, },
{
"name": "terminal_failed_at",
"type": "timestamp"
},
{ {
"name": "storage_key", "name": "storage_key",
"type": "text" "type": "text"
@@ -7893,6 +7901,10 @@
"name": "verification_token", "name": "verification_token",
"type": "text" "type": "text"
}, },
{
"name": "oauth_grant_id",
"type": "text"
},
{ {
"name": "verified_at", "name": "verified_at",
"type": "timestamp" "type": "timestamp"
@@ -7905,9 +7917,25 @@
"name": "created_at", "name": "created_at",
"type": "timestamp" "type": "timestamp"
}, },
{
"name": "revision",
"type": "text"
},
{ {
"name": "version", "name": "version",
"type": "int" "type": "int"
},
{
"name": "membership",
"type": "text"
},
{
"name": "credential_payload",
"type": "text"
},
{
"name": "credential_expires_at",
"type": "timestamp"
} }
], ],
"primary_key": "((user_id), connection_type, connection_id)", "primary_key": "((user_id), connection_type, connection_id)",
@@ -8172,6 +8200,10 @@
"name": "requested_at", "name": "requested_at",
"type": "timestamp" "type": "timestamp"
}, },
{
"name": "attempt_id",
"type": "text"
},
{ {
"name": "started_at", "name": "started_at",
"type": "timestamp" "type": "timestamp"
@@ -8184,6 +8216,10 @@
"name": "failed_at", "name": "failed_at",
"type": "timestamp" "type": "timestamp"
}, },
{
"name": "terminal_failed_at",
"type": "timestamp"
},
{ {
"name": "storage_key", "name": "storage_key",
"type": "text" "type": "text"
@@ -8562,6 +8598,10 @@
"name": "pending_deletion_at", "name": "pending_deletion_at",
"type": "timestamp" "type": "timestamp"
}, },
{
"name": "deletion_started_at",
"type": "timestamp"
},
{ {
"name": "pending_bulk_message_deletion_at", "name": "pending_bulk_message_deletion_at",
"type": "timestamp" "type": "timestamp"
+1 -13
View File
@@ -79,19 +79,7 @@ pub fn which(name: &str) -> Option<PathBuf> {
} }
fn is_writable(path: &Path) -> bool { fn is_writable(path: &Path) -> bool {
let probe = path.join(".fluxer-write-test"); tempfile::NamedTempFile::new_in(path).is_ok()
match std::fs::OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(&probe)
{
Ok(_) => {
let _ = std::fs::remove_file(probe);
true
}
Err(_) => false,
}
} }
#[cfg(unix)] #[cfg(unix)]
+14 -21
View File
@@ -130,7 +130,7 @@ pub fn run_command(args: &[&str], options: RunOptions<'_>) -> Result<Output> {
.current_dir(options.cwd) .current_dir(options.cwd)
.env_clear() .env_clear()
.envs(env); .envs(env);
if options.capture { let output = if options.capture {
command.stdout(Stdio::piped()).stderr(Stdio::piped()); command.stdout(Stdio::piped()).stderr(Stdio::piped());
let output = command let output = command
.output() .output()
@@ -142,27 +142,20 @@ pub fn run_command(args: &[&str], options: RunOptions<'_>) -> Result<Output> {
if !text.trim_end().is_empty() { if !text.trim_end().is_empty() {
println!("{}", text.trim_end()); println!("{}", text.trim_end());
} }
if options.check && !output.status.success() { output
let code = output.status.code().unwrap_or(-1); } else {
bail!( command
"Command failed with exit code {code}: {}", .stdin(Stdio::inherit())
format_command(args) .stdout(Stdio::inherit())
); .stderr(Stdio::inherit());
let status = command
.status()
.with_context(|| format!("failed to run {}", format_command(args)))?;
Output {
status,
stdout: Vec::new(),
stderr: Vec::new(),
} }
return Ok(output);
}
command
.stdin(Stdio::inherit())
.stdout(Stdio::inherit())
.stderr(Stdio::inherit());
let status = command
.status()
.with_context(|| format!("failed to run {}", format_command(args)))?;
let output = Output {
status,
stdout: Vec::new(),
stderr: Vec::new(),
}; };
if options.check && !output.status.success() { if options.check && !output.status.success() {
let code = output.status.code().unwrap_or(-1); let code = output.status.code().unwrap_or(-1);