Skip to main content

mobius_gateway/command/
lifecycle.rs

1use super::*;
2
3#[cfg(any(unix, test))]
4#[derive(Debug, Serialize, Deserialize)]
5#[serde(deny_unknown_fields)]
6pub(super) struct ProcessRecord {
7    pub(super) pid: u32,
8    pub(super) endpoint: Option<String>,
9}
10
11pub(super) struct ProcessRecordGuard {
12    #[cfg(unix)]
13    pub(super) path: PathBuf,
14    #[cfg(unix)]
15    pub(super) file: File,
16}
17
18#[derive(Debug)]
19pub(super) struct StartupGuard {
20    #[cfg(unix)]
21    pub(super) file: File,
22}
23
24pub(super) async fn serve(
25    state_dir: PathBuf,
26    lock_startup: bool,
27    save_local_client: fn(&Endpoint, String) -> Result<()>,
28    load_local_client: fn(&Endpoint) -> Result<Option<String>>,
29) -> Result<()> {
30    let (store, config) = ConfigStore::open(state_dir)?;
31    let state_dir = store.state_dir().to_path_buf();
32    let startup = lock_startup
33        .then(|| StartupGuard::create(&state_dir))
34        .transpose()?;
35    #[cfg(not(unix))]
36    let _ = load_local_client;
37    #[cfg(unix)]
38    if reuse_current_gateway(&store, &config, None, load_local_client)
39        .await?
40        .is_some()
41    {
42        println!("gateway is already running at the same or a newer version");
43        return Ok(());
44    }
45    #[cfg(unix)]
46    ensure_gateway_stopped(&store, &config)?;
47    let auth = AuthStore::open(store.auth_path(), config.auth)?;
48    if let Some((endpoint, token)) = provision_cloudflare_local_client(&auth, &config)? {
49        save_local_client(&endpoint, token)?;
50    }
51    let mut server = GatewayServer::open(state_dir.clone()).await?;
52    let ready = server.notify_ready();
53    let mut tunnel = CloudflareTunnel::start(&store, &config)?;
54    let endpoint = match &mut tunnel {
55        Some(tunnel) => Some(tunnel.endpoint().await?),
56        None => None,
57    };
58    let serving = async {
59        match &endpoint {
60            Some(endpoint) => server.serve_cloudflare(endpoint.host().to_owned()).await,
61            None => server.serve().await,
62        }
63    };
64    tokio::pin!(serving);
65    tokio::select! {
66        biased;
67        result = &mut serving => return result,
68        result = ready => result.map_err(|_| Error::Config("gateway stopped before becoming ready".into()))?,
69    }
70    // The parent treats this locked record as readiness, so publish it only after
71    // the serving loop has completed initialization and can accept connections.
72    let _process_record = ProcessRecordGuard::create(&state_dir, endpoint.as_ref())?;
73    drop(startup);
74    println!("gateway serving in foreground");
75    print_listener(&config, endpoint.as_ref());
76    tokio::select! {
77        result = &mut serving => result,
78        result = async {
79            match &mut tunnel {
80                Some(tunnel) => tunnel.wait().await,
81                None => std::future::pending().await,
82            }
83        } => result,
84    }
85}
86
87#[cfg(unix)]
88pub(super) async fn serve_in_background(
89    state_dir: PathBuf,
90    load_local_client: fn(&Endpoint) -> Result<Option<String>>,
91) -> Result<()> {
92    let (store, config) = ConfigStore::open(state_dir)?;
93    let _startup = StartupGuard::create(store.state_dir())?;
94    let mut interrupts = signal(SignalKind::interrupt())?;
95    let mut terminations = signal(SignalKind::terminate())?;
96    let Some(process) = start_background_gateway(
97        store.state_dir(),
98        &mut interrupts,
99        &mut terminations,
100        load_local_client,
101    )
102    .await?
103    else {
104        println!("gateway start cancelled");
105        return Ok(());
106    };
107    println!("gateway running (pid {})", process.pid);
108    print_listener(&config, process.endpoint()?.as_ref());
109    Ok(())
110}
111
112/// Starts the configured detached gateway, replacing an older running release.
113#[cfg(unix)]
114/// # Errors
115///
116/// Returns an error if the supplied value is invalid.
117pub async fn ensure_background_gateway(
118    state_dir: PathBuf,
119    load_local_client: fn(&Endpoint) -> Result<Option<String>>,
120) -> Result<()> {
121    serve_in_background(state_dir, load_local_client).await
122}
123
124#[cfg(unix)]
125pub(super) async fn start_background_gateway(
126    state_dir: &Path,
127    interrupts: &mut TokioSignal,
128    terminations: &mut TokioSignal,
129    load_local_client: fn(&Endpoint) -> Result<Option<String>>,
130) -> Result<Option<ProcessRecord>> {
131    let state_dir = fs::canonicalize(state_dir)?;
132    let process_path = state_dir.join(PROCESS_FILE);
133    let (store, config) = ConfigStore::open(state_dir.clone())?;
134    if let Some(process) = reuse_current_gateway(&store, &config, None, load_local_client).await? {
135        return Ok(Some(process));
136    }
137
138    let log = tempfile::NamedTempFile::new_in(&state_dir)?;
139    log.as_file().set_permissions(mobius::owner_only::file())?;
140    let mut command = TokioCommand::new(std::env::current_exe()?);
141    command
142        .arg("__serve")
143        .arg("--state-dir")
144        .arg(&state_dir)
145        .current_dir(&state_dir)
146        .stdin(Stdio::null())
147        .stdout(Stdio::null())
148        .stderr(Stdio::from(log.reopen()?));
149    #[cfg(target_os = "macos")]
150    command.env(
151        "PATH",
152        macos_gateway_path(std::env::var_os("PATH").as_deref()),
153    );
154    command.as_std_mut().process_group(0);
155
156    let mut child = command.spawn()?;
157    let Some(pid) = child.id() else {
158        stop_background_child(&mut child, &process_path).await;
159        return Err(Error::Config("background gateway has no process ID".into()));
160    };
161    let started = Instant::now();
162    loop {
163        match child.try_wait() {
164            Ok(Some(status)) => {
165                return Err(background_startup_error(
166                    format!("background gateway exited during startup with {status}"),
167                    &log,
168                ));
169            }
170            Ok(None) => {}
171            Err(error) => {
172                stop_background_child(&mut child, &process_path).await;
173                return Err(error.into());
174            }
175        }
176
177        let process_error = match running_process_record(&process_path) {
178            Ok(Some(record)) if record.pid == pid => return Ok(Some(record)),
179            Ok(Some(record)) => {
180                stop_background_child(&mut child, &process_path).await;
181                return Err(background_startup_error(
182                    format!(
183                        "gateway process {} claimed the process record during startup",
184                        record.pid
185                    ),
186                    &log,
187                ));
188            }
189            Ok(None) => None,
190            Err(error) => Some(error),
191        };
192
193        if started.elapsed() >= BACKGROUND_START_TIMEOUT {
194            stop_background_child(&mut child, &process_path).await;
195            let message = process_error.map_or_else(
196                || {
197                    format!(
198                        "background gateway did not start within {} seconds",
199                        BACKGROUND_START_TIMEOUT.as_secs()
200                    )
201                },
202                |error| format!("background gateway process record is invalid: {error}"),
203            );
204            return Err(background_startup_error(message, &log));
205        }
206        tokio::select! {
207            () = shutdown_signal(interrupts, terminations) => {
208                stop_background_child(&mut child, &process_path).await;
209                return Ok(None);
210            }
211            () = tokio::time::sleep(BACKGROUND_START_POLL_INTERVAL) => {}
212        }
213    }
214}
215
216#[cfg(target_os = "macos")]
217pub(super) fn macos_gateway_path(inherited: Option<&std::ffi::OsStr>) -> OsString {
218    // Finder-launched apps omit Homebrew; retain the user's existing tool priority.
219    let mut path = inherited
220        .unwrap_or(std::ffi::OsStr::new("/usr/bin:/bin:/usr/sbin:/sbin"))
221        .to_os_string();
222    path.push(":/opt/homebrew/bin:/usr/local/bin");
223    path
224}
225
226#[cfg(unix)]
227pub(super) async fn shutdown_signal(interrupts: &mut TokioSignal, terminations: &mut TokioSignal) {
228    tokio::select! {
229        _ = interrupts.recv() => {}
230        _ = terminations.recv() => {}
231    }
232}
233
234#[cfg(not(unix))]
235pub(super) async fn serve_in_background(
236    _state_dir: PathBuf,
237    _load_local_client: fn(&Endpoint) -> Result<Option<String>>,
238) -> Result<()> {
239    Err(unsupported_lifecycle())
240}
241
242#[cfg(not(unix))]
243pub async fn ensure_background_gateway(
244    _state_dir: PathBuf,
245    _load_local_client: fn(&Endpoint) -> Result<Option<String>>,
246) -> Result<()> {
247    Err(unsupported_lifecycle())
248}
249
250#[cfg(unix)]
251pub(super) async fn reuse_current_gateway(
252    store: &ConfigStore,
253    config: &GatewayConfig,
254    configured_endpoint: Option<&Endpoint>,
255    load_local_client: fn(&Endpoint) -> Result<Option<String>>,
256) -> Result<Option<ProcessRecord>> {
257    let Some(process) = running_process_record(&store.state_dir().join(PROCESS_FILE))? else {
258        return Ok(None);
259    };
260    let endpoint = if config.tls.is_some() {
261        let endpoint = match configured_endpoint {
262            Some(endpoint) => std::borrow::Cow::Borrowed(endpoint),
263            None => std::borrow::Cow::Owned(Endpoint::from_env()?),
264        };
265        if endpoint.is_plaintext() || endpoint.is_websocket() {
266            return Err(Error::Config(
267                "TLS gateway upgrades require MOBIUS_GATEWAY_ENDPOINT with the certificate hostname".into(),
268            ));
269        }
270        endpoint
271    } else {
272        std::borrow::Cow::Owned(loopback_endpoint(config)?)
273    };
274    let token = load_local_client(&endpoint)?.ok_or_else(|| {
275        Error::Config("local gateway credential is unavailable; cannot check its version".into())
276    })?;
277    let version = tokio::time::timeout(
278        Duration::from_secs(2),
279        endpoint.local_gateway_version(config.listen, &token),
280    )
281    .await
282    .map_err(|_| Error::Config("gateway version check timed out".into()))??;
283    if !gateway_version_is_older(&version, env!("CARGO_PKG_VERSION"))? {
284        return Ok(Some(process));
285    }
286    let state_dir = store.state_dir().to_path_buf();
287    tokio::task::spawn_blocking(move || stop_gateway(&state_dir, Some(process.pid)))
288        .await
289        .map_err(|error| Error::Config(format!("gateway stop task failed: {error}")))??;
290    Ok(None)
291}
292
293#[cfg(any(unix, test))]
294pub(super) fn gateway_version_is_older(running: &str, starting: &str) -> Result<bool> {
295    let parse = |version: &str| {
296        semver::Version::parse(version)
297            .map_err(|error| Error::Config(format!("invalid gateway version `{version}`: {error}")))
298    };
299    Ok(parse(running)?.cmp_precedence(&parse(starting)?).is_lt())
300}
301
302#[cfg(unix)]
303pub(super) async fn stop_background_child(child: &mut Child, process_path: &Path) {
304    if let Some(pid) = child.id() {
305        terminate_process_group(pid);
306    }
307    let _ = child.wait().await;
308    remove_unlocked_process_record(process_path);
309}
310
311#[cfg(unix)]
312pub(super) fn remove_unlocked_process_record(path: &Path) {
313    let Ok(file) = OpenOptions::new().read(true).write(true).open(path) else {
314        return;
315    };
316    if file.try_lock().is_ok() {
317        let _ = fs::remove_file(path);
318    }
319}
320
321#[cfg(unix)]
322pub(super) fn background_startup_error(
323    message: impl std::fmt::Display,
324    log: &tempfile::NamedTempFile,
325) -> Error {
326    startup_error(message, log, MAX_BACKGROUND_ERROR_BYTES)
327}
328
329/// Adds a bounded diagnostic excerpt from a gateway startup log.
330pub fn startup_error(
331    message: impl std::fmt::Display,
332    log: &tempfile::NamedTempFile,
333    max_bytes: u64,
334) -> Error {
335    let mut details = String::new();
336    if let Ok(file) = File::open(log.path()) {
337        let _ = file.take(max_bytes).read_to_string(&mut details);
338    }
339    let details = details.trim();
340    Error::Config(if details.is_empty() {
341        message.to_string()
342    } else {
343        format!("{message}: {details}")
344    })
345}
346
347pub(super) fn print_listener(config: &GatewayConfig, runtime_endpoint: Option<&Endpoint>) {
348    if let Some(cloudflare) = &config.cloudflare {
349        if let Some(endpoint) = runtime_endpoint
350            .map(ToString::to_string)
351            .or_else(|| cloudflare.endpoint())
352        {
353            println!("public endpoint: {endpoint}");
354        } else {
355            println!("public endpoint: assigned when the gateway starts");
356        }
357        println!("local endpoint: tcp://{}", config.listen);
358        println!("tunnel origin: http://{}", config.listen);
359        return;
360    }
361    let scheme = if config.tls.is_some() { "tls" } else { "tcp" };
362    println!("listener: {scheme}://{}", config.listen);
363}
364
365#[cfg(unix)]
366pub(super) fn exit_gateway(state_dir: PathBuf) -> Result<()> {
367    let (store, _) = ConfigStore::open(state_dir)?;
368    let _startup = StartupGuard::create(store.state_dir())?;
369    stop_gateway(store.state_dir(), None)
370}
371
372#[cfg(unix)]
373pub(super) fn stop_gateway(state_dir: &Path, expected_pid: Option<u32>) -> Result<()> {
374    let path = state_dir.join(PROCESS_FILE);
375    let Some((record, file)) = open_process_record(&path)? else {
376        eprintln!("gateway is stopped");
377        return Ok(());
378    };
379    if !process_is_running(&file)? {
380        eprintln!("gateway is stopped");
381        return Ok(());
382    }
383    if let Some(expected_pid) = expected_pid
384        && record.pid != expected_pid
385    {
386        return Err(Error::Config(format!(
387            "gateway process changed from {expected_pid} to {}",
388            record.pid
389        )));
390    }
391    let pid = i32::try_from(record.pid)
392        .map(Pid::from_raw)
393        .map_err(|_| Error::Config("invalid gateway process record".into()))?;
394    if let Err(error) = kill(pid, Signal::SIGINT) {
395        if !process_is_running(&file)? {
396            eprintln!("gateway is stopped");
397            return Ok(());
398        }
399        return Err(Error::Config(format!(
400            "failed to interrupt gateway: {error}"
401        )));
402    }
403    let started = Instant::now();
404    while process_is_running(&file)? {
405        if started.elapsed() >= EXIT_TIMEOUT {
406            return Err(Error::Config(format!(
407                "gateway process {} did not stop within {} seconds",
408                record.pid,
409                EXIT_TIMEOUT.as_secs()
410            )));
411        }
412        std::thread::sleep(EXIT_POLL_INTERVAL);
413    }
414    eprintln!("gateway stopped");
415    Ok(())
416}
417
418#[cfg(not(unix))]
419pub(super) fn exit_gateway(_state_dir: PathBuf) -> Result<()> {
420    Err(unsupported_lifecycle())
421}
422
423#[cfg(any(unix, test))]
424impl ProcessRecord {
425    pub(super) fn validate(&self) -> Result<()> {
426        if self.pid == 0 || i32::try_from(self.pid).is_err() {
427            return Err(Error::Config("invalid gateway process record".into()));
428        }
429        if let Some(endpoint) = self.endpoint()?
430            && !endpoint.is_websocket()
431        {
432            return Err(Error::Config(
433                "gateway process endpoint must use wss://".into(),
434            ));
435        }
436        Ok(())
437    }
438
439    pub(super) fn endpoint(&self) -> Result<Option<Endpoint>> {
440        self.endpoint.as_deref().map(str::parse).transpose()
441    }
442}
443
444impl ProcessRecordGuard {
445    #[cfg(unix)]
446    pub(super) fn create(state_dir: &Path, endpoint: Option<&Endpoint>) -> Result<Self> {
447        let state_dir = fs::canonicalize(state_dir)?;
448        let path = state_dir.join(PROCESS_FILE);
449        let mut file = OpenOptions::new()
450            .create(true)
451            .truncate(false)
452            .read(true)
453            .write(true)
454            .open(&path)?;
455        file.set_permissions(mobius::owner_only::file())?;
456        file.try_lock().map_err(|error| match error {
457            TryLockError::WouldBlock => Error::Config("gateway is already running".into()),
458            TryLockError::Error(error) => error.into(),
459        })?;
460        file.set_len(0)?;
461        file.seek(SeekFrom::Start(0))?;
462        let record = ProcessRecord {
463            pid: std::process::id(),
464            endpoint: endpoint.map(ToString::to_string),
465        };
466        serde_json::to_writer(&mut file, &record)?;
467        file.flush()?;
468        file.sync_all()?;
469        Ok(Self { path, file })
470    }
471
472    #[cfg(not(unix))]
473    pub(super) fn create(_state_dir: &Path, _endpoint: Option<&Endpoint>) -> Result<Self> {
474        Ok(Self {})
475    }
476}
477
478impl Drop for ProcessRecordGuard {
479    fn drop(&mut self) {
480        #[cfg(unix)]
481        {
482            let _ = fs::remove_file(&self.path);
483            let _ = self.file.unlock();
484        }
485    }
486}
487
488impl StartupGuard {
489    #[cfg(unix)]
490    pub(super) fn create(state_dir: &Path) -> Result<Self> {
491        let path = fs::canonicalize(state_dir)?.join(STARTUP_FILE);
492        let file = OpenOptions::new()
493            .create(true)
494            .truncate(false)
495            .read(true)
496            .write(true)
497            .open(path)?;
498        file.set_permissions(mobius::owner_only::file())?;
499        file.try_lock().map_err(|error| match error {
500            TryLockError::WouldBlock => {
501                Error::Config("gateway startup is already in progress".into())
502            }
503            TryLockError::Error(error) => error.into(),
504        })?;
505        Ok(Self { file })
506    }
507
508    #[cfg(not(unix))]
509    fn create(_state_dir: &Path) -> Result<Self> {
510        Ok(Self {})
511    }
512}
513
514impl Drop for StartupGuard {
515    fn drop(&mut self) {
516        #[cfg(unix)]
517        let _ = self.file.unlock();
518    }
519}
520
521#[cfg(any(unix, test))]
522pub(super) fn open_process_record(path: &Path) -> Result<Option<(ProcessRecord, File)>> {
523    let mut file = match OpenOptions::new().read(true).write(true).open(path) {
524        Ok(file) => file,
525        Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
526        Err(error) => return Err(error.into()),
527    };
528    if file.metadata()?.len() > MAX_PROCESS_RECORD_BYTES as u64 {
529        return Err(Error::Config("gateway process record is too large".into()));
530    }
531    file.seek(SeekFrom::Start(0))?;
532    let mut contents = Vec::new();
533    (&mut file)
534        .take(MAX_PROCESS_RECORD_BYTES as u64 + 1)
535        .read_to_end(&mut contents)?;
536    if contents.len() > MAX_PROCESS_RECORD_BYTES {
537        return Err(Error::Config("gateway process record is too large".into()));
538    }
539    let record: ProcessRecord = serde_json::from_slice(&contents)?;
540    record.validate()?;
541    Ok(Some((record, file)))
542}
543
544#[cfg(any(unix, test))]
545pub(super) fn process_is_running(file: &File) -> Result<bool> {
546    match file.try_lock() {
547        Ok(()) => {
548            file.unlock()?;
549            Ok(false)
550        }
551        Err(TryLockError::WouldBlock) => Ok(true),
552        Err(TryLockError::Error(error)) => Err(error.into()),
553    }
554}
555
556#[cfg(unix)]
557pub(super) fn running_process_pid(path: &Path) -> Result<Option<u32>> {
558    Ok(running_process_record(path)?.map(|record| record.pid))
559}
560
561#[cfg(unix)]
562pub(super) fn running_process_record(path: &Path) -> Result<Option<ProcessRecord>> {
563    let Some((record, file)) = open_process_record(path)? else {
564        return Ok(None);
565    };
566    Ok(process_is_running(&file)?.then_some(record))
567}
568
569#[cfg(not(unix))]
570pub(super) fn unsupported_lifecycle() -> Error {
571    Error::Config("gateway process lifecycle commands require macOS or Linux".into())
572}