diff --git a/crates/dhcp-server/src/command_line.rs b/crates/dhcp-server/src/command_line.rs index dc811b3a23..19eda53310 100644 --- a/crates/dhcp-server/src/command_line.rs +++ b/crates/dhcp-server/src/command_line.rs @@ -53,6 +53,23 @@ pub struct Args { )] pub host_config: Option, + #[arg(long, help = "Root CA certificate used to connect to the Carbide API.")] + pub forge_root_ca_path: Option, + + #[arg( + long, + requires = "client_key_path", + help = "Client certificate used to connect to the Carbide API." + )] + pub client_cert_path: Option, + + #[arg( + long, + requires = "client_cert_path", + help = "Client private key used to connect to the Carbide API." + )] + pub client_key_path: Option, + #[arg(short, long, value_enum, default_value_t=ServerMode::Dpu)] pub mode: ServerMode, @@ -99,6 +116,9 @@ mod tests { SocketAddrV4::new(Ipv4Addr::UNSPECIFIED, 67) ); assert_eq!(defaults.relay_response_port, 67); + assert_eq!(defaults.forge_root_ca_path, None); + assert_eq!(defaults.client_cert_path, None); + assert_eq!(defaults.client_key_path, None); let overridden = Args::try_parse_from([ "forge-dhcp-server", @@ -117,5 +137,28 @@ mod tests { assert!( Args::try_parse_from(["forge-dhcp-server", "--listen-addr", "[::]:6767",]).is_err() ); + + let tls = Args::try_parse_from([ + "forge-dhcp-server", + "--forge-root-ca-path", + "/local/ca.crt", + "--client-cert-path", + "/local/client.crt", + "--client-key-path", + "/local/client.key", + ]) + .unwrap(); + assert_eq!(tls.forge_root_ca_path.as_deref(), Some("/local/ca.crt")); + assert_eq!(tls.client_cert_path.as_deref(), Some("/local/client.crt")); + assert_eq!(tls.client_key_path.as_deref(), Some("/local/client.key")); + + assert!( + Args::try_parse_from([ + "forge-dhcp-server", + "--client-cert-path", + "/local/client.crt", + ]) + .is_err() + ); } } diff --git a/crates/dhcp-server/src/main.rs b/crates/dhcp-server/src/main.rs index ac3565db5d..50685522fd 100644 --- a/crates/dhcp-server/src/main.rs +++ b/crates/dhcp-server/src/main.rs @@ -31,12 +31,15 @@ use std::net::SocketAddr; use std::sync::Arc; use ::rpc::forge::{DhcpDiscovery, DhcpRecord}; +use ::rpc::forge_tls_client::ForgeClientConfig; use cache::CacheEntry; use carbide_instrument::emit; use carbide_rpc_utils::dhcp::{DhcpConfig, DhcpTimestamps, DhcpTimestampsFilePath, HostConfig}; use chrono::Utc; use command_line::{Args, ServerMode}; use errors::DhcpError; +use forge_tls::client_config::ClientCert; +use forge_tls::default::{default_client_cert, default_client_key, default_root_ca}; use grpc_server::{ControlRequest, run_grpc_server}; use lru::LruCache; use metrics::{DhcpPacketDropped, DhcpReplySent, DropReason}; @@ -558,9 +561,11 @@ pub struct Config { dhcp_config: DhcpConfig, host_config: Option, // Valid only for Dpu mode. relay_response_port: u16, + forge_client_config: ForgeClientConfig, } async fn init(args: Args) -> Result { + let forge_client_config = forge_client_config(&args)?; let f = tokio::fs::read_to_string(args.dhcp_config).await?; let dhcp_config: DhcpConfig = serde_yaml::from_str(&f)?; @@ -575,9 +580,34 @@ async fn init(args: Args) -> Result { dhcp_config, host_config, relay_response_port: args.relay_response_port, + forge_client_config, }) } +fn forge_client_config(args: &Args) -> Result { + let root_ca_path = args + .forge_root_ca_path + .clone() + .unwrap_or_else(|| default_root_ca().to_string()); + let client_cert = match (&args.client_cert_path, &args.client_key_path) { + (Some(cert_path), Some(key_path)) => ClientCert { + cert_path: cert_path.clone(), + key_path: key_path.clone(), + }, + (None, None) => ClientCert { + cert_path: default_client_cert().to_string(), + key_path: default_client_key().to_string(), + }, + _ => { + return Err(DhcpError::MissingArgument( + "client_cert_path and client_key_path must be configured together".to_string(), + )); + } + }; + + Ok(ForgeClientConfig::new(root_ca_path, Some(client_cert))) +} + #[derive(Debug)] pub struct TestArm {} @@ -733,7 +763,10 @@ mod test { use crate::command_line::{Args, ServerMode}; use crate::errors::DhcpError; - use crate::{DhcpMode, Test, TestArm, cache, handle_reload, init, packet_handler, process}; + use crate::{ + DhcpMode, Test, TestArm, cache, forge_client_config, handle_reload, init, packet_handler, + process, + }; fn make_reload_args(td: &TempDir, interfaces: Vec) -> Args { Args { @@ -742,6 +775,9 @@ mod test { relay_response_port: 67, dhcp_config: td.path().join("dhcp.yaml").display().to_string(), host_config: Some(td.path().join("host.yaml").display().to_string()), + forge_root_ca_path: None, + client_cert_path: None, + client_key_path: None, mode: ServerMode::Dpu, grpc_listen_addr: None, metrics_listen_addr: None, @@ -914,6 +950,9 @@ mod test { .display() .to_string(), ), + forge_root_ca_path: None, + client_cert_path: None, + client_key_path: None, mode: crate::command_line::ServerMode::Dpu, grpc_listen_addr: None, metrics_listen_addr: None, @@ -925,6 +964,28 @@ mod test { init(get_test_args()).await.unwrap(); } + #[test] + fn forge_client_tls_paths_are_configurable() { + let defaults = forge_client_config(&get_test_args()).unwrap(); + assert_eq!(defaults.root_ca_path, forge_tls::default::ROOT_CA); + let default_identity = defaults.client_cert.unwrap(); + assert_eq!(default_identity.cert_path, forge_tls::default::CLIENT_CERT); + assert_eq!(default_identity.key_path, forge_tls::default::CLIENT_KEY); + + let mut explicit = get_test_args(); + explicit.forge_root_ca_path = Some("/local/ca.crt".to_string()); + explicit.client_cert_path = Some("/local/client.crt".to_string()); + explicit.client_key_path = Some("/local/client.key".to_string()); + let configured = forge_client_config(&explicit).unwrap(); + assert_eq!(configured.root_ca_path, "/local/ca.crt"); + let configured_identity = configured.client_cert.unwrap(); + assert_eq!(configured_identity.cert_path, "/local/client.crt"); + assert_eq!(configured_identity.key_path, "/local/client.key"); + + explicit.client_key_path = None; + assert!(forge_client_config(&explicit).is_err()); + } + #[tokio::test] async fn test_arm_non_relayed_packet() { let byte_stream = get_byte_stream(Ipv4Addr::new(0, 0, 0, 0), None, MessageType::Request); diff --git a/crates/dhcp-server/src/rpc/client.rs b/crates/dhcp-server/src/rpc/client.rs index 44bce165ce..f76835ba98 100644 --- a/crates/dhcp-server/src/rpc/client.rs +++ b/crates/dhcp-server/src/rpc/client.rs @@ -14,10 +14,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -use forge_tls::client_config::ClientCert; -use forge_tls::default::{default_client_cert, default_client_key, default_root_ca}; use rpc::forge::{DhcpDiscovery, DhcpRecord}; -use rpc::forge_tls_client::{ApiConfig, ForgeClientConfig, ForgeTlsClient}; +use rpc::forge_tls_client::{ApiConfig, ForgeTlsClient}; use crate::Config; use crate::errors::DhcpError; @@ -32,15 +30,7 @@ pub async fn discover_dhcp( )); }; - let client_config = ForgeClientConfig::new( - default_root_ca().to_string(), - Some(ClientCert { - cert_path: default_client_cert().to_string(), - key_path: default_client_key().to_string(), - }), - ); - - let api_config = ApiConfig::new(carbide_api_url, &client_config); + let api_config = ApiConfig::new(carbide_api_url, &config.forge_client_config); let mut client = ForgeTlsClient::retry_build(&api_config) .await