| use std::path::Path; |
| use std::sync::Arc; |
|
|
| use anyhow::{Context, Result}; |
| use forge_domain::*; |
| use schemars::JsonSchema; |
| use serde::Deserialize; |
|
|
| use crate::services::{ |
| AgentRegistry, AppConfigService, ProviderAuthService, ProviderService, ShellService, |
| TemplateService, |
| }; |
| use crate::{AgentProviderResolver, EnvironmentInfra, Services}; |
|
|
| |
| #[derive(thiserror::Error, Debug)] |
| pub enum GitAppError { |
| #[error("nothing to commit, working tree clean")] |
| NoChangesToCommit, |
| } |
|
|
| |
| pub struct GitApp<S> { |
| services: Arc<S>, |
| } |
|
|
| |
| #[derive(Debug, Clone)] |
| pub struct CommitResult { |
| |
| pub message: String, |
| |
| pub committed: bool, |
| |
| pub has_staged_files: bool, |
| |
| pub git_output: String, |
| } |
|
|
| |
| #[derive(Debug, Clone)] |
| struct CommitMessageDetails { |
| |
| message: String, |
| |
| has_staged_files: bool, |
| } |
|
|
| |
| #[derive(Debug, Clone, Deserialize, JsonSchema)] |
| #[serde(rename_all = "snake_case")] |
| #[schemars(title = "commit_message")] |
| pub struct CommitMessageResponse { |
| |
| pub commit_message: String, |
| } |
|
|
| |
| #[derive(Debug, Clone)] |
| struct DiffContext { |
| diff_content: String, |
| branch_name: String, |
| recent_commits: String, |
| has_staged_files: bool, |
| additional_context: Option<String>, |
| } |
|
|
| impl<S> GitApp<S> { |
| |
| pub fn new(services: Arc<S>) -> Self { |
| Self { services } |
| } |
|
|
| |
| fn truncate_diff( |
| &self, |
| diff_content: String, |
| max_diff_size: Option<usize>, |
| original_size: usize, |
| ) -> (String, bool) { |
| match max_diff_size { |
| Some(max_size) if original_size > max_size => { |
| |
| let truncated = diff_content |
| .char_indices() |
| .take_while(|(idx, _)| *idx < max_size) |
| .map(|(_, c)| c) |
| .collect::<String>(); |
| (truncated, true) |
| } |
| _ => (diff_content, false), |
| } |
| } |
| } |
|
|
| impl<S: Services + EnvironmentInfra<Config = forge_config::ForgeConfig>> GitApp<S> { |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| pub async fn commit_message( |
| &self, |
| max_diff_size: Option<usize>, |
| diff: Option<String>, |
| additional_context: Option<String>, |
| ) -> Result<CommitResult> { |
| let CommitMessageDetails { message, has_staged_files } = self |
| .generate_commit_message(max_diff_size, diff, additional_context) |
| .await?; |
|
|
| Ok(CommitResult { |
| message, |
| committed: false, |
| has_staged_files, |
| git_output: String::new(), |
| }) |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| pub async fn commit( |
| &self, |
| message: String, |
| has_staged_files: bool, |
| use_forge_committer: bool, |
| ) -> Result<CommitResult> { |
| let cwd = self.services.get_environment().cwd; |
| let flags = if has_staged_files { "" } else { " -a" }; |
| let commit_command = build_commit_command(&message, flags, use_forge_committer); |
|
|
| let commit_result = self |
| .services |
| .execute(commit_command, cwd, false, true, None, None) |
| .await |
| .context("Failed to commit changes")?; |
|
|
| if !commit_result.output.success() { |
| anyhow::bail!("Git commit failed: {}", commit_result.output.stderr); |
| } |
|
|
| |
| let git_output = if commit_result.output.stdout.is_empty() { |
| commit_result.output.stderr.clone() |
| } else if commit_result.output.stderr.is_empty() { |
| commit_result.output.stdout.clone() |
| } else { |
| format!( |
| "{}\n{}", |
| commit_result.output.stdout, commit_result.output.stderr |
| ) |
| }; |
|
|
| Ok(CommitResult { message, committed: true, has_staged_files, git_output }) |
| } |
|
|
| |
| |
| async fn generate_commit_message( |
| &self, |
| max_diff_size: Option<usize>, |
| diff: Option<String>, |
| additional_context: Option<String>, |
| ) -> Result<CommitMessageDetails> { |
| |
| let cwd = self.services.get_environment().cwd; |
|
|
| |
| let (recent_commits, branch_name) = self.fetch_git_context(&cwd).await?; |
|
|
| |
| let (diff_content, original_size, has_staged_files) = if let Some(piped_diff) = diff { |
| |
| let size = piped_diff.len(); |
| (piped_diff, size, false) |
| } else { |
| |
| self.fetch_git_diff(&cwd).await? |
| }; |
|
|
| |
| let (truncated_diff, _) = self.truncate_diff(diff_content, max_diff_size, original_size); |
|
|
| let ctx = DiffContext { |
| diff_content: truncated_diff, |
| branch_name, |
| recent_commits, |
| has_staged_files, |
| additional_context, |
| }; |
|
|
| let retry_config = self.services.get_config()?.retry.unwrap_or_default(); |
| crate::retry::retry_with_config( |
| &retry_config, |
| || self.generate_message_from_diff(ctx.clone()), |
| None::<fn(&anyhow::Error, std::time::Duration)>, |
| ) |
| .await |
| } |
|
|
| |
| async fn fetch_git_context(&self, cwd: &Path) -> Result<(String, String)> { |
| let max_commit_count = self.services.get_config()?.max_commit_count; |
| let git_log_cmd = |
| format!("git log --pretty=format:%s --abbrev-commit --max-count={max_commit_count}"); |
| let (recent_commits, branch_name) = tokio::join!( |
| self.services |
| .execute(git_log_cmd, cwd.to_path_buf(), false, true, None, None,), |
| self.services.execute( |
| "git rev-parse --abbrev-ref HEAD".into(), |
| cwd.to_path_buf(), |
| false, |
| true, |
| None, |
| None, |
| ), |
| ); |
|
|
| let recent_commits = recent_commits.context("Failed to get recent commits")?; |
| let branch_name = branch_name.context("Failed to get branch name")?; |
|
|
| Ok((recent_commits.output.stdout, branch_name.output.stdout)) |
| } |
|
|
| |
| async fn fetch_git_diff(&self, cwd: &Path) -> Result<(String, usize, bool)> { |
| let (staged_diff, unstaged_diff) = tokio::join!( |
| self.services.execute( |
| "git diff --staged".into(), |
| cwd.to_path_buf(), |
| false, |
| true, |
| None, |
| None, |
| ), |
| self.services.execute( |
| "git diff".into(), |
| cwd.to_path_buf(), |
| false, |
| true, |
| None, |
| None, |
| ) |
| ); |
|
|
| let staged_diff = staged_diff.context("Failed to get staged changes")?; |
| let unstaged_diff = unstaged_diff.context("Failed to get unstaged changes")?; |
|
|
| |
| let has_staged_files = !staged_diff.output.stdout.trim().is_empty(); |
| let diff_output = if has_staged_files { |
| staged_diff |
| } else if !unstaged_diff.output.stdout.trim().is_empty() { |
| unstaged_diff |
| } else { |
| return Err(GitAppError::NoChangesToCommit.into()); |
| }; |
|
|
| let size = diff_output.output.stdout.len(); |
| Ok((diff_output.output.stdout, size, has_staged_files)) |
| } |
|
|
| |
| async fn resolve_agent_provider_and_model( |
| &self, |
| resolver: &AgentProviderResolver<S>, |
| agent_id: Option<AgentId>, |
| ) -> Result<(Provider<url::Url>, ModelId)> { |
| let (provider_template, model) = tokio::try_join!( |
| resolver.get_provider(agent_id.clone()), |
| resolver.get_model(agent_id) |
| )?; |
| let provider = self |
| .services |
| .refresh_provider_credential(provider_template) |
| .await?; |
| Ok((provider, model)) |
| } |
|
|
| |
| async fn generate_message_from_diff(&self, ctx: DiffContext) -> Result<CommitMessageDetails> { |
| let (agent_id, commit_config) = tokio::try_join!( |
| self.services.get_active_agent_id(), |
| self.services.get_commit_config() |
| )?; |
| let agent_provider_resolver = AgentProviderResolver::new(self.services.clone()); |
|
|
| |
| |
| |
| let (provider, model) = match commit_config { |
| Some(mc) => match self.services.get_provider(mc.provider).await { |
| Ok(provider) => match self.services.refresh_provider_credential(provider).await { |
| Ok(provider) => (provider, mc.model), |
| Err(err) => { |
| tracing::warn!( |
| error = %err, |
| "Failed to refresh credentials for configured commit provider. Falling back to the active provider." |
| ); |
| self.resolve_agent_provider_and_model(&agent_provider_resolver, agent_id) |
| .await? |
| } |
| }, |
| Err(err) => { |
| tracing::warn!( |
| error = %err, |
| "Configured commit provider unavailable. Falling back to the active provider." |
| ); |
| self.resolve_agent_provider_and_model(&agent_provider_resolver, agent_id) |
| .await? |
| } |
| }, |
| None => { |
| self.resolve_agent_provider_and_model(&agent_provider_resolver, agent_id) |
| .await? |
| } |
| }; |
|
|
| let rendered_prompt = self |
| .services |
| .render_template(Template::new("{{> forge-commit-message-prompt.md }}"), &()) |
| .await?; |
|
|
| |
| let user_data = serde_json::json!({ |
| "branch_name": ctx.branch_name, |
| "recent_commit_messages": ctx.recent_commits, |
| "git_diff": ctx.diff_content, |
| "additional_context": ctx.additional_context |
| }); |
|
|
| |
| let schema = schemars::schema_for!(CommitMessageResponse); |
|
|
| let context = forge_domain::Context::default() |
| .add_message(ContextMessage::system(rendered_prompt)) |
| .add_message(ContextMessage::user( |
| serde_json::to_string(&user_data)?, |
| Some(model.clone()), |
| )) |
| .response_format(ResponseFormat::JsonSchema(Box::new(schema))); |
|
|
| |
| let stream = self.services.chat(&model, context, provider).await?; |
| let message = stream.into_full(false).await?; |
|
|
| |
| |
| let commit_message = match serde_json::from_str::<CommitMessageResponse>(&message.content) { |
| Ok(response) => response.commit_message, |
| Err(_) => { |
| |
| message.content.trim().to_string() |
| } |
| }; |
|
|
| if commit_message.is_empty() { |
| return Err(Error::Retryable(anyhow::anyhow!("Empty commit message generated")).into()); |
| } |
|
|
| Ok(CommitMessageDetails { |
| message: commit_message, |
| has_staged_files: ctx.has_staged_files, |
| }) |
| } |
| } |
|
|
| |
| |
| |
| |
| |
| fn build_commit_command(message: &str, flags: &str, use_forge_committer: bool) -> String { |
| |
| let escaped_message = message.replace('\'', r"'\''"); |
| if use_forge_committer { |
| format!( |
| "GIT_COMMITTER_NAME='ForgeCode' GIT_COMMITTER_EMAIL='noreply@forgecode.dev' git commit {flags} -m '{escaped_message}'" |
| ) |
| } else { |
| format!("git commit {flags} -m '{escaped_message}'") |
| } |
| } |
|
|
| #[cfg(test)] |
| mod tests { |
| use pretty_assertions::assert_eq; |
|
|
| use super::*; |
|
|
| #[test] |
| fn test_build_commit_command_with_forge_committer_staged() { |
| let actual = build_commit_command("feat: add feature", "", true); |
| let expected = "GIT_COMMITTER_NAME='ForgeCode' GIT_COMMITTER_EMAIL='noreply@forgecode.dev' git commit -m 'feat: add feature'"; |
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_build_commit_command_with_forge_committer_unstaged() { |
| let actual = build_commit_command("fix: bug", " -a", true); |
| let expected = "GIT_COMMITTER_NAME='ForgeCode' GIT_COMMITTER_EMAIL='noreply@forgecode.dev' git commit -a -m 'fix: bug'"; |
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_build_commit_command_without_forge_committer_staged() { |
| let actual = build_commit_command("chore: update", "", false); |
| let expected = "git commit -m 'chore: update'"; |
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_build_commit_command_without_forge_committer_unstaged() { |
| let actual = build_commit_command("docs: readme", " -a", false); |
| let expected = "git commit -a -m 'docs: readme'"; |
| assert_eq!(actual, expected); |
| } |
|
|
| #[test] |
| fn test_build_commit_command_escapes_single_quotes() { |
| let actual = build_commit_command("feat: it's done", "", true); |
| let expected = "GIT_COMMITTER_NAME='ForgeCode' GIT_COMMITTER_EMAIL='noreply@forgecode.dev' git commit -m 'feat: it'\\''s done'"; |
| assert_eq!(actual, expected); |
| } |
| } |
|
|