diff --git a/Cargo.lock b/Cargo.lock index 2edd188c..d4a3e49c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -614,7 +614,6 @@ checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" name = "contextforge-gateway-rs" version = "0.1.0" dependencies = [ - "anyhow", "clap", "contextforge-gateway-rs-lib", "futures", @@ -663,6 +662,7 @@ dependencies = [ "rmcp", "rmp-serde", "rustls", + "rustls-pemfile", "rustls-pki-types", "serde", "serde_json", @@ -1552,9 +1552,9 @@ dependencies = [ [[package]] name = "hashbrown" -version = "0.17.0" +version = "0.17.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4f467dd6dccf739c208452f8014c75c18bb8301b050ad1cfb27153803edb0f51" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" [[package]] name = "headers" @@ -1905,7 +1905,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" dependencies = [ "equivalent", - "hashbrown 0.17.0", + "hashbrown 0.17.1", "serde", "serde_core", ] @@ -2019,9 +2019,9 @@ dependencies = [ [[package]] name = "jsonwebtoken" -version = "10.3.0" +version = "10.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0529410abe238729a60b108898784df8984c87f6054c9c4fcacc47e4803c1ce1" +checksum = "eba32bfb4ffdeaca3e34431072faf01745c9b26d25504aa7a6cf5684334fc4fc" dependencies = [ "base64", "ed25519-dalek", @@ -2038,6 +2038,7 @@ dependencies = [ "sha2", "signature", "simple_asn1", + "zeroize", ] [[package]] @@ -2214,9 +2215,9 @@ dependencies = [ [[package]] name = "nix" -version = "0.31.2" +version = "0.31.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d6d0705320c1e6ba1d912b5e37cf18071b6c2e9b7fa8215a1e8a7651966f5d3" +checksum = "cf20d2fde8ff38632c426f1165ed7436270b44f199fc55284c38276f9db47c3d" dependencies = [ "bitflags", "cfg-if", @@ -3328,6 +3329,15 @@ dependencies = [ "security-framework", ] +[[package]] +name = "rustls-pemfile" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "rustls-pki-types" version = "1.14.1" @@ -4016,9 +4026,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.52.2" +version = "1.52.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "110a78583f19d5cdb2c5ccf321d1290344e71313c6c37d43520d386027d18386" +checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" dependencies = [ "bytes", "libc", @@ -5080,6 +5090,20 @@ name = "zeroize" version = "1.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0" +dependencies = [ + "zeroize_derive", +] + +[[package]] +name = "zeroize_derive" +version = "1.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85a5b4158499876c763cb03bc4e49185d3cccbabb15b33c627f7884f43db852e" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.117", +] [[package]] name = "zerotrie" diff --git a/Cargo.toml b/Cargo.toml index b318f9dc..2de8c54e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -18,10 +18,12 @@ repository = "https://github.com/contextforge-gateway-rs/contextforge-gateway-rs [workspace.dependencies] + contextforge-gateway-rs = { path = "./crates/contextforge-gateway-rs" } contextforge-gateway-rs-lib = { path = "./crates/contextforge-gateway-rs-lib", features = [ "with_tools", ] } + rmcp = { version = "1.6.0", features = [ "server", "client", @@ -32,7 +34,6 @@ rmcp = { version = "1.6.0", features = [ "elicitation", ], git = "https://github.com/contextforge-gateway-rs/mcp-rust-sdk", branch = "enabling_propagation_of_new_session_id_2", commit="d07fa4f8d7978cd0829a09fab181437b99a81878" } - serde = "1.0" serde_json = "1.0" tracing = "0.1" @@ -51,7 +52,11 @@ axum-jwt-auth = "0.6.3" futures = { version = "0.3", features = ["std", "alloc"] } jsonwebtoken = "10.3.0" chrono = "0.4.44" -redis = { version = "1.1.0", features = ["default", "tokio-rustls-comp"] } +redis = { version = "1.2.1", features = [ + "default", + "tokio-rustls-comp", + "tls-rustls", +] } clap = { version = "4.5.60", features = ["derive", "env"] } thiserror = "2.0.18" openid = "0.23.0" diff --git a/README.md b/README.md index 593a3146..cbd7e094 100644 --- a/README.md +++ b/README.md @@ -9,7 +9,8 @@ docker compose -f docker/docker-compose-local.yaml up -d 2. Run gateway ```bash - cargo run --release --bin contextforge-gateway-rs -- --address 0.0.0.0:8001 --redis-port 6379 --redis-address 127.0.0.1 --token-verification-public-key assets/jwt.key.pub --token-verification-private-key assets/jwt.key --number-of-cpus 16 + cargo run --bin contextforge-gateway-rs -- --address 0.0.0.0:8001 --redis-port 6379 --redis-address 127.0.0.1 --token-verification-public-key assets/jwt.key.pub --token-verification-private-key assets/jwt.key --number-of-cpus 16 --redis-mode=plain-text --upstream-connection-mode=plain-text-or-tls + ``` This should spin up Redis instance and two mcp-gateways: a simple counter and a conformance test server from mcp-rust-sdk @@ -55,4 +56,4 @@ cargo run --release --bin contextforge-load-test -- --host 'http://127.0.0.1:800 ``` -[Performance reports](./reports) \ No newline at end of file +[Performance reports](./reports) diff --git a/crates/contextforge-gateway-rs-lib/Cargo.toml b/crates/contextforge-gateway-rs-lib/Cargo.toml index d68c88b6..e97dd5f2 100644 --- a/crates/contextforge-gateway-rs-lib/Cargo.toml +++ b/crates/contextforge-gateway-rs-lib/Cargo.toml @@ -46,8 +46,11 @@ hyper-util = "0.1.20" hyper = { version = "1.4.0" } rustls.workspace = true rustls-pki-types = { version = "1.14.1", features = ["std"] } +rustls-pemfile = "2.2.0" tokio-rustls = "0.26.4" typed-builder.workspace = true + + [features] default = [] with_tools = [] diff --git a/crates/contextforge-gateway-rs-lib/src/common.rs b/crates/contextforge-gateway-rs-lib/src/common.rs index 09a8ba25..de6841ec 100644 --- a/crates/contextforge-gateway-rs-lib/src/common.rs +++ b/crates/contextforge-gateway-rs-lib/src/common.rs @@ -1,6 +1,6 @@ use std::{ fs::{self, File}, - io::Read, + io::{Cursor, Read}, net::SocketAddr, path::PathBuf, sync::Arc, @@ -11,7 +11,7 @@ use chrono::{Duration, Utc}; use clap::{Parser, ValueEnum}; use http::uri::Authority; use openid::{CompactJson, CustomClaims, StandardClaims}; -use redis::{ConnectionAddr, IntoConnectionInfo}; +use redis::{ConnectionAddr, IntoConnectionInfo, RedisError}; use serde::{Deserialize, Serialize}; use thiserror::Error; use url::Url; @@ -43,35 +43,59 @@ impl CompactJson for ContextForgeGatewayClaims {} pub type RedisClient = redis::Client; #[derive(Debug, Clone)] -pub struct RedisConfig { - address: String, - port: u16, +pub enum RedisConfig { + PlainText { host: String, port: u16 }, + + Tls { host: String, port: u16, trust_bundle: Vec }, + MTls { host: String, port: u16, trust_bundle: Vec, client_cert: Vec, client_key: Vec }, } -impl IntoConnectionInfo for RedisConfig { - fn into_connection_info(self) -> redis::RedisResult { - ConnectionAddr::Tcp(self.address, self.port).into_connection_info() +impl TryFrom for RedisClient { + type Error = RedisError; + + fn try_from(redis_config: RedisConfig) -> Result { + match redis_config { + RedisConfig::PlainText { host, port } => { + Ok(RedisClient::open(ConnectionAddr::Tcp(host, port).into_connection_info()?)?) + }, + RedisConfig::Tls { host, port, trust_bundle } => RedisClient::build_with_tls( + format!("rediss://{host}:{port}"), + redis::TlsCertificates { client_tls: None, root_cert: Some(trust_bundle) }, + ), + RedisConfig::MTls { host, port, trust_bundle, client_cert, client_key } => RedisClient::build_with_tls( + format!("rediss://{host}:{port}"), + redis::TlsCertificates { + client_tls: Some(redis::ClientTlsConfig { client_cert, client_key }), + root_cert: Some(trust_bundle), + }, + ), + } } } #[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, ValueEnum)] pub enum UpstreamConnectionMode { - PlainTextAndTls, - PlainTextAndMTls, + PlainTextOrTls, + PlainTextOrMTls, TlsOnly, MtlsOnly, } +#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, ValueEnum)] +#[derive(Default)] +pub enum RedisConnectionMode { + PlainText, + #[default] + Tls, + Mtls, +} + #[derive(Debug, Clone, Parser, Default)] #[command(name = "contextforge-gateway-rs")] #[command(about = "Minimal, fast and experimental Gateway/Dataplane for ContextForge")] pub struct Config { #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_ADDRESS")] pub address: Option, - #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_HOSTNAME")] - pub redis_address: String, - #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_PORT")] - pub redis_port: u16, #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_TOKEN_VERIFICATION_PUBLIC_KEY")] pub token_verification_public_key: PathBuf, @@ -109,6 +133,23 @@ pub struct Config { #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_TLS_UPSTREAM_TRUST_BUNDLE")] pub upstream_trust_bundle: Option, + + #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_HOSTNAME")] + pub redis_address: String, + #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_PORT")] + pub redis_port: u16, + + #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_CONNECTION_MODE")] + pub redis_mode: RedisConnectionMode, + + #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_TLS_REDIS_TRUST_BUNDLE")] + pub redis_tls_trust_bundle: Option, + + #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_TLS_REDIS_CLIENT_PRIVATE_KEY")] + pub redis_tls_client_private_key: Option, + + #[arg(long, env = "CONTEXTFORGE_GATEWAY_RS_REDIS_TLS_REDIS_CLIENT_CERTIFICATE")] + pub redis_tls_client_certificate: Option, } #[derive(Error, Debug)] @@ -122,12 +163,102 @@ impl TryFrom<&Config> for RedisConfig { let _: Authority = format!("{}:{}", value.redis_address, value.redis_port) .parse::() .map_err(|e| ConfigValidationError::RedisConfigurationError(e.to_string()))?; - Ok(Self { address: value.redis_address.clone(), port: value.redis_port }) + + match value.redis_mode { + RedisConnectionMode::PlainText => { + Ok(Self::PlainText { host: value.redis_address.clone(), port: value.redis_port }) + }, + RedisConnectionMode::Tls => { + let Some(trust_bundle) = &value.redis_tls_trust_bundle else { + return Err(ConfigValidationError::RedisConfigurationError(format!( + "Trust bundle is required for Redis {:?}", + value.redis_mode + ))); + }; + + let trust_bundle = validate_certs(trust_bundle)?; + + Ok(Self::Tls { host: value.redis_address.clone(), port: value.redis_port, trust_bundle }) + }, + RedisConnectionMode::Mtls => { + let Some(trust_bundle) = &value.redis_tls_trust_bundle else { + return Err(ConfigValidationError::RedisConfigurationError(format!( + "Trust bundle is required for Redis {:?}", + value.redis_mode + ))); + }; + + let trust_bundle = validate_certs(trust_bundle)?; + + let Some(certificate) = &value.redis_tls_client_certificate else { + return Err(ConfigValidationError::RedisConfigurationError(format!( + "Client certificate is required for Redis {:?}", + value.redis_mode + ))); + }; + + let client_cert = validate_certs(certificate)?; + + let Some(key) = &value.redis_tls_client_private_key else { + return Err(ConfigValidationError::RedisConfigurationError(format!( + "Client key is required for Redis {:?}", + value.redis_mode + ))); + }; + + let client_key = validate_key(key)?; + + Ok(Self::MTls { + host: value.redis_address.clone(), + port: value.redis_port, + trust_bundle, + client_cert, + client_key, + }) + }, + } } type Error = ConfigValidationError; } +fn validate_certs(path: &PathBuf) -> Result, ConfigValidationError> { + let mut buf = Vec::new(); + File::open(path) + .map_err(|e| ConfigValidationError::RedisConfigurationError(e.to_string()))? + .read_to_end(&mut buf) + .map_err(|e| ConfigValidationError::RedisConfigurationError(e.to_string()))?; + let mut cursor = Cursor::new(buf); + + let mut count = 0; + for cert in rustls_pemfile::certs(&mut cursor) { + if let Err(e) = cert { + return Err(ConfigValidationError::RedisConfigurationError(e.to_string())); + } + count += 1; + } + if count == 0 { + Err(ConfigValidationError::RedisConfigurationError("No certificates provided".to_owned())) + } else { + Ok(cursor.into_inner()) + } +} + +fn validate_key(path: &PathBuf) -> Result, ConfigValidationError> { + let mut buf = Vec::new(); + File::open(path) + .map_err(|e| ConfigValidationError::RedisConfigurationError(e.to_string()))? + .read_to_end(&mut buf) + .map_err(|e| ConfigValidationError::RedisConfigurationError(e.to_string()))?; + let mut cursor = Cursor::new(buf); + + if let Ok(Some(_)) = rustls_pemfile::private_key(&mut cursor) { + Ok(cursor.into_inner()) + } else { + Err(ConfigValidationError::RedisConfigurationError("Private key is wrong".to_owned())) + } +} + impl TryFrom<&Config> for reqwest::Client { type Error = crate::Error; @@ -135,8 +266,8 @@ impl TryFrom<&Config> for reqwest::Client { let builder = reqwest::Client::builder(); let builder = match config.upstream_connection_mode.as_ref() { None | Some(UpstreamConnectionMode::TlsOnly) => builder.https_only(true), - Some(UpstreamConnectionMode::PlainTextAndTls) => builder.https_only(false), - Some(UpstreamConnectionMode::PlainTextAndMTls) => { + Some(UpstreamConnectionMode::PlainTextOrTls) => builder.https_only(false), + Some(UpstreamConnectionMode::PlainTextOrMTls) => { builder.https_only(false).identity(extract_identity(config)?) }, Some(UpstreamConnectionMode::MtlsOnly) => builder.https_only(true).identity(extract_identity(config)?), diff --git a/crates/contextforge-gateway-rs-lib/src/lib.rs b/crates/contextforge-gateway-rs-lib/src/lib.rs index 8e319b09..4a1e382b 100644 --- a/crates/contextforge-gateway-rs-lib/src/lib.rs +++ b/crates/contextforge-gateway-rs-lib/src/lib.rs @@ -136,3 +136,8 @@ impl Gateway { Ok(()) } } + +pub fn get_config_store(config: &Config) -> Result { + let redis_config = RedisConfig::try_from(config)?; + Ok(RedisUserConfigStore::new(RedisClient::try_from(redis_config)?)) +} diff --git a/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs b/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs index dd10cee2..81af0d1d 100644 --- a/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs +++ b/crates/contextforge-gateway-rs-lib/src/tests/gateway_end_to_end.rs @@ -257,7 +257,7 @@ async fn plaintext_list_tools_end_to_end_test() -> crate::Result<()> { let config = Config { address: Some(format!("127.0.0.1:{gateway_port}").parse().expect("This should work")), token_verification_public_key: "../../assets/jwt.key.pub".into(), - upstream_connection_mode: Some(crate::common::UpstreamConnectionMode::PlainTextAndTls), + upstream_connection_mode: Some(crate::common::UpstreamConnectionMode::PlainTextOrTls), ..Default::default() }; @@ -335,7 +335,7 @@ async fn tls_list_tools_end_to_end_test() -> crate::Result<()> { let config = Config { token_verification_public_key: "../../assets/jwt.key.pub".into(), - upstream_connection_mode: Some(crate::common::UpstreamConnectionMode::PlainTextAndTls), + upstream_connection_mode: Some(crate::common::UpstreamConnectionMode::PlainTextOrTls), tls_address: Some(server_socket_addr), server_private_key: Some("../../assets/contextforgeCA/contextforge-server.key.pem".into()), server_certificate: Some("../../assets/contextforgeCA/contextforge-server.cert.pem".into()), diff --git a/crates/contextforge-gateway-rs/Cargo.toml b/crates/contextforge-gateway-rs/Cargo.toml index ff1a3015..7f01ed81 100644 --- a/crates/contextforge-gateway-rs/Cargo.toml +++ b/crates/contextforge-gateway-rs/Cargo.toml @@ -9,7 +9,6 @@ keywords.workspace = true [dependencies] contextforge-gateway-rs-lib.workspace = true -anyhow.workspace = true clap.workspace = true thiserror.workspace = true tracing.workspace = true diff --git a/crates/contextforge-gateway-rs/src/main.rs b/crates/contextforge-gateway-rs/src/main.rs index 85be9566..e57ee98e 100644 --- a/crates/contextforge-gateway-rs/src/main.rs +++ b/crates/contextforge-gateway-rs/src/main.rs @@ -4,7 +4,7 @@ mod runtime; use std::sync::Arc; use clap::Parser; -use contextforge_gateway_rs_lib::{Config, Gateway, RedisClient, RedisConfig, RedisUserConfigStore}; +use contextforge_gateway_rs_lib::{Config, Gateway, get_config_store}; use rmcp::transport::streamable_http_server::session::local::LocalSessionManager; use rustls::crypto; use tikv_jemallocator::Jemalloc; @@ -22,7 +22,7 @@ fn main() -> Result<(), Box> { let runtime = runtime::Runtime::from(&config); - let user_config_store = RedisUserConfigStore::new(RedisClient::open(RedisConfig::try_from(&config)?)?); + let user_config_store = get_config_store(&config)?; let gateway = Gateway::builder() .with_config(config) .with_user_config_store(Arc::new(user_config_store)) diff --git a/docker/docker-compose-local.yaml b/docker/docker-compose-local.yaml index 76fd1b72..33736793 100644 --- a/docker/docker-compose-local.yaml +++ b/docker/docker-compose-local.yaml @@ -105,8 +105,24 @@ services: - "300" - "--maxclients" - "10000" + - "--port" + - "6379" + - "--tls-port" + - "16379" + - "--tls-cert-file" + - "/tls_config/contextforge-server.cert.pem" + - "--tls-key-file" + - "/tls_config/contextforge-server.key.pem" + - "--tls-auth-clients" + - "no" + # - "--tls-ca-cert-file" + # - "/tls_config/contextforge.intermediate.cert.pem" ports: - "6379:6379" # expose only if you want host access + - "16379:16379" # expose only if you want host access + volumes: + - ../assets/contextforgeCA/:/tls_config:z + networks: [gateways-net] deploy: