-
Notifications
You must be signed in to change notification settings - Fork 6
Expand file tree
/
Copy pathmain.rs
More file actions
147 lines (127 loc) · 4.95 KB
/
Copy pathmain.rs
File metadata and controls
147 lines (127 loc) · 4.95 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use agent_gateway::proxy::MakeProxyService;
use agent_gateway::rate_limit::BucketStore;
use agent_gateway::{config, observability, policy, proxy, tls};
use anyhow::Context;
use clap::Parser;
use hyper_util::rt::TokioExecutor;
use sqlx::PgPool;
use tokio::net::TcpListener;
use tokio::time::MissedTickBehavior;
use tracing::{error, info, warn};
/// How often the background task polls `permission_registry` for newly
/// revoked rows and propagates their `revoked` flag to the matching
/// in-process buckets. Bounded latency for explicit revocation.
const REVOCATION_POLL_INTERVAL: Duration = Duration::from_secs(30);
#[derive(Parser)]
#[command(name = "agent_gateway", about = "mTLS HTTP/2 CONNECT proxy")]
struct Cli {
/// Path to the TOML configuration file
#[arg(short, long, default_value = "config.toml")]
config: PathBuf,
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
rustls::crypto::aws_lc_rs::default_provider()
.install_default()
.expect("failed to install default crypto provider");
let cli = Cli::parse();
let config = config::Config::load(&cli.config)
.with_context(|| format!("loading config from {}", cli.config.display()))?;
serve(config).await
}
async fn serve(config: config::Config) -> anyhow::Result<()> {
observability::init(&config.observability)?;
let server_tls = tls::build_server_config(&config.server)?;
let tls_acceptor = tls::TlsAcceptor::from(server_tls);
let policy_engine = policy::build_engine(&config.policy).await?;
let bucket_store = Arc::new(BucketStore::new());
let make_service = Arc::new(MakeProxyService::new(
policy_engine,
bucket_store.clone(),
));
let revocation_pool = policy::build_pool(&config.policy).await?;
let revocation_task = tokio::spawn(run_revocation_poll(revocation_pool, bucket_store));
let listen_addr: std::net::SocketAddr = config.server.listen_addr.parse()?;
let listener = TcpListener::bind(listen_addr).await?;
info!(%listen_addr, "listening");
tokio::select! {
result = serve_loop(&listener, &tls_acceptor, &make_service) => {
result?;
}
_ = tokio::signal::ctrl_c() => {
info!("received shutdown signal");
}
}
revocation_task.abort();
observability::shutdown();
Ok(())
}
/// Background task that propagates explicit permission revocations
/// (`revoked_at` set on `permission_registry`) to the matching in-process
/// buckets. Bounded latency = `REVOCATION_POLL_INTERVAL`.
async fn run_revocation_poll(pool: PgPool, store: Arc<BucketStore>) {
let mut ticker = tokio::time::interval(REVOCATION_POLL_INTERVAL);
ticker.set_missed_tick_behavior(MissedTickBehavior::Skip);
loop {
ticker.tick().await;
match query_revoked_permission_ids(&pool).await {
Ok(ids) => {
for id in ids {
store.mark_revoked(&id);
}
}
Err(e) => warn!(error = ?e, "revocation poll query failed"),
}
}
}
/// Read-only query that returns every `permission_id` with `revoked_at` set.
/// The gateway calls this on a 30-second timer; integration tests call it
/// directly so they don't have to wait on the timer.
async fn query_revoked_permission_ids(pool: &PgPool) -> anyhow::Result<Vec<String>> {
let rows = sqlx::query_scalar!(
"SELECT permission_id FROM permission_registry WHERE revoked_at IS NOT NULL"
)
.fetch_all(pool)
.await?;
Ok(rows)
}
async fn serve_loop(
listener: &TcpListener,
tls_acceptor: &tls::TlsAcceptor,
make_service: &Arc<MakeProxyService>,
) -> anyhow::Result<()> {
loop {
let (tcp_stream, peer_addr) = match listener.accept().await {
Ok(conn) => conn,
Err(e) => {
error!(error = %e, "TCP accept failed");
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
continue;
}
};
let acceptor = tls_acceptor.clone();
let make_svc = make_service.clone();
tokio::spawn(async move {
let tls_stream = match acceptor.accept(tcp_stream).await {
Ok(s) => s,
Err(e) => {
error!(source_peer_addr = %peer_addr, error = %e, "TLS handshake failed");
return;
}
};
let peer_certs = proxy::extract_peer_certs(tls_stream.get_ref().1);
let service = make_svc.make_service(peer_certs, peer_addr);
let io = hyper_util::rt::TokioIo::new(tls_stream);
if let Err(e) = hyper_util::server::conn::auto::Builder::new(TokioExecutor::new())
.http2_only()
.serve_connection_with_upgrades(io, service)
.await
{
error!(source_peer_addr = %peer_addr, error = %e, "connection error");
}
});
}
}