use crate::persistence::{
CrossProcessFileLock, atomic_write_with_permissions, open_regular_file,
read_regular_file_bounded,
};
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use std::{
fs,
path::{Path, PathBuf},
time::{Duration, Instant},
};
const MAX_ENTRIES: usize = 200;
const MAX_PROMPT_BYTES: usize = 4096;
const MAX_STORE_BYTES: u64 = 4 * 1024 * 1024;
#[derive(Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub(crate) struct PromptHistory {
entries: Vec<HistoryEntry>,
}
impl std::fmt::Debug for PromptHistory {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PromptHistory")
.field("entries", &self.entries.len())
.finish()
}
}
#[derive(Clone, PartialEq, Eq, Serialize, Deserialize)]
struct HistoryEntry {
text: String,
count: u64,
last_used: i64,
}
pub(crate) struct PromptHistoryRecorder {
path: PathBuf,
initial_prompt: Option<String>,
}
impl PromptHistoryRecorder {
pub(crate) fn new(path: PathBuf, prompt: String) -> Self {
Self {
path,
initial_prompt: Some(prompt),
}
}
pub(crate) fn record_event(
&mut self,
event: &crate::output::OutputEvent,
) -> anyhow::Result<Option<PromptHistory>> {
use crate::output::OutputEvent;
let prompt = match event {
OutputEvent::UserPrompt { .. } => self.initial_prompt.take(),
_ => None,
};
prompt
.map(|prompt| PromptHistory::remember(&self.path, &prompt))
.transpose()
}
pub(crate) fn record_steering_prompt(&self, prompt: &str) -> anyhow::Result<PromptHistory> {
PromptHistory::remember(&self.path, prompt)
}
}
pub(crate) fn enabled() -> bool {
std::env::var_os("MC_PROMPT_HISTORY").is_none_or(|value| value != "0")
}
fn scope_identity(cwd: &Path) -> anyhow::Result<PathBuf> {
let cwd = cwd.canonicalize()?;
for parent in cwd.ancestors() {
let marker = parent.join(".git");
let git = if marker.is_dir() {
marker
} else if marker.is_file() {
let text = read_text(&marker, 8192)?;
let target = text
.trim()
.strip_prefix("gitdir:")
.ok_or_else(|| anyhow::anyhow!("invalid Git directory marker"))?;
parent.join(target.trim())
} else {
continue;
};
let common = git.join("commondir");
let identity = if common.exists() {
git.join(read_text(&common, 8192)?.trim())
} else {
git
};
return Ok(identity.canonicalize()?);
}
Ok(cwd)
}
fn read_text(path: &Path, limit: u64) -> anyhow::Result<String> {
let file = open_regular_file(path, false)?;
Ok(String::from_utf8(read_regular_file_bounded(&file, limit)?)?)
}
pub(crate) fn store_path(state_dir: &Path, cwd: &Path) -> anyhow::Result<PathBuf> {
let identity = scope_identity(cwd)?;
let digest = Sha256::digest(identity.as_os_str().as_encoded_bytes());
let digest: String = digest.iter().map(|byte| format!("{byte:02x}")).collect();
Ok(state_dir
.join("prompt-history")
.join(format!("{digest}.json")))
}
fn check_history_directory(path: &Path) -> anyhow::Result<()> {
let parent = path
.parent()
.ok_or_else(|| anyhow::anyhow!("missing history directory"))?;
match fs::symlink_metadata(parent) {
Ok(metadata) if metadata.file_type().is_symlink() || !metadata.is_dir() => {
anyhow::bail!("history directory must be a regular directory")
}
Err(error) if error.kind() != std::io::ErrorKind::NotFound => return Err(error.into()),
_ => {}
}
Ok(())
}
impl PromptHistory {
pub(crate) fn load(path: &Path) -> anyhow::Result<Self> {
check_history_directory(path)?;
let text = match read_text(path, MAX_STORE_BYTES) {
Ok(text) => text,
Err(error)
if error
.downcast_ref::<std::io::Error>()
.is_some_and(|e| e.kind() == std::io::ErrorKind::NotFound) =>
{
return Ok(Self::default());
}
Err(error) => return Err(error),
};
let mut history: Self = serde_json::from_str(&text)?;
history.entries.retain(|entry| eligible(&entry.text));
history.bound();
Ok(history)
}
pub(crate) fn remember(path: &Path, text: &str) -> anyhow::Result<Self> {
check_history_directory(path)?;
fs::create_dir_all(
path.parent()
.ok_or_else(|| anyhow::anyhow!("missing history directory"))?,
)?;
let _lock =
CrossProcessFileLock::acquire_until(path, Instant::now() + Duration::from_millis(100))?;
let mut history = Self::load(path)?;
history.record(text, chrono::Utc::now().timestamp_millis());
atomic_write_with_permissions(path, &serde_json::to_vec(&history)?, Some(0o600))?;
Ok(history)
}
fn record(&mut self, text: &str, now: i64) {
if !eligible(text) {
return;
}
if let Some(entry) = self.entries.iter_mut().find(|entry| entry.text == text) {
entry.count = entry.count.saturating_add(1);
entry.last_used = now;
} else {
self.entries.push(HistoryEntry {
text: text.to_owned(),
count: 1,
last_used: now,
});
}
self.bound();
}
fn bound(&mut self) {
self.entries.sort_by(|a, b| {
b.count
.cmp(&a.count)
.then_with(|| b.last_used.cmp(&a.last_used))
.then_with(|| a.text.cmp(&b.text))
});
self.entries.truncate(MAX_ENTRIES);
}
pub(crate) fn suffix(&self, prefix: &str) -> Option<&str> {
if prefix.is_empty() || prefix.contains(['\n', '\r']) {
return None;
}
self.entries.iter().find_map(|entry| {
entry
.text
.strip_prefix(prefix)
.filter(|suffix| !suffix.trim().is_empty())
})
}
}
fn eligible(text: &str) -> bool {
!text.trim().is_empty() && text.len() <= MAX_PROMPT_BYTES
&& !text.chars().any(char::is_control)
&& !text.starts_with(['/', '!'])
&& crate::output::redact_sensitive_text(text) == text
&& !text.split(|ch: char| ch.is_whitespace() || "=\"'([{@`".contains(ch)).any(|word| {
word.starts_with('/') || word.starts_with("~/") || word.contains(":\\")
|| (word.as_bytes().get(1..3) == Some(b":/"))
})
}