diff --git a/Cargo.lock b/Cargo.lock index d0702e09d4..70a8af4e7d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4296,6 +4296,7 @@ dependencies = [ "camino-tempfile", "chrono", "clap", + "compute_api", "consumption_metrics", "dashmap", "ecdsa 0.16.9", diff --git a/libs/compute_api/src/spec.rs b/libs/compute_api/src/spec.rs index 883c624f71..525a1572ff 100644 --- a/libs/compute_api/src/spec.rs +++ b/libs/compute_api/src/spec.rs @@ -268,6 +268,22 @@ pub struct GenericOption { /// declare a `trait` on it. pub type GenericOptions = Option>; +/// Configured the local-proxy application with the relevant JWKS and roles it should +/// use for authorizing connect requests using JWT. +#[derive(Clone, Debug, Deserialize, Serialize)] +pub struct LocalProxySpec { + pub jwks: Vec, +} + +#[derive(Clone, Debug, Deserialize, Serialize)] +pub struct JwksSettings { + pub id: String, + pub role_names: Vec, + pub jwks_url: String, + pub provider_name: String, + pub jwt_audience: Option, +} + #[cfg(test)] mod tests { use super::*; diff --git a/proxy/Cargo.toml b/proxy/Cargo.toml index 501ce050e0..04e0f9d4f5 100644 --- a/proxy/Cargo.toml +++ b/proxy/Cargo.toml @@ -24,6 +24,7 @@ bytes = { workspace = true, features = ["serde"] } camino.workspace = true chrono.workspace = true clap.workspace = true +compute_api.workspace = true consumption_metrics.workspace = true dashmap.workspace = true env_logger.workspace = true diff --git a/proxy/src/auth/backend/jwt.rs b/proxy/src/auth/backend/jwt.rs index 94e5999a5f..ab848551a9 100644 --- a/proxy/src/auth/backend/jwt.rs +++ b/proxy/src/auth/backend/jwt.rs @@ -12,7 +12,10 @@ use serde::{Deserialize, Deserializer}; use signature::Verifier; use tokio::time::Instant; -use crate::{context::RequestMonitoring, http::parse_json_body_with_limit, EndpointId, RoleName}; +use crate::{ + context::RequestMonitoring, http::parse_json_body_with_limit, intern::RoleNameInt, EndpointId, + RoleName, +}; // TODO(conrad): make these configurable. const CLOCK_SKEW_LEEWAY: Duration = Duration::from_secs(30); @@ -27,7 +30,6 @@ pub(crate) trait FetchAuthRules: Clone + Send + Sync + 'static { &self, ctx: &RequestMonitoring, endpoint: EndpointId, - role_name: RoleName, ) -> impl Future>> + Send; } @@ -35,10 +37,11 @@ pub(crate) struct AuthRule { pub(crate) id: String, pub(crate) jwks_url: url::Url, pub(crate) audience: Option, + pub(crate) role_names: Vec, } #[derive(Default)] -pub(crate) struct JwkCache { +pub struct JwkCache { client: reqwest::Client, map: DashMap<(EndpointId, RoleName), Arc>, @@ -54,18 +57,28 @@ pub(crate) struct JwkCacheEntry { } impl JwkCacheEntry { - fn find_jwk_and_audience(&self, key_id: &str) -> Option<(&jose_jwk::Jwk, Option<&str>)> { - self.key_sets.values().find_map(|key_set| { - key_set - .find_key(key_id) - .map(|jwk| (jwk, key_set.audience.as_deref())) - }) + fn find_jwk_and_audience( + &self, + key_id: &str, + role_name: &RoleName, + ) -> Option<(&jose_jwk::Jwk, Option<&str>)> { + self.key_sets + .values() + // make sure our requested role has access to the key set + .filter(|key_set| key_set.role_names.iter().any(|role| **role == **role_name)) + // try and find the requested key-id in the key set + .find_map(|key_set| { + key_set + .find_key(key_id) + .map(|jwk| (jwk, key_set.audience.as_deref())) + }) } } struct KeySet { jwks: jose_jwk::JwkSet, audience: Option, + role_names: Vec, } impl KeySet { @@ -106,7 +119,6 @@ impl JwkCacheEntryLock { ctx: &RequestMonitoring, client: &reqwest::Client, endpoint: EndpointId, - role_name: RoleName, auth_rules: &F, ) -> anyhow::Result> { // double check that no one beat us to updating the cache. @@ -119,11 +131,10 @@ impl JwkCacheEntryLock { } } - let rules = auth_rules - .fetch_auth_rules(ctx, endpoint, role_name) - .await?; + let rules = auth_rules.fetch_auth_rules(ctx, endpoint).await?; let mut key_sets = ahash::HashMap::with_capacity_and_hasher(rules.len(), ahash::RandomState::new()); + // TODO(conrad): run concurrently // TODO(conrad): strip the JWKs urls (should be checked by cplane as well - cloud#16284) for rule in rules { @@ -151,6 +162,7 @@ impl JwkCacheEntryLock { KeySet { jwks, audience: rule.audience, + role_names: rule.role_names, }, ); } @@ -173,7 +185,6 @@ impl JwkCacheEntryLock { ctx: &RequestMonitoring, client: &reqwest::Client, endpoint: EndpointId, - role_name: RoleName, fetch: &F, ) -> Result, anyhow::Error> { let now = Instant::now(); @@ -183,9 +194,7 @@ impl JwkCacheEntryLock { let Some(cached) = guard else { let _paused = ctx.latency_timer_pause(crate::metrics::Waiting::Compute); let permit = self.acquire_permit().await; - return self - .renew_jwks(permit, ctx, client, endpoint, role_name, fetch) - .await; + return self.renew_jwks(permit, ctx, client, endpoint, fetch).await; }; let last_update = now.duration_since(cached.last_retrieved); @@ -196,9 +205,7 @@ impl JwkCacheEntryLock { let permit = self.acquire_permit().await; // it's been too long since we checked the keys. wait for them to update. - return self - .renew_jwks(permit, ctx, client, endpoint, role_name, fetch) - .await; + return self.renew_jwks(permit, ctx, client, endpoint, fetch).await; } // every 5 minutes we should spawn a job to eagerly update the token. @@ -212,7 +219,7 @@ impl JwkCacheEntryLock { let ctx = ctx.clone(); tokio::spawn(async move { if let Err(e) = entry - .renew_jwks(permit, &ctx, &client, endpoint, role_name, &fetch) + .renew_jwks(permit, &ctx, &client, endpoint, &fetch) .await { tracing::warn!(error=?e, "could not fetch JWKs in background job"); @@ -232,7 +239,7 @@ impl JwkCacheEntryLock { jwt: &str, client: &reqwest::Client, endpoint: EndpointId, - role_name: RoleName, + role_name: &RoleName, fetch: &F, ) -> Result<(), anyhow::Error> { // JWT compact form is defined to be @@ -254,30 +261,26 @@ impl JwkCacheEntryLock { let sig = base64::decode_config(signature, base64::URL_SAFE_NO_PAD) .context("Provided authentication token is not a valid JWT encoding")?; - ensure!(header.typ == "JWT"); + ensure!( + header.typ == "JWT", + "Provided authentication token is not a valid JWT encoding" + ); let kid = header.key_id.context("missing key id")?; let mut guard = self - .get_or_update_jwk_cache(ctx, client, endpoint.clone(), role_name.clone(), fetch) + .get_or_update_jwk_cache(ctx, client, endpoint.clone(), fetch) .await?; // get the key from the JWKs if possible. If not, wait for the keys to update. let (jwk, expected_audience) = loop { - match guard.find_jwk_and_audience(kid) { + match guard.find_jwk_and_audience(kid, role_name) { Some(jwk) => break jwk, None if guard.last_retrieved.elapsed() > MIN_RENEW => { let _paused = ctx.latency_timer_pause(crate::metrics::Waiting::Compute); let permit = self.acquire_permit().await; guard = self - .renew_jwks( - permit, - ctx, - client, - endpoint.clone(), - role_name.clone(), - fetch, - ) + .renew_jwks(permit, ctx, client, endpoint.clone(), fetch) .await?; } _ => { @@ -320,11 +323,14 @@ impl JwkCacheEntryLock { let now = SystemTime::now(); if let Some(exp) = payload.expiration { - ensure!(now < exp + CLOCK_SKEW_LEEWAY); + ensure!(now < exp + CLOCK_SKEW_LEEWAY, "JWT token has expired"); } if let Some(nbf) = payload.not_before { - ensure!(nbf < now + CLOCK_SKEW_LEEWAY); + ensure!( + nbf < now + CLOCK_SKEW_LEEWAY, + "JWT token is not yet ready to use" + ); } Ok(()) @@ -336,7 +342,7 @@ impl JwkCache { &self, ctx: &RequestMonitoring, endpoint: EndpointId, - role_name: RoleName, + role_name: &RoleName, fetch: &F, jwt: &str, ) -> Result<(), anyhow::Error> { @@ -572,7 +578,7 @@ mod tests { format!("{header}.{body}") } - fn new_ec_jwt(kid: String, key: p256::SecretKey) -> String { + fn new_ec_jwt(kid: String, key: &p256::SecretKey) -> String { use p256::ecdsa::{Signature, SigningKey}; let payload = build_jwt_payload(kid, jose_jwa::Signing::Es256); @@ -660,11 +666,6 @@ X0n5X2/pBLJzxZc62ccvZYVnctBiFs6HbSnxpuMQCfkt/BcR/ttIepBQQIW86wHL let (ec1, jwk3) = new_ec_jwk("3".into()); let (ec2, jwk4) = new_ec_jwk("4".into()); - let jwt1 = new_rsa_jwt("1".into(), rs1); - let jwt2 = new_rsa_jwt("2".into(), rs2); - let jwt3 = new_ec_jwt("3".into(), ec1); - let jwt4 = new_ec_jwt("4".into(), ec2); - let foo_jwks = jose_jwk::JwkSet { keys: vec![jwk1, jwk3], }; @@ -706,47 +707,98 @@ X0n5X2/pBLJzxZc62ccvZYVnctBiFs6HbSnxpuMQCfkt/BcR/ttIepBQQIW86wHL let client = reqwest::Client::new(); #[derive(Clone)] - struct Fetch(SocketAddr); + struct Fetch(SocketAddr, Vec); impl FetchAuthRules for Fetch { async fn fetch_auth_rules( &self, _ctx: &RequestMonitoring, _endpoint: EndpointId, - _role_name: RoleName, ) -> anyhow::Result> { Ok(vec![ AuthRule { id: "foo".to_owned(), jwks_url: format!("http://{}/foo", self.0).parse().unwrap(), audience: None, + role_names: self.1.clone(), }, AuthRule { id: "bar".to_owned(), jwks_url: format!("http://{}/bar", self.0).parse().unwrap(), audience: None, + role_names: self.1.clone(), }, ]) } } - let role_name = RoleName::from("user"); + let role_name1 = RoleName::from("anonymous"); + let role_name2 = RoleName::from("authenticated"); + + let fetch = Fetch( + addr, + vec![ + RoleNameInt::from(&role_name1), + RoleNameInt::from(&role_name2), + ], + ); + let endpoint = EndpointId::from("ep"); let jwk_cache = Arc::new(JwkCacheEntryLock::default()); - for token in [jwt1, jwt2, jwt3, jwt4] { - jwk_cache - .check_jwt( - &RequestMonitoring::test(), - &token, - &client, - endpoint.clone(), - role_name.clone(), - &Fetch(addr), - ) - .await - .unwrap(); + let jwt1 = new_rsa_jwt("1".into(), rs1); + let jwt2 = new_rsa_jwt("2".into(), rs2); + let jwt3 = new_ec_jwt("3".into(), &ec1); + let jwt4 = new_ec_jwt("4".into(), &ec2); + + // had the wrong kid, therefore will have the wrong ecdsa signature + let bad_jwt = new_ec_jwt("3".into(), &ec2); + // this role_name is not accepted + let bad_role_name = RoleName::from("cloud_admin"); + + let err = jwk_cache + .check_jwt( + &RequestMonitoring::test(), + &bad_jwt, + &client, + endpoint.clone(), + &role_name1, + &fetch, + ) + .await + .unwrap_err(); + assert!(err.to_string().contains("signature error")); + + let err = jwk_cache + .check_jwt( + &RequestMonitoring::test(), + &jwt1, + &client, + endpoint.clone(), + &bad_role_name, + &fetch, + ) + .await + .unwrap_err(); + assert!(err.to_string().contains("jwk not found")); + + let tokens = [jwt1, jwt2, jwt3, jwt4]; + let role_names = [role_name1, role_name2]; + for role in &role_names { + for token in &tokens { + jwk_cache + .check_jwt( + &RequestMonitoring::test(), + token, + &client, + endpoint.clone(), + role, + &fetch, + ) + .await + .unwrap(); + } } } } diff --git a/proxy/src/auth/backend/local.rs b/proxy/src/auth/backend/local.rs index 2ff2ca00f0..2ab53f2c6a 100644 --- a/proxy/src/auth/backend/local.rs +++ b/proxy/src/auth/backend/local.rs @@ -1,4 +1,4 @@ -use std::{collections::HashMap, net::SocketAddr}; +use std::net::SocketAddr; use anyhow::Context; use arc_swap::ArcSwapOption; @@ -10,8 +10,8 @@ use crate::{ NodeInfo, }, context::RequestMonitoring, - intern::{BranchIdInt, BranchIdTag, EndpointIdTag, InternId, ProjectIdInt, ProjectIdTag}, - EndpointId, RoleName, + intern::{BranchIdTag, EndpointIdTag, InternId, ProjectIdTag}, + EndpointId, }; use super::jwt::{AuthRule, FetchAuthRules, JwkCache}; @@ -48,26 +48,17 @@ impl LocalBackend { #[derive(Clone, Copy)] pub(crate) struct StaticAuthRules; -pub static JWKS_ROLE_MAP: ArcSwapOption = ArcSwapOption::const_empty(); - -#[derive(Debug, Clone)] -pub struct JwksRoleSettings { - pub roles: HashMap, - pub project_id: ProjectIdInt, - pub branch_id: BranchIdInt, -} +pub static JWKS_ROLE_MAP: ArcSwapOption = ArcSwapOption::const_empty(); impl FetchAuthRules for StaticAuthRules { async fn fetch_auth_rules( &self, _ctx: &RequestMonitoring, _endpoint: EndpointId, - role_name: RoleName, ) -> anyhow::Result> { let mappings = JWKS_ROLE_MAP.load(); let role_mappings = mappings .as_deref() - .and_then(|m| m.roles.get(&role_name)) .context("JWKs settings for this role were not configured")?; let mut rules = vec![]; for setting in &role_mappings.jwks { @@ -75,6 +66,7 @@ impl FetchAuthRules for StaticAuthRules { id: setting.id.clone(), jwks_url: setting.jwks_url.clone(), audience: setting.jwt_audience.clone(), + role_names: setting.role_names.clone(), }); } diff --git a/proxy/src/bin/local_proxy.rs b/proxy/src/bin/local_proxy.rs index 94365ddf05..1b3f465686 100644 --- a/proxy/src/bin/local_proxy.rs +++ b/proxy/src/bin/local_proxy.rs @@ -1,34 +1,35 @@ -use std::{ - net::SocketAddr, - path::{Path, PathBuf}, - pin::pin, - sync::Arc, - time::Duration, -}; +use std::{net::SocketAddr, pin::pin, str::FromStr, sync::Arc, time::Duration}; -use anyhow::{bail, ensure}; +use anyhow::{bail, ensure, Context}; +use camino::{Utf8Path, Utf8PathBuf}; +use compute_api::spec::LocalProxySpec; use dashmap::DashMap; -use futures::{future::Either, FutureExt}; +use futures::future::Either; use proxy::{ - auth::backend::local::{JwksRoleSettings, LocalBackend, JWKS_ROLE_MAP}, + auth::backend::local::{LocalBackend, JWKS_ROLE_MAP}, cancellation::CancellationHandlerMain, config::{self, AuthenticationConfig, HttpConfig, ProxyConfig, RetryConfig}, - console::{locks::ApiLocks, messages::JwksRoleMapping}, + console::{ + locks::ApiLocks, + messages::{EndpointJwksResponse, JwksSettings}, + }, http::health_server::AppMetrics, + intern::RoleNameInt, metrics::{Metrics, ThreadPoolMetrics}, rate_limiter::{BucketRateLimiter, EndpointRateLimiter, LeakyBucketConfig, RateBucketInfo}, scram::threadpool::ThreadPool, serverless::{self, cancel_set::CancelSet, GlobalConnPoolOptions}, + RoleName, }; project_git_version!(GIT_VERSION); project_build_tag!(BUILD_TAG); use clap::Parser; -use tokio::{net::TcpListener, task::JoinSet}; +use tokio::{net::TcpListener, sync::Notify, task::JoinSet}; use tokio_util::sync::CancellationToken; use tracing::{error, info, warn}; -use utils::{project_build_tag, project_git_version, sentry_init::init_sentry}; +use utils::{pid_file, project_build_tag, project_git_version, sentry_init::init_sentry}; #[global_allocator] static GLOBAL: tikv_jemallocator::Jemalloc = tikv_jemallocator::Jemalloc; @@ -72,9 +73,12 @@ struct LocalProxyCliArgs { /// Address of the postgres server #[clap(long, default_value = "127.0.0.1:5432")] compute: SocketAddr, - /// File address of the local proxy config file + /// Path of the local proxy config file #[clap(long, default_value = "./localproxy.json")] - config_path: PathBuf, + config_path: Utf8PathBuf, + /// Path of the local proxy PID file + #[clap(long, default_value = "./localproxy.pid")] + pid_path: Utf8PathBuf, } #[derive(clap::Args, Clone, Copy, Debug)] @@ -126,6 +130,24 @@ async fn main() -> anyhow::Result<()> { let args = LocalProxyCliArgs::parse(); let config = build_config(&args)?; + // before we bind to any ports, write the process ID to a file + // so that compute-ctl can find our process later + // in order to trigger the appropriate SIGHUP on config change. + // + // This also claims a "lock" that makes sure only one instance + // of local-proxy runs at a time. + let _process_guard = loop { + match pid_file::claim_for_current_process(&args.pid_path) { + Ok(guard) => break guard, + Err(e) => { + // compute-ctl might have tried to read the pid-file to let us + // know about some config change. We should try again. + error!(path=?args.pid_path, "could not claim PID file guard: {e:?}"); + tokio::time::sleep(Duration::from_secs(1)).await; + } + } + }; + let metrics_listener = TcpListener::bind(args.metrics).await?.into_std()?; let http_listener = TcpListener::bind(args.http).await?; let shutdown = CancellationToken::new(); @@ -139,12 +161,30 @@ async fn main() -> anyhow::Result<()> { 16, )); - refresh_config(args.config_path.clone()).await; + // write the process ID to a file so that compute-ctl can find our process later + // in order to trigger the appropriate SIGHUP on config change. + let pid = std::process::id(); + info!("process running in PID {pid}"); + std::fs::write(args.pid_path, format!("{pid}\n")).context("writing PID to file")?; let mut maintenance_tasks = JoinSet::new(); - maintenance_tasks.spawn(proxy::handle_signals(shutdown.clone(), move || { - refresh_config(args.config_path.clone()).map(Ok) + + let refresh_config_notify = Arc::new(Notify::new()); + maintenance_tasks.spawn(proxy::handle_signals(shutdown.clone(), { + let refresh_config_notify = Arc::clone(&refresh_config_notify); + move || { + refresh_config_notify.notify_one(); + } })); + + // trigger the first config load **after** setting up the signal hook + // to avoid the race condition where: + // 1. No config file registered when local-proxy starts up + // 2. The config file is written but the signal hook is not yet received + // 3. local-proxy completes startup but has no config loaded, despite there being a registerd config. + refresh_config_notify.notify_one(); + tokio::spawn(refresh_config_loop(args.config_path, refresh_config_notify)); + maintenance_tasks.spawn(proxy::http::health_server::task_main( metrics_listener, AppMetrics { @@ -245,81 +285,84 @@ fn build_config(args: &LocalProxyCliArgs) -> anyhow::Result<&'static ProxyConfig }))) } -async fn refresh_config(path: PathBuf) { - match refresh_config_inner(&path).await { - Ok(()) => {} - Err(e) => { - error!(error=?e, ?path, "could not read config file"); +async fn refresh_config_loop(path: Utf8PathBuf, rx: Arc) { + loop { + rx.notified().await; + + match refresh_config_inner(&path).await { + Ok(()) => {} + Err(e) => { + error!(error=?e, ?path, "could not read config file"); + } } } } -async fn refresh_config_inner(path: &Path) -> anyhow::Result<()> { +async fn refresh_config_inner(path: &Utf8Path) -> anyhow::Result<()> { let bytes = tokio::fs::read(&path).await?; - let mut data: JwksRoleMapping = serde_json::from_slice(&bytes)?; + let data: LocalProxySpec = serde_json::from_slice(&bytes)?; - let mut settings = None; + let mut jwks_set = vec![]; - for mapping in data.roles.values_mut() { - for jwks in &mut mapping.jwks { - ensure!( - jwks.jwks_url.has_authority() - && (jwks.jwks_url.scheme() == "http" || jwks.jwks_url.scheme() == "https"), - "Invalid JWKS url. Must be HTTP", - ); + for jwks in data.jwks { + let mut jwks_url = url::Url::from_str(&jwks.jwks_url).context("parsing JWKS url")?; - ensure!( - jwks.jwks_url - .host() - .is_some_and(|h| h != url::Host::Domain("")), - "Invalid JWKS url. No domain listed", - ); + ensure!( + jwks_url.has_authority() + && (jwks_url.scheme() == "http" || jwks_url.scheme() == "https"), + "Invalid JWKS url. Must be HTTP", + ); - // clear username, password and ports - jwks.jwks_url.set_username("").expect( + ensure!( + jwks_url.host().is_some_and(|h| h != url::Host::Domain("")), + "Invalid JWKS url. No domain listed", + ); + + // clear username, password and ports + jwks_url + .set_username("") + .expect("url can be a base and has a valid host and is not a file. should not error"); + jwks_url + .set_password(None) + .expect("url can be a base and has a valid host and is not a file. should not error"); + // local testing is hard if we need to have a specific restricted port + if cfg!(not(feature = "testing")) { + jwks_url.set_port(None).expect( "url can be a base and has a valid host and is not a file. should not error", ); - jwks.jwks_url.set_password(None).expect( - "url can be a base and has a valid host and is not a file. should not error", - ); - // local testing is hard if we need to have a specific restricted port - if cfg!(not(feature = "testing")) { - jwks.jwks_url.set_port(None).expect( - "url can be a base and has a valid host and is not a file. should not error", - ); - } - - // clear query params - jwks.jwks_url.set_fragment(None); - jwks.jwks_url.query_pairs_mut().clear().finish(); - - if jwks.jwks_url.scheme() != "https" { - // local testing is hard if we need to set up https support. - if cfg!(not(feature = "testing")) { - jwks.jwks_url - .set_scheme("https") - .expect("should not error to set the scheme to https if it was http"); - } else { - warn!(scheme = jwks.jwks_url.scheme(), "JWKS url is not HTTPS"); - } - } - - let (pr, br) = settings.get_or_insert((jwks.project_id, jwks.branch_id)); - ensure!( - *pr == jwks.project_id, - "inconsistent project IDs configured" - ); - ensure!(*br == jwks.branch_id, "inconsistent branch IDs configured"); } + + // clear query params + jwks_url.set_fragment(None); + jwks_url.query_pairs_mut().clear().finish(); + + if jwks_url.scheme() != "https" { + // local testing is hard if we need to set up https support. + if cfg!(not(feature = "testing")) { + jwks_url + .set_scheme("https") + .expect("should not error to set the scheme to https if it was http"); + } else { + warn!(scheme = jwks_url.scheme(), "JWKS url is not HTTPS"); + } + } + + jwks_set.push(JwksSettings { + id: jwks.id, + jwks_url, + provider_name: jwks.provider_name, + jwt_audience: jwks.jwt_audience, + role_names: jwks + .role_names + .into_iter() + .map(RoleName::from) + .map(|s| RoleNameInt::from(&s)) + .collect(), + }) } - if let Some((project_id, branch_id)) = settings { - JWKS_ROLE_MAP.store(Some(Arc::new(JwksRoleSettings { - roles: data.roles, - project_id, - branch_id, - }))); - } + info!("successfully loaded new config"); + JWKS_ROLE_MAP.store(Some(Arc::new(EndpointJwksResponse { jwks: jwks_set }))); Ok(()) } diff --git a/proxy/src/bin/pg_sni_router.rs b/proxy/src/bin/pg_sni_router.rs index 20d2d3df9a..53f1586abe 100644 --- a/proxy/src/bin/pg_sni_router.rs +++ b/proxy/src/bin/pg_sni_router.rs @@ -133,9 +133,7 @@ async fn main() -> anyhow::Result<()> { proxy_listener, cancellation_token.clone(), )); - let signals_task = tokio::spawn(proxy::handle_signals(cancellation_token, || async { - Ok(()) - })); + let signals_task = tokio::spawn(proxy::handle_signals(cancellation_token, || {})); // the signal task cant ever succeed. // the main task can error, or can succeed on cancellation. diff --git a/proxy/src/bin/proxy.rs b/proxy/src/bin/proxy.rs index 2ac66ffe8c..141005788d 100644 --- a/proxy/src/bin/proxy.rs +++ b/proxy/src/bin/proxy.rs @@ -461,10 +461,7 @@ async fn main() -> anyhow::Result<()> { // maintenance tasks. these never return unless there's an error let mut maintenance_tasks = JoinSet::new(); - maintenance_tasks.spawn(proxy::handle_signals( - cancellation_token.clone(), - || async { Ok(()) }, - )); + maintenance_tasks.spawn(proxy::handle_signals(cancellation_token.clone(), || {})); maintenance_tasks.spawn(http::health_server::task_main( http_listener, AppMetrics { diff --git a/proxy/src/console/messages.rs b/proxy/src/console/messages.rs index 85683acb82..1696e229ce 100644 --- a/proxy/src/console/messages.rs +++ b/proxy/src/console/messages.rs @@ -1,13 +1,11 @@ use measured::FixedCardinalityLabel; use serde::{Deserialize, Serialize}; -use std::collections::HashMap; use std::fmt::{self, Display}; use crate::auth::IpPattern; -use crate::intern::{BranchIdInt, EndpointIdInt, ProjectIdInt}; +use crate::intern::{BranchIdInt, EndpointIdInt, ProjectIdInt, RoleNameInt}; use crate::proxy::retry::CouldRetry; -use crate::RoleName; /// Generic error response with human-readable description. /// Note that we can't always present it to user as is. @@ -348,11 +346,6 @@ impl ColdStartInfo { } } -#[derive(Debug, Deserialize, Clone)] -pub struct JwksRoleMapping { - pub roles: HashMap, -} - #[derive(Debug, Deserialize, Clone)] pub struct EndpointJwksResponse { pub jwks: Vec, @@ -361,11 +354,10 @@ pub struct EndpointJwksResponse { #[derive(Debug, Deserialize, Clone)] pub struct JwksSettings { pub id: String, - pub project_id: ProjectIdInt, - pub branch_id: BranchIdInt, pub jwks_url: url::Url, pub provider_name: String, pub jwt_audience: Option, + pub role_names: Vec, } #[cfg(test)] diff --git a/proxy/src/intern.rs b/proxy/src/intern.rs index e5144cfe2e..108420d7d7 100644 --- a/proxy/src/intern.rs +++ b/proxy/src/intern.rs @@ -130,14 +130,14 @@ impl Default for StringInterner { } #[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)] -pub(crate) struct RoleNameTag; +pub struct RoleNameTag; impl InternId for RoleNameTag { fn get_interner() -> &'static StringInterner { static ROLE_NAMES: OnceLock> = OnceLock::new(); ROLE_NAMES.get_or_init(Default::default) } } -pub(crate) type RoleNameInt = InternedString; +pub type RoleNameInt = InternedString; impl From<&RoleName> for RoleNameInt { fn from(value: &RoleName) -> Self { RoleNameTag::get_interner().get_or_intern(value) diff --git a/proxy/src/lib.rs b/proxy/src/lib.rs index 0070839aa8..ea0a9beced 100644 --- a/proxy/src/lib.rs +++ b/proxy/src/lib.rs @@ -82,7 +82,7 @@ impl_trait_overcaptures, )] -use std::{convert::Infallible, future::Future}; +use std::convert::Infallible; use anyhow::{bail, Context}; use intern::{EndpointIdInt, EndpointIdTag, InternId}; @@ -117,13 +117,12 @@ pub mod usage_metrics; pub mod waiters; /// Handle unix signals appropriately. -pub async fn handle_signals( +pub async fn handle_signals( token: CancellationToken, mut refresh_config: F, ) -> anyhow::Result where - F: FnMut() -> Fut, - Fut: Future>, + F: FnMut(), { use tokio::signal::unix::{signal, SignalKind}; @@ -136,7 +135,7 @@ where // Hangup is commonly used for config reload. _ = hangup.recv() => { warn!("received SIGHUP"); - refresh_config().await?; + refresh_config(); } // Shut down the whole application. _ = interrupt.recv() => { diff --git a/proxy/src/scram/threadpool.rs b/proxy/src/scram/threadpool.rs index 2702aeebfe..c027a0cd20 100644 --- a/proxy/src/scram/threadpool.rs +++ b/proxy/src/scram/threadpool.rs @@ -43,6 +43,13 @@ impl ThreadPool { pub fn new(n_workers: u8) -> Arc { // rayon would be nice here, but yielding in rayon does not work well afaict. + if n_workers == 0 { + return Arc::new(Self { + runtime: None, + metrics: Arc::new(ThreadPoolMetrics::new(n_workers as usize)), + }); + } + Arc::new_cyclic(|pool| { let pool = pool.clone(); let worker_id = AtomicUsize::new(0); diff --git a/proxy/src/serverless/backend.rs b/proxy/src/serverless/backend.rs index aa236907db..607eb0caf6 100644 --- a/proxy/src/serverless/backend.rs +++ b/proxy/src/serverless/backend.rs @@ -119,7 +119,7 @@ impl PoolingBackend { .check_jwt( ctx, user_info.endpoint.clone(), - user_info.user.clone(), + &user_info.user, &StaticAuthRules, jwt, )