use crate::messages::InputResponse;
use anyhow::{Context, Result, anyhow};
use interprocess::local_socket::{
GenericFilePath, Listener, ListenerNonblockingMode, ListenerOptions, Stream, ToFsName,
traits::{Listener as _, Stream as _},
};
use std::{
collections::HashMap,
env,
ffi::OsString,
io::{self, BufRead, BufReader, ErrorKind, Write},
path::PathBuf,
sync::{
Arc, Mutex,
atomic::{AtomicBool, Ordering},
},
thread::{self, JoinHandle},
time::Duration,
};
use uuid::Uuid;
pub(crate) struct AskpassThread {
pub environment: HashMap<OsString, OsString>,
stop_flag: Arc<AtomicBool>,
handle: Option<JoinHandle<()>>,
}
impl Drop for AskpassThread {
fn drop(&mut self) {
self.stop_flag.store(true, Ordering::Relaxed);
if let Some(handle) = self.handle.take() {
let _ = handle.join();
}
}
}
pub fn run_askpass() -> Option<Result<()>> {
let socket_path = env::var("GG_ASKPASS_SOCKET").ok()?;
Some(askpass_client(socket_path).context("askpass_client"))
}
pub(crate) fn serve_askpass(
input: Option<InputResponse>,
prompts: Arc<Mutex<Vec<String>>>,
) -> AskpassThread {
let stop_flag = Arc::new(AtomicBool::new(false));
let socket_name = format!("gg-askpass-{}", Uuid::new_v4());
let socket_path = if cfg!(windows) {
PathBuf::from(format!(r"\\.\pipe\{}", socket_name))
} else {
env::temp_dir().join(format!("{}.sock", socket_name))
};
let handle = match socket_path.clone().to_fs_name::<GenericFilePath>() {
Ok(name) => match ListenerOptions::new().name(name).create_sync() {
Ok(listener) => {
let response = input;
let prompts_clone = Arc::clone(&prompts);
let stop_flag_clone = Arc::clone(&stop_flag);
#[cfg(unix)]
let socket_path_clone = socket_path.clone();
Some(thread::spawn(move || {
if let Err(err) =
askpass_server(listener, response, prompts_clone, stop_flag_clone)
{
log::error!("askpass_server: {:#?}", err);
}
#[cfg(unix)]
let _ = std::fs::remove_file(&socket_path_clone);
}))
}
Err(e) => {
log::warn!("Failed to start askpass server: {}", e);
None
}
},
Err(e) => {
log::warn!("Invalid askpass socket path: {}", e);
None
}
};
let environment = if handle.is_some()
&& let Ok(exe_path) = env::current_exe()
{
HashMap::from([
("GIT_ASKPASS".into(), exe_path.clone().into()),
("GIT_TERMINAL_PROMPT".into(), "0".into()),
("SSH_ASKPASS".into(), exe_path.into()),
("SSH_ASKPASS_REQUIRE".into(), "force".into()),
("GG_ASKPASS_SOCKET".into(), socket_path.into()),
])
} else {
HashMap::new()
};
AskpassThread {
environment,
stop_flag,
handle,
}
}
fn askpass_client(socket_path: String) -> Result<()> {
let prompt = env::args().nth(1).unwrap_or_default();
let name = socket_path
.to_fs_name::<GenericFilePath>()
.map_err(|e| anyhow!("invalid socket path: {}", e))?;
let stream =
Stream::connect(name).map_err(|e| anyhow!("failed to connect to askpass socket: {}", e))?;
let mut writer = &stream;
writeln!(writer, "{}", prompt)?;
writer.flush()?;
let mut response = String::new();
BufReader::new(&stream).read_line(&mut response)?;
let response = response.trim();
if let Some(credential) = response.strip_prefix("OK:") {
println!("{credential}");
Ok(())
} else {
Err(anyhow!("credential unavailable"))
}
}
pub(crate) fn askpass_server(
listener: Listener,
response: Option<InputResponse>,
prompts: Arc<Mutex<Vec<String>>>,
stop_flag: Arc<AtomicBool>,
) -> Result<()> {
listener
.set_nonblocking(ListenerNonblockingMode::Both)
.context("set_nonblocking")?;
while !stop_flag.load(Ordering::Relaxed) {
match listener.accept() {
Ok(stream) => {
handle_askpass_request(stream, &response, &prompts)
.context("handle_askpass_request")?;
}
Err(ref e) if e.kind() == ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(100));
}
Err(e) => {
return Err(e.into());
}
}
}
Ok(())
}
fn handle_askpass_request(
mut stream: Stream,
response: &Option<InputResponse>,
prompts: &Mutex<Vec<String>>,
) -> io::Result<()> {
stream.set_nonblocking(false)?;
let mut reader = BufReader::new(&stream);
let mut prompt = String::new();
reader.read_line(&mut prompt)?;
let prompt = prompt.trim();
log::debug!("askpass prompt: {}", prompt);
let credential = if let Some(input) = response
&& let Some(field) = input.fields.get(prompt)
{
format!("OK:{}", field)
} else {
prompts.lock().unwrap().push(prompt.to_owned());
"NO".to_string()
};
writeln!(stream, "{credential}")?;
stream.flush()?;
log::debug!("askpass credential: {}", credential);
Ok(())
}