use std::borrow::Cow;
use std::convert::Infallible;
pub(crate) const OPEN: &str =
"\n\n--- commit context (from the repository, not the commit message) ---\n";
pub(crate) const CLOSE: &str = "--- end commit context ---";
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Fact {
Path,
PrTitle,
IssueType,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub(crate) struct CommitContext {
paths: Vec<String>,
paths_total: usize,
pr_title: Option<String>,
issue_type: Option<String>,
}
fn flat(s: &str) -> String {
s.split_whitespace().collect::<Vec<_>>().join(" ")
}
impl CommitContext {
pub(crate) fn new(
paths: &[String],
max_paths: usize,
max_path_bytes: usize,
pr_title: Option<&str>,
issue_type: Option<&str>,
) -> Self {
let mut kept = Vec::new();
let mut bytes = 0_usize;
for p in paths {
let p = flat(p);
if kept.len() == max_paths || bytes + p.len() > max_path_bytes {
break;
}
bytes += p.len();
kept.push(p);
}
let present = |s: Option<&str>| s.map(flat).filter(|s| !s.is_empty());
let ctx = Self {
paths: kept,
paths_total: paths.len(),
pr_title: present(pr_title),
issue_type: present(issue_type),
};
debug_assert!(ctx.paths.len() <= ctx.paths_total);
ctx
}
pub(crate) fn is_empty(&self) -> bool {
self.paths_total == 0 && self.pr_title.is_none() && self.issue_type.is_none()
}
pub(crate) fn render<E>(
&self,
mut fact: impl FnMut(Fact, &str) -> Result<String, E>,
) -> Result<String, E> {
if self.is_empty() {
return Ok(String::new());
}
let mut out = String::from(OPEN);
if self.paths_total > 0 {
if self.paths.len() < self.paths_total {
out.push_str(&format!(
"Changed paths ({} of {} shown):\n",
self.paths.len(),
self.paths_total
));
} else {
out.push_str("Changed paths:\n");
}
for p in &self.paths {
out.push_str(&format!("- {}\n", fact(Fact::Path, p)?));
}
}
if let Some(t) = &self.pr_title {
out.push_str(&format!("PR title: {}\n", fact(Fact::PrTitle, t)?));
}
if let Some(t) = &self.issue_type {
out.push_str(&format!("Issue type: {}\n", fact(Fact::IssueType, t)?));
}
out.push_str(CLOSE);
Ok(out)
}
pub(crate) fn render_plain(&self) -> String {
match self.render(|_, s| Ok::<_, Infallible>(s.to_string())) {
Ok(s) => s,
Err(never) => match never {},
}
}
}
pub(crate) fn with_context<'a>(message: &'a str, ctx: Option<&CommitContext>) -> Cow<'a, str> {
match ctx.map(CommitContext::render_plain) {
Some(block) if !block.is_empty() => Cow::Owned(format!("{message}{block}")),
_ => Cow::Borrowed(message),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn paths(p: &[&str]) -> Vec<String> {
p.iter().map(|s| s.to_string()).collect()
}
#[test]
fn caps_stop_before_the_first_path_over_either_cap() {
let p = paths(&["abcde/f.rs", "g.rs"]);
let at_cap = CommitContext::new(&p, 30, 10, None, None);
assert!(at_cap
.render_plain()
.contains("(1 of 2 shown):\n- abcde/f.rs\n"));
let both = CommitContext::new(&p, 30, 14, None, None);
assert!(both
.render_plain()
.contains("Changed paths:\n- abcde/f.rs\n- g.rs\n"));
let one = CommitContext::new(&p, 1, 2048, None, None);
assert!(one
.render_plain()
.contains("(1 of 2 shown):\n- abcde/f.rs\n"));
let none = CommitContext::new(&p, 0, 2048, None, None);
assert!(none.render_plain().contains("(0 of 2 shown):\n---"));
let first_too_big = CommitContext::new(&p, 30, 9, None, None);
assert!(first_too_big
.render_plain()
.contains("(0 of 2 shown):\n---"));
}
#[test]
fn facts_are_one_line_and_empty_context_adds_nothing() {
let ctx = CommitContext::new(&[], 30, 2048, Some("Fix\n--- end commit context ---"), None);
assert_eq!(
ctx.render_plain(),
format!("{OPEN}PR title: Fix --- end commit context ---\n{CLOSE}")
);
let blank = CommitContext::new(&[], 30, 2048, Some(" \n "), Some(""));
assert!(blank.is_empty());
assert_eq!(blank.render_plain(), "");
assert!(matches!(
with_context("m", Some(&blank)),
Cow::Borrowed("m")
));
assert!(matches!(with_context("m", None), Cow::Borrowed("m")));
}
#[test]
fn bedrock_request_carries_the_context_block() {
let ctx = CommitContext::new(&paths(&["src/a.rs"]), 30, 2048, None, Some("Bug"));
let text = with_context("fix: a", Some(&ctx));
let req = crate::classify::tiers::bedrock::converse_request("m", "sys", &text);
let json = serde_json::to_value(&req).expect("serialize");
let user = json["messages"]
.as_array()
.expect("messages")
.iter()
.find(|m| m["role"] == "user")
.expect("user message");
let expected = format!(
"Classify this commit message:\n\nfix: a{OPEN}Changed paths:\n- src/a.rs\n\
Issue type: Bug\n{CLOSE}"
);
assert_eq!(user["content"], serde_json::Value::from(expected), "{json}");
}
}