summaryrefslogtreecommitdiff
path: root/src-repo/src/lib.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src-repo/src/lib.rs')
-rw-r--r--src-repo/src/lib.rs121
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)
+ }
+ }
+ }
+}