|
| 1 | +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. |
| 2 | +// SPDX-License-Identifier: Apache-2.0 |
| 3 | + |
| 4 | +//! Optional adapter from the standalone VM driver to the gateway registry. |
| 5 | +
|
| 6 | +use crate::{VmComputeConfig, spawn_managed_vm_driver}; |
| 7 | +use openshell_core::telemetry::TelemetryComputeDriver; |
| 8 | +use openshell_core::{Error, Result}; |
| 9 | +use openshell_server::{ |
| 10 | + ComputeDriverBuildContext, ComputeDriverConfigContext, ComputeDriverFactory, |
| 11 | + ComputeDriverInstance, ComputeDriverRegistration, connect_managed_compute_driver, |
| 12 | +}; |
| 13 | +use std::path::{Path, PathBuf}; |
| 14 | + |
| 15 | +const DRIVER_NAME: &str = "vm"; |
| 16 | + |
| 17 | +/// Build the VM driver's self-contained gateway registration. |
| 18 | +pub fn gateway_registration() -> Result<ComputeDriverRegistration> { |
| 19 | + ComputeDriverRegistration::new(DRIVER_NAME, u16::MAX, None, VmFactory).map(|registration| { |
| 20 | + registration |
| 21 | + .with_telemetry_category(TelemetryComputeDriver::anonymous_category(DRIVER_NAME)) |
| 22 | + .with_local_singleplayer() |
| 23 | + }) |
| 24 | +} |
| 25 | + |
| 26 | +#[derive(Clone, Copy)] |
| 27 | +struct VmFactory; |
| 28 | + |
| 29 | +#[async_trait::async_trait] |
| 30 | +impl ComputeDriverFactory for VmFactory { |
| 31 | + fn supports_config_preflight(&self) -> bool { |
| 32 | + true |
| 33 | + } |
| 34 | + |
| 35 | + fn validate_config(&self, context: ComputeDriverConfigContext<'_>) -> Result<()> { |
| 36 | + let mut config = vm_config(context)?; |
| 37 | + apply_default_grpc_endpoint( |
| 38 | + &mut config, |
| 39 | + context.gateway_tls_enabled(), |
| 40 | + context.gateway_port(), |
| 41 | + ); |
| 42 | + config.validate_configuration() |
| 43 | + } |
| 44 | + |
| 45 | + async fn build(&self, context: ComputeDriverBuildContext<'_>) -> Result<ComputeDriverInstance> { |
| 46 | + let mut config = vm_config(context.config_context())?; |
| 47 | + require_guest_tls(&context)?; |
| 48 | + if !context.gateway_tls_enabled() || context.guest_tls_paths().is_some() { |
| 49 | + apply_default_grpc_endpoint( |
| 50 | + &mut config, |
| 51 | + context.gateway_tls_enabled(), |
| 52 | + context.gateway_port(), |
| 53 | + ); |
| 54 | + } |
| 55 | + apply_guest_tls( |
| 56 | + &mut config.guest_tls_ca, |
| 57 | + &mut config.guest_tls_cert, |
| 58 | + &mut config.guest_tls_key, |
| 59 | + context.guest_tls_paths(), |
| 60 | + ); |
| 61 | + let launch = spawn_managed_vm_driver( |
| 62 | + context.gateway_log_level(), |
| 63 | + context.gateway_name(), |
| 64 | + &config, |
| 65 | + context.otlp_config().map(|config| config.endpoint.as_str()), |
| 66 | + )?; |
| 67 | + let (child, socket_path) = launch.into_parts(); |
| 68 | + let endpoint = connect_managed_compute_driver(DRIVER_NAME, socket_path, child) |
| 69 | + .await |
| 70 | + .map_err(|error| Error::execution(error.to_string()))?; |
| 71 | + Ok(ComputeDriverInstance::ManagedRemote(endpoint)) |
| 72 | + } |
| 73 | +} |
| 74 | + |
| 75 | +fn vm_config(context: ComputeDriverConfigContext<'_>) -> Result<VmComputeConfig> { |
| 76 | + let mut config: VmComputeConfig = context.driver_config()?; |
| 77 | + if config.state_dir.as_os_str().is_empty() { |
| 78 | + config.state_dir = VmComputeConfig::default_state_dir(); |
| 79 | + } |
| 80 | + Ok(config) |
| 81 | +} |
| 82 | + |
| 83 | +fn apply_default_grpc_endpoint(config: &mut VmComputeConfig, tls_enabled: bool, port: u16) { |
| 84 | + if config.grpc_endpoint.trim().is_empty() { |
| 85 | + let scheme = if tls_enabled { "https" } else { "http" }; |
| 86 | + config.grpc_endpoint = format!("{scheme}://127.0.0.1:{port}"); |
| 87 | + } |
| 88 | +} |
| 89 | + |
| 90 | +fn require_guest_tls(context: &ComputeDriverBuildContext<'_>) -> Result<()> { |
| 91 | + if context.gateway_tls_enabled() && context.guest_tls_paths().is_none() { |
| 92 | + return Err(Error::config(format!( |
| 93 | + "gateway TLS requires guest_tls_ca, guest_tls_cert, and guest_tls_key in [openshell.gateway] when using the {DRIVER_NAME} compute driver" |
| 94 | + ))); |
| 95 | + } |
| 96 | + Ok(()) |
| 97 | +} |
| 98 | + |
| 99 | +fn apply_guest_tls( |
| 100 | + ca: &mut Option<PathBuf>, |
| 101 | + cert: &mut Option<PathBuf>, |
| 102 | + key: &mut Option<PathBuf>, |
| 103 | + defaults: Option<(&Path, &Path, &Path)>, |
| 104 | +) { |
| 105 | + if ca.is_none() |
| 106 | + && cert.is_none() |
| 107 | + && key.is_none() |
| 108 | + && let Some((default_ca, default_cert, default_key)) = defaults |
| 109 | + { |
| 110 | + *ca = Some(default_ca.to_owned()); |
| 111 | + *cert = Some(default_cert.to_owned()); |
| 112 | + *key = Some(default_key.to_owned()); |
| 113 | + } |
| 114 | +} |
| 115 | + |
| 116 | +#[cfg(test)] |
| 117 | +mod tests { |
| 118 | + use super::apply_guest_tls; |
| 119 | + use std::path::{Path, PathBuf}; |
| 120 | + |
| 121 | + #[test] |
| 122 | + fn package_managed_guest_bundle_is_injected_when_driver_paths_are_absent() { |
| 123 | + let mut ca = None; |
| 124 | + let mut cert = None; |
| 125 | + let mut key = None; |
| 126 | + apply_guest_tls( |
| 127 | + &mut ca, |
| 128 | + &mut cert, |
| 129 | + &mut key, |
| 130 | + Some(( |
| 131 | + Path::new("ca.pem"), |
| 132 | + Path::new("client.pem"), |
| 133 | + Path::new("client-key.pem"), |
| 134 | + )), |
| 135 | + ); |
| 136 | + |
| 137 | + assert_eq!(ca, Some(PathBuf::from("ca.pem"))); |
| 138 | + assert_eq!(cert, Some(PathBuf::from("client.pem"))); |
| 139 | + assert_eq!(key, Some(PathBuf::from("client-key.pem"))); |
| 140 | + } |
| 141 | +} |
0 commit comments