1
#[cfg(target_os = "linux")]
2
use std::{
3
    io::{self, IoSlice, Read},
4
    mem::MaybeUninit,
5
    os::{fd::AsFd, unix::fs::PermissionsExt},
6
    path::PathBuf,
7
    sync::LazyLock,
8
};
9

            
10
#[cfg(target_os = "linux")]
11
use rustix::{
12
    fs::{MemfdFlags, SealFlags},
13
    net::{SendAncillaryBuffer, SendAncillaryMessage, SendFlags},
14
};
15

            
16
#[cfg(target_os = "linux")]
17
const HELPER_TIMEOUT_SECS: u64 = 120;
18
#[cfg(target_os = "linux")]
19
const BINARY_NAME: &str = env!("CARGO_BIN_NAME");
20

            
21
#[cfg(target_os = "linux")]
22
pub static SOCKET_PATH: LazyLock<PathBuf> = LazyLock::new(|| {
23
    let uid = rustix::process::getuid().as_raw();
24
    PathBuf::from(format!("/run/user/{uid}/oo7-daemon-login.sock"))
25
});
26

            
27
#[cfg(target_os = "linux")]
28
fn main() {
29
    tracing_subscriber::fmt::init();
30

            
31
    let mut secret = vec![];
32
    io::stdin()
33
        .lock()
34
        .read_to_end(&mut secret)
35
        .unwrap_or_else(|e| {
36
            tracing::error!("Failed to read secret from stdin: {e}");
37
            std::process::exit(1);
38
        });
39

            
40
    if secret.is_empty() {
41
        tracing::error!("No secret provided on stdin");
42
        std::process::exit(1);
43
    }
44

            
45
    tracing::info!("Starting {BINARY_NAME}");
46

            
47
    let socket_path = &*SOCKET_PATH;
48

            
49
    let _ = std::fs::remove_file(socket_path);
50

            
51
    let listener = std::os::unix::net::UnixListener::bind(socket_path).unwrap_or_else(|e| {
52
        tracing::error!("Failed to bind {}: {e}", socket_path.display());
53
        std::process::exit(1);
54
    });
55

            
56
    std::fs::set_permissions(socket_path, std::fs::Permissions::from_mode(0o600)).unwrap_or_else(
57
        |e| {
58
            tracing::error!("Failed to set socket permissions: {e}");
59
            std::process::exit(1);
60
        },
61
    );
62

            
63
    let uid = rustix::process::getuid();
64

            
65
    listener.set_nonblocking(false).ok();
66
    tracing::info!("Listening on {}", socket_path.display());
67

            
68
    // Poll with timeout
69
    let timeout = rustix::event::Timespec {
70
        tv_sec: HELPER_TIMEOUT_SECS as i64,
71
        tv_nsec: 0,
72
    };
73
    let poll_ret = rustix::event::poll(
74
        &mut [rustix::event::PollFd::from_borrowed_fd(
75
            listener.as_fd(),
76
            rustix::event::PollFlags::IN,
77
        )],
78
        Some(&timeout),
79
    );
80
    match poll_ret {
81
        Ok(0) => {
82
            tracing::info!("Timed out after {HELPER_TIMEOUT_SECS}s, no daemon connected");
83
            let _ = std::fs::remove_file(socket_path);
84
            std::process::exit(0);
85
        }
86
        Err(e) => {
87
            tracing::error!("Poll failed: {e}");
88
            let _ = std::fs::remove_file(socket_path);
89
            std::process::exit(1);
90
        }
91
        _ => {}
92
    }
93

            
94
    let (stream, _addr) = listener.accept().unwrap_or_else(|e| {
95
        tracing::error!("Failed to accept connection: {e}");
96
        let _ = std::fs::remove_file(socket_path);
97
        std::process::exit(1);
98
    });
99

            
100
    let peer_cred = rustix::net::sockopt::socket_peercred(&stream).unwrap_or_else(|e| {
101
        tracing::error!("Failed to get peer credentials: {e}");
102
        let _ = std::fs::remove_file(socket_path);
103
        std::process::exit(1);
104
    });
105
    if peer_cred.uid != uid {
106
        tracing::error!(
107
            "Rejected connection from UID {} (expected {})",
108
            peer_cred.uid.as_raw(),
109
            uid.as_raw()
110
        );
111
        let _ = std::fs::remove_file(socket_path);
112
        std::process::exit(1);
113
    }
114

            
115
    // Create memfd, write secret, seal it
116
    let memfd = rustix::fs::memfd_create(
117
        c"oo7-login-secret",
118
        MemfdFlags::CLOEXEC | MemfdFlags::ALLOW_SEALING,
119
    )
120
    .unwrap_or_else(|e| {
121
        tracing::error!("Failed to create memfd: {e}");
122
        let _ = std::fs::remove_file(socket_path);
123
        std::process::exit(1);
124
    });
125

            
126
    rustix::io::write(&memfd, &secret).unwrap_or_else(|e| {
127
        tracing::error!("Failed to write to memfd: {e}");
128
        let _ = std::fs::remove_file(socket_path);
129
        std::process::exit(1);
130
    });
131

            
132
    zeroize::Zeroize::zeroize(&mut secret);
133

            
134
    rustix::fs::fcntl_add_seals(
135
        &memfd,
136
        SealFlags::WRITE | SealFlags::SHRINK | SealFlags::GROW | SealFlags::SEAL,
137
    )
138
    .unwrap_or_else(|e| {
139
        tracing::error!("Failed to seal memfd: {e}");
140
        let _ = std::fs::remove_file(socket_path);
141
        std::process::exit(1);
142
    });
143

            
144
    // Send the memfd via SCM_RIGHTS
145
    let fds = [std::os::fd::AsFd::as_fd(&memfd)];
146
    let mut space = [MaybeUninit::uninit(); rustix::cmsg_space!(ScmRights(1))];
147
    let mut cmsg_buf = SendAncillaryBuffer::new(&mut space);
148
    cmsg_buf.push(SendAncillaryMessage::ScmRights(&fds));
149

            
150
    let iov = [IoSlice::new(&[0u8])];
151
    rustix::net::sendmsg(&stream, &iov, &mut cmsg_buf, SendFlags::empty()).unwrap_or_else(|e| {
152
        tracing::error!("Failed to send memfd: {e}");
153
        let _ = std::fs::remove_file(socket_path);
154
        std::process::exit(1);
155
    });
156

            
157
    tracing::info!("Secret delivered to daemon");
158
    let _ = std::fs::remove_file(socket_path);
159
}
160

            
161
#[cfg(not(target_os = "linux"))]
162
fn main() {}