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, } impl PamModule 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::>() .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::>(); 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) } } } }