diff options
Diffstat (limited to 'src-repo/src/lib.rs')
| -rw-r--r-- | src-repo/src/lib.rs | 121 |
1 files changed, 121 insertions, 0 deletions
diff --git a/src-repo/src/lib.rs b/src-repo/src/lib.rs new file mode 100644 index 0000000..a32eb5e --- /dev/null +++ b/src-repo/src/lib.rs @@ -0,0 +1,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) + } + } + } +} |
