summaryrefslogtreecommitdiff
path: root/src-repo/src/lib.rs
blob: a32eb5e157fb33e3917925e27f3cf20201636db4 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
use std::collections::HashMap;
use std::ffi::CStr;
use std::sync::mpsc::channel;
use std::sync::{Arc, Mutex};
use std::thread;

use nonstick::handle::PamHandleModule;
use nonstick::{pam_hooks, ErrorCode, Flags, PamModule, Result as PamResult};

use crate::pam_client::{Client, PamResult as PamClientResult};
use serde::{Deserialize, Serialize};

use crate::mode::Mode;
use crate::pam_any_conversation::PamAnyConversation;
use crate::raw_conv::RawConv;
use crate::un_hide_input::un_hide_input;

mod mode;
mod pam_any_conversation;
mod pam_client;
mod raw_conv;
mod un_hide_input;

struct PamAny;
pam_hooks!(PamAny);

#[derive(Serialize, Deserialize, Debug)]
struct Input {
    mode: Mode,
    modules: HashMap<String, String>,
}

impl<T: PamHandleModule> PamModule<T> for PamAny {
    fn authenticate(handle: &mut T, args: Vec<&CStr>, _flags: Flags) -> PamResult<()> {
        let arg_string = args
            .iter()
            .map(|s| s.to_str().unwrap_or(""))
            .collect::<Vec<_>>()
            .join(" ");

        // Support reading config from a file path (for NixOS compatibility)
        let config_string = if arg_string.starts_with('/') {
            let path = arg_string.split_whitespace().next().unwrap_or("");
            std::fs::read_to_string(path).map_err(|_| ErrorCode::ServiceError)?
        } else {
            arg_string
        };

        let input: Input =
            serde_json::from_str(&config_string).map_err(|_| ErrorCode::ServiceError)?;

        let user = handle.get_user(None)?.to_owned();

        // Get the raw PAM conversation to share across threads.
        let handle_ptr = handle as *mut T as *mut libc::c_void;
        let conv = RawConv::from_pam_handle(handle_ptr)
            .ok_or(ErrorCode::ConversationError)?;
        let conv = Arc::new(Mutex::new(conv));

        let (tx, rx) = channel::<PamClientResult<()>>();
        let _handles: Vec<_> = input
            .modules
            .iter()
            .map(|(service, service_display_name)| {
                let service = service.to_owned();
                let tx = tx.clone();
                let conv = conv.clone();
                let user = user.clone();
                let service_display_name = service_display_name.to_owned();
                thread::spawn(move || {
                    let client = Client::with_conversation(
                        &service,
                        PamAnyConversation {
                            service_display_name,
                            user,
                            conv,
                        },
                    );
                    let result = match client {
                        Ok(mut c) => c.authenticate(),
                        Err(e) => Err(e),
                    };
                    let _ = tx.send(result);
                })
            })
            .collect();

        match input.mode {
            Mode::One => {
                let mut failed_modules = 0;
                for result in rx {
                    if result.is_ok() {
                        let _ = un_hide_input();
                        return Ok(());
                    } else {
                        failed_modules += 1;
                        if failed_modules == input.modules.len() {
                            return Err(ErrorCode::AuthenticationError);
                        }
                    }
                }
                Err(ErrorCode::AuthenticationError)
            }
            Mode::All => {
                let mut successful_modules = 0;
                for result in rx {
                    if result.is_ok() {
                        successful_modules += 1;
                        if successful_modules == input.modules.len() {
                            let _ = un_hide_input();
                            return Ok(());
                        }
                    } else {
                        return Err(ErrorCode::AuthenticationError);
                    }
                }
                Err(ErrorCode::AuthenticationError)
            }
        }
    }
}