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 =
262            configured_endpoint.map_or_else(Endpoint::from_env, |endpoint| Ok(endpoint.clone()))?;
263        if endpoint.is_plaintext() || endpoint.is_websocket() {
264            return Err(Error::Config(
265                "TLS gateway upgrades require MOBIUS_GATEWAY_ENDPOINT with the certificate hostname".into(),
266            ));
267        }
268        endpoint
269    } else {
270        loopback_endpoint(config)?
271    };
272    let token = load_local_client(&endpoint)?.ok_or_else(|| {
273        Error::Config("local gateway credential is unavailable; cannot check its version".into())
274    })?;
275    let version = tokio::time::timeout(
276        Duration::from_secs(2),
277        endpoint.local_gateway_version(config.listen, &token),
278    )
279    .await
280    .map_err(|_| Error::Config("gateway version check timed out".into()))??;
281    if !gateway_version_is_older(&version, env!("CARGO_PKG_VERSION"))? {
282        return Ok(Some(process));
283    }
284    let state_dir = store.state_dir().to_path_buf();
285    tokio::task::spawn_blocking(move || stop_gateway(&state_dir, Some(process.pid)))
286        .await
287        .map_err(|error| Error::Config(format!("gateway stop task failed: {error}")))??;
288    Ok(None)
289}
290
291#[cfg(any(unix, test))]
292pub(super) fn gateway_version_is_older(running: &str, starting: &str) -> Result<bool> {
293    let parse = |version: &str| {
294        semver::Version::parse(version)
295            .map_err(|error| Error::Config(format!("invalid gateway version `{version}`: {error}")))
296    };
297    Ok(parse(running)?.cmp_precedence(&parse(starting)?).is_lt())
298}
299
300#[cfg(unix)]
301pub(super) async fn stop_background_child(child: &mut Child, process_path: &Path) {
302    if let Some(pid) = child.id() {
303        terminate_process_group(pid);
304    }
305    let _ = child.wait().await;
306    remove_unlocked_process_record(process_path);
307}
308
309#[cfg(unix)]
310pub(super) fn remove_unlocked_process_record(path: &Path) {
311    let Ok(file) = OpenOptions::new().read(true).write(true).open(path) else {
312        return;
313    };
314    if file.try_lock().is_ok() {
315        let _ = fs::remove_file(path);
316    }
317}
318
319#[cfg(unix)]
320pub(super) fn background_startup_error(
321    message: impl std::fmt::Display,
322    log: &tempfile::NamedTempFile,
323) -> Error {
324    startup_error(message, log, MAX_BACKGROUND_ERROR_BYTES)
325}
326
327/// Adds a bounded diagnostic excerpt from a gateway startup log.
328pub fn startup_error(
329    message: impl std::fmt::Display,
330    log: &tempfile::NamedTempFile,
331    max_bytes: u64,
332) -> Error {
333    let mut details = String::new();
334    if let Ok(file) = File::open(log.path()) {
335        let _ = file.take(max_bytes).read_to_string(&mut details);
336    }
337    let details = details.trim();
338    Error::Config(if details.is_empty() {
339        message.to_string()
340    } else {
341        format!("{message}: {details}")
342    })
343}
344
345pub(super) fn print_listener(config: &GatewayConfig, runtime_endpoint: Option<&Endpoint>) {
346    if let Some(cloudflare) = &config.cloudflare {
347        if let Some(endpoint) = runtime_endpoint
348            .map(ToString::to_string)
349            .or_else(|| cloudflare.endpoint())
350        {
351            println!("public endpoint: {endpoint}");
352        } else {
353            println!("public endpoint: assigned when the gateway starts");
354        }
355        println!("local endpoint: tcp://{}", config.listen);
356        println!("tunnel origin: http://{}", config.listen);
357        return;
358    }
359    let scheme = if config.tls.is_some() { "tls" } else { "tcp" };
360    println!("listener: {scheme}://{}", config.listen);
361}
362
363#[cfg(unix)]
364pub(super) fn exit_gateway(state_dir: PathBuf) -> Result<()> {
365    let (store, _) = ConfigStore::open(state_dir)?;
366    let _startup = StartupGuard::create(store.state_dir())?;
367    stop_gateway(store.state_dir(), None)
368}
369
370#[cfg(unix)]
371pub(super) fn stop_gateway(state_dir: &Path, expected_pid: Option<u32>) -> Result<()> {
372    let path = state_dir.join(PROCESS_FILE);
373    let Some((record, file)) = open_process_record(&path)? else {
374        eprintln!("gateway is stopped");
375        return Ok(());
376    };
377    if !process_is_running(&file)? {
378        eprintln!("gateway is stopped");
379        return Ok(());
380    }
381    if let Some(expected_pid) = expected_pid
382        && record.pid != expected_pid
383    {
384        return Err(Error::Config(format!(
385            "gateway process changed from {expected_pid} to {}",
386            record.pid
387        )));
388    }
389    let pid = i32::try_from(record.pid)
390        .map(Pid::from_raw)
391        .map_err(|_| Error::Config("invalid gateway process record".into()))?;
392    if let Err(error) = kill(pid, Signal::SIGINT) {
393        if !process_is_running(&file)? {
394            eprintln!("gateway is stopped");
395            return Ok(());
396        }
397        return Err(Error::Config(format!(
398            "failed to interrupt gateway: {error}"
399        )));
400    }
401    let started = Instant::now();
402    while process_is_running(&file)? {
403        if started.elapsed() >= EXIT_TIMEOUT {
404            return Err(Error::Config(format!(
405                "gateway process {} did not stop within {} seconds",
406                record.pid,
407                EXIT_TIMEOUT.as_secs()
408            )));
409        }
410        std::thread::sleep(EXIT_POLL_INTERVAL);
411    }
412    eprintln!("gateway stopped");
413    Ok(())
414}
415
416#[cfg(not(unix))]
417pub(super) fn exit_gateway(_state_dir: PathBuf) -> Result<()> {
418    Err(unsupported_lifecycle())
419}
420
421#[cfg(any(unix, test))]
422impl ProcessRecord {
423    pub(super) fn validate(&self) -> Result<()> {
424        if self.pid == 0 || i32::try_from(self.pid).is_err() {
425            return Err(Error::Config("invalid gateway process record".into()));
426        }
427        if let Some(endpoint) = self.endpoint()?
428            && !endpoint.is_websocket()
429        {
430            return Err(Error::Config(
431                "gateway process endpoint must use wss://".into(),
432            ));
433        }
434        Ok(())
435    }
436
437    pub(super) fn endpoint(&self) -> Result<Option<Endpoint>> {
438        self.endpoint.as_deref().map(str::parse).transpose()
439    }
440}
441
442impl ProcessRecordGuard {
443    #[cfg(unix)]
444    pub(super) fn create(state_dir: &Path, endpoint: Option<&Endpoint>) -> Result<Self> {
445        let state_dir = fs::canonicalize(state_dir)?;
446        let path = state_dir.join(PROCESS_FILE);
447        let mut file = OpenOptions::new()
448            .create(true)
449            .truncate(false)
450            .read(true)
451            .write(true)
452            .open(&path)?;
453        file.set_permissions(mobius::owner_only::file())?;
454        file.try_lock().map_err(|error| match error {
455            TryLockError::WouldBlock => Error::Config("gateway is already running".into()),
456            TryLockError::Error(error) => error.into(),
457        })?;
458        file.set_len(0)?;
459        file.seek(SeekFrom::Start(0))?;
460        let record = ProcessRecord {
461            pid: std::process::id(),
462            endpoint: endpoint.map(ToString::to_string),
463        };
464        serde_json::to_writer(&mut file, &record)?;
465        file.flush()?;
466        file.sync_all()?;
467        Ok(Self { path, file })
468    }
469
470    #[cfg(not(unix))]
471    pub(super) fn create(_state_dir: &Path, _endpoint: Option<&Endpoint>) -> Result<Self> {
472        Ok(Self {})
473    }
474}
475
476impl Drop for ProcessRecordGuard {
477    fn drop(&mut self) {
478        #[cfg(unix)]
479        {
480            let _ = fs::remove_file(&self.path);
481            let _ = self.file.unlock();
482        }
483    }
484}
485
486impl StartupGuard {
487    #[cfg(unix)]
488    pub(super) fn create(state_dir: &Path) -> Result<Self> {
489        let path = fs::canonicalize(state_dir)?.join(STARTUP_FILE);
490        let file = OpenOptions::new()
491            .create(true)
492            .truncate(false)
493            .read(true)
494            .write(true)
495            .open(path)?;
496        file.set_permissions(mobius::owner_only::file())?;
497        file.try_lock().map_err(|error| match error {
498            TryLockError::WouldBlock => {
499                Error::Config("gateway startup is already in progress".into())
500            }
501            TryLockError::Error(error) => error.into(),
502        })?;
503        Ok(Self { file })
504    }
505
506    #[cfg(not(unix))]
507    fn create(_state_dir: &Path) -> Result<Self> {
508        Ok(Self {})
509    }
510}
511
512impl Drop for StartupGuard {
513    fn drop(&mut self) {
514        #[cfg(unix)]
515        let _ = self.file.unlock();
516    }
517}
518
519#[cfg(any(unix, test))]
520pub(super) fn open_process_record(path: &Path) -> Result<Option<(ProcessRecord, File)>> {
521    let mut file = match OpenOptions::new().read(true).write(true).open(path) {
522        Ok(file) => file,
523        Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(None),
524        Err(error) => return Err(error.into()),
525    };
526    if file.metadata()?.len() > MAX_PROCESS_RECORD_BYTES as u64 {
527        return Err(Error::Config("gateway process record is too large".into()));
528    }
529    file.seek(SeekFrom::Start(0))?;
530    let mut contents = Vec::new();
531    (&mut file)
532        .take(MAX_PROCESS_RECORD_BYTES as u64 + 1)
533        .read_to_end(&mut contents)?;
534    if contents.len() > MAX_PROCESS_RECORD_BYTES {
535        return Err(Error::Config("gateway process record is too large".into()));
536    }
537    let record: ProcessRecord = serde_json::from_slice(&contents)?;
538    record.validate()?;
539    Ok(Some((record, file)))
540}
541
542#[cfg(any(unix, test))]
543pub(super) fn process_is_running(file: &File) -> Result<bool> {
544    match file.try_lock() {
545        Ok(()) => {
546            file.unlock()?;
547            Ok(false)
548        }
549        Err(TryLockError::WouldBlock) => Ok(true),
550        Err(TryLockError::Error(error)) => Err(error.into()),
551    }
552}
553
554#[cfg(unix)]
555pub(super) fn running_process_pid(path: &Path) -> Result<Option<u32>> {
556    Ok(running_process_record(path)?.map(|record| record.pid))
557}
558
559#[cfg(unix)]
560pub(super) fn running_process_record(path: &Path) -> Result<Option<ProcessRecord>> {
561    let Some((record, file)) = open_process_record(path)? else {
562        return Ok(None);
563    };
564    Ok(process_is_running(&file)?.then_some(record))
565}
566
567#[cfg(not(unix))]
568pub(super) fn unsupported_lifecycle() -> Error {
569    Error::Config("gateway process lifecycle commands require macOS or Linux".into())
570}