feat(cli): add gai pr subcommand for AI-powered pull request creation

This commit is contained in:
2026-06-11 00:24:51 +08:00
parent 76a5bac11a
commit d9b037c1fc
5 changed files with 457 additions and 11 deletions
+170
View File
@@ -19,6 +19,18 @@ import { generateCommitMessage } from "./src/ai";
import { copyToClipboard } from "./src/clipboard";
import { BOLD, GREEN, YELLOW, CYAN, RED, DIM, RESET } from "./src/terminal";
import type { Config } from "./src/types";
import {
getDefaultBranch,
getBranchName,
getBranchCommits,
getBranchDiff,
detectPlatform,
checkCLI,
checkAuth,
createPR,
} from "./src/pr";
import { PR_SYSTEM_PROMPT, buildPRPrompt } from "./src/prompt";
import { generatePRMessage } from "./src/ai";
const args = process.argv.slice(2);
@@ -31,6 +43,8 @@ ${BOLD}Usage:${RESET}
gai commit Generate commit message for staged/changed files
gai commit --auto Auto-stage all changed files
gai commit -d Generate message without committing
gai pr Create a PR with AI-generated title and body
gai pr --draft Create a draft PR
gai config Configure API settings
gai --help Show this help message
gai --version Show version
@@ -277,6 +291,7 @@ interface MenuAction {
const MENU_ACTIONS: MenuAction[] = [
{ key: "commit", label: "commit", description: "Generate AI commit message" },
{ key: "pr", label: "pr", description: "Create a PR with AI-generated title" },
{ key: "config", label: "config", description: "Configure API settings" },
];
@@ -385,6 +400,8 @@ async function showMenu(): Promise<void> {
if (selected.key === "commit") {
handleCommit(false, false).then(resolve);
} else if (selected.key === "pr") {
handlePR(false).then(resolve);
} else if (selected.key === "config") {
handleConfig().then(resolve);
} else {
@@ -517,6 +534,153 @@ async function handleCommit(autoMode: boolean, dryRun: boolean): Promise<void> {
}
}
async function handlePR(draft: boolean): Promise<void> {
const config = await loadConfig();
if (!config.apiKey) {
console.error(
` ${RED}Error: API key not set. Run ${BOLD}gai config${RESET}${RED} to configure.${RESET}`,
);
process.exit(1);
}
if (!(await isGitRepo())) {
console.error(` ${RED}Error: Not a git repository.${RESET}`);
process.exit(1);
}
const platform = await detectPlatform();
if (!platform) {
console.error(
` ${RED}Error: Could not detect GitHub or Gitea from origin remote URL.${RESET}`,
);
process.exit(1);
}
const platformLabel = platform === "github" ? "GitHub" : "Gitea";
console.log(` Remote platform: ${CYAN}${platformLabel}${RESET}`);
const cliError = checkCLI(platform);
if (cliError) {
console.error(` ${RED}Error: ${cliError}${RESET}`);
process.exit(1);
}
const authError = await checkAuth(platform);
if (authError) {
console.error(` ${RED}Error: ${authError}${RESET}`);
process.exit(1);
}
const baseBranch = await getDefaultBranch();
const branchName = await getBranchName();
if (branchName === baseBranch) {
console.error(
` ${RED}Error: You are on the default branch (${baseBranch}). Switch to a feature branch first.${RESET}`,
);
process.exit(1);
}
console.log(
` Branch: ${CYAN}${branchName}${RESET} → base: ${CYAN}${baseBranch}${RESET}`,
);
const commits = await getBranchCommits(baseBranch);
if (commits.length === 0) {
console.error(
` ${RED}Error: No commits on ${branchName} compared to ${baseBranch}. Commit something first.${RESET}`,
);
process.exit(1);
}
console.log(
` ${commits.length} commit${commits.length > 1 ? "s" : ""} on this branch`,
);
const diff = await getBranchDiff(baseBranch);
if (!diff) {
console.error(` ${RED}Error: No diff from base branch.${RESET}`);
process.exit(1);
}
const MAX_DIFF_SIZE = 15000;
const truncatedDiff =
diff.length > MAX_DIFF_SIZE
? diff.substring(0, MAX_DIFF_SIZE) + "\n... (truncated)"
: diff;
const repoRoot = await getRepoRoot();
const projectCtx = await collectProjectContext(repoRoot);
const userPrompt = buildPRPrompt({
readme: projectCtx.readme,
packageDescription: projectCtx.packageDescription,
structure: projectCtx.structure,
branchName,
baseBranch,
branchCommits: commits,
diff: truncatedDiff,
});
console.log("\n Generating PR title...");
let title: string;
let body: string;
try {
const result = await generatePRMessage(config, PR_SYSTEM_PROMPT, userPrompt);
title = result.title;
body = result.body;
} catch (err) {
console.error(
` ${RED}AI request failed: ${err instanceof Error ? err.message : err}${RESET}`,
);
process.exit(1);
}
console.log(`\n ${BOLD}Generated PR:${RESET}`);
console.log(` Title: ${GREEN}${title}${RESET}`);
if (body) {
console.log(
` Body: ${DIM}${body.replace(/\n/g, "\n ")}${RESET}`,
);
}
console.log("");
const answer = await ask(` Create this PR? [${GREEN}Y${RESET}/n/e] `);
const lower = answer.toLowerCase();
if (lower === "n") {
console.log(" Aborted.");
return;
}
if (lower === "e") {
const newTitle = await ask(" Title: ");
const newBody = await ask(" Body (optional): ");
if (!newTitle.trim()) {
console.log(" Aborted.");
return;
}
title = newTitle;
body = newBody;
}
console.log(`\n Creating PR...`);
try {
const url = await createPR(platform, title, body, baseBranch, draft);
console.log(` ${GREEN}${BOLD}✔ PR created!${RESET}`);
console.log(` ${CYAN}${url}${RESET}`);
} catch (err) {
console.error(
` ${RED}PR creation failed: ${err instanceof Error ? err.message : err}${RESET}`,
);
process.exit(1);
}
}
async function main() {
if (args.includes("--help") || args.includes("-h")) {
showHelp();
@@ -547,6 +711,12 @@ async function main() {
return;
}
if (subcommand === "pr") {
const draft = args.includes("--draft");
await handlePR(draft);
return;
}
if (!subcommand) {
await showMenu();
return;
+40 -10
View File
@@ -39,11 +39,10 @@ async function sleep(ms: number) {
return new Promise((resolve) => setTimeout(resolve, ms));
}
export async function generateCommitMessage(
export async function callAI(
config: Config,
systemPrompt: string,
userPrompt: string,
retries = MAX_RETRIES,
): Promise<string> {
const url = `${config.apiBase.replace(/\/$/, "")}/chat/completions`;
@@ -52,7 +51,7 @@ export async function generateCommitMessage(
{ role: "user", content: userPrompt },
];
for (let attempt = 1; attempt <= retries; attempt++) {
for (let attempt = 1; attempt <= MAX_RETRIES; attempt++) {
try {
const response = await fetch(url, {
method: "POST",
@@ -70,7 +69,7 @@ export async function generateCommitMessage(
if (!response.ok) {
const text = await response.text();
if (response.status === 429 && attempt < retries) {
if (response.status === 429 && attempt < MAX_RETRIES) {
await sleep(RETRY_DELAY * attempt);
continue;
}
@@ -89,7 +88,7 @@ export async function generateCommitMessage(
const finishReason = data.choices?.[0]?.finish_reason;
if (raw && raw.trim()) {
return cleanMessage(raw);
return raw;
}
if (finishReason === "length") {
@@ -102,22 +101,53 @@ export async function generateCommitMessage(
throw new Error("Response blocked by content filter.");
}
if (attempt < retries) {
if (attempt < MAX_RETRIES) {
await sleep(RETRY_DELAY * attempt);
continue;
}
throw new Error(
`Empty response from AI after ${retries} attempts. finish_reason: ${finishReason ?? "unknown"}`,
`Empty response from AI after ${MAX_RETRIES} attempts. finish_reason: ${finishReason ?? "unknown"}`,
);
} catch (err) {
if (attempt >= retries) throw err;
if (attempt >= MAX_RETRIES) throw err;
if (err instanceof Error && err.message.startsWith("API error")) throw err;
if (err instanceof Error && err.message.includes("max_tokens")) throw err;
if (err instanceof Error && err.message.includes("content filter")) throw err;
if (err instanceof Error && err.message.includes("content filter"))
throw err;
await sleep(RETRY_DELAY * attempt);
}
}
throw new Error("Failed to generate commit message");
throw new Error("Failed to generate response");
}
export async function generateCommitMessage(
config: Config,
systemPrompt: string,
userPrompt: string,
): Promise<string> {
const raw = await callAI(config, systemPrompt, userPrompt);
return cleanMessage(raw);
}
export async function generatePRMessage(
config: Config,
systemPrompt: string,
userPrompt: string,
): Promise<{ title: string; body: string }> {
const raw = await callAI(config, systemPrompt, userPrompt);
const cleaned = cleanMessage(raw);
const lines = cleaned.split("\n");
const title = lines[0]?.trim() || "Update";
let bodyStart = 1;
while (bodyStart < lines.length && lines[bodyStart]?.trim() === "") {
bodyStart++;
}
const body = lines.slice(bodyStart).join("\n").trim();
return { title, body };
}
+174
View File
@@ -0,0 +1,174 @@
export type Platform = "github" | "gitea";
export async function getDefaultBranch(): Promise<string> {
try {
const result =
await Bun.$`git symbolic-ref refs/remotes/origin/HEAD`.quiet().text();
return result.trim().replace("refs/remotes/origin/", "");
} catch {
try {
const branches = await Bun.$`git branch -r`.quiet().text();
for (const line of branches.split("\n")) {
const trimmed = line.trim();
if (trimmed === "origin/main" || trimmed === "origin/master") {
return trimmed.replace("origin/", "");
}
}
} catch {}
return "main";
}
}
export async function getBranchName(): Promise<string> {
const result =
await Bun.$`git rev-parse --abbrev-ref HEAD`.quiet().text();
return result.trim();
}
export async function getBranchCommits(base: string): Promise<string[]> {
try {
const result =
await Bun.$`git log --oneline origin/${base}..HEAD`.quiet().text();
return result.trim().split("\n").filter(Boolean);
} catch {
try {
const result =
await Bun.$`git log --oneline ${base}..HEAD`.quiet().text();
return result.trim().split("\n").filter(Boolean);
} catch {
return [];
}
}
}
export async function getBranchDiff(base: string): Promise<string> {
try {
const result =
await Bun.$`git diff ${base}...HEAD`.quiet().text();
return result.trim();
} catch {
return "";
}
}
export async function detectPlatform(): Promise<Platform | null> {
try {
const url = await Bun.$`git remote get-url origin`.quiet().text();
const trimmed = url.trim().toLowerCase();
const hostname = trimmed
.replace(/^(https?:\/\/|ssh:\/\/|git:\/\/)/, "")
.replace(/^[^@]+@/, "")
.split(/[:/]/)[0];
if (!hostname) return null;
if (hostname === "github.com") return "github";
if (hostname.includes("gitea")) return "gitea";
if (Bun.which("tea")) return "gitea";
if (Bun.which("gh")) return "github";
return null;
} catch {
return null;
}
}
export function checkCLI(platform: Platform): string | null {
const bin = platform === "github" ? "gh" : "tea";
const path = Bun.which(bin);
if (!path) {
if (platform === "github") {
return "GitHub CLI (gh) not found. Install: brew install gh";
}
return "Gitea CLI (tea) not found. Install from: https://gitea.com/gitea/tea";
}
return null;
}
export async function checkAuth(platform: Platform): Promise<string | null> {
if (platform === "github") {
try {
await Bun.$`gh auth status`.quiet();
return null;
} catch {
return "Not authenticated with GitHub CLI. Run: gh auth login";
}
}
try {
const result = await Bun.$`tea logins list`.quiet().text();
if (result.trim()) return null;
return "Not authenticated with Gitea CLI. Run: tea login add";
} catch {
return "Not authenticated with Gitea CLI. Run: tea login add";
}
}
export async function createPR(
platform: Platform,
title: string,
body: string,
base: string,
draft: boolean,
): Promise<string> {
if (platform === "github") {
const args = [
"pr",
"create",
"--title",
title,
"--body",
body,
"--base",
base,
];
if (draft) args.push("--draft");
const proc = Bun.spawn(["gh", ...args], {
stdout: "pipe",
stderr: "pipe",
});
const exitCode = await proc.exited;
const stdout = await new Response(proc.stdout).text();
const stderr = await new Response(proc.stderr).text();
if (exitCode !== 0) {
throw new Error(
stderr.trim() || `gh pr create failed (exit code ${exitCode})`,
);
}
const match = stdout.match(/(https?:\/\/[^\s]+)/);
return match ? match[1] : stdout.trim();
}
const args = [
"pulls",
"create",
"--title",
title,
"--description",
body,
"--base",
base,
];
const proc = Bun.spawn(["tea", ...args], {
stdout: "pipe",
stderr: "pipe",
});
const exitCode = await proc.exited;
const stdout = await new Response(proc.stdout).text();
const stderr = await new Response(proc.stderr).text();
if (exitCode !== 0) {
throw new Error(
stderr.trim() || `tea pulls create failed (exit code ${exitCode})`,
);
}
const match = stdout.match(/(https?:\/\/[^\s]+)/);
return match ? match[1] : stdout.trim();
}
+63 -1
View File
@@ -1,4 +1,4 @@
import type { ProjectContext } from "./types";
import type { PRContext, ProjectContext } from "./types";
export const SYSTEM_PROMPT = `You are an expert at writing concise, meaningful git commit messages following the Conventional Commits specification.
@@ -54,3 +54,65 @@ export function buildPrompt(context: ProjectContext): string {
return parts.join("\n");
}
export const PR_SYSTEM_PROMPT = `You are an expert at writing clear, concise pull request titles and descriptions.
Format:
<pr title>
<blank line>
<pr body>
Rules:
1. Title must be under 72 characters, in imperative mood
2. Follow the Conventional Commits style for the title (e.g., "feat(api): add user authentication")
3. Body should be 2-3 sentences in plain text explaining WHAT was changed and WHY
4. Be specific — avoid vague messages
5. Match the language and style of recent commits if provided
6. If the branch name hints at the type (e.g., "feat/..." or "fix/..."), reflect that in the title
7. Output ONLY the PR text — no markdown, no code blocks, no prefixes`;
export function buildPRPrompt(context: PRContext): string {
const parts: string[] = [];
if (
context.packageDescription ||
context.readme ||
context.structure
) {
parts.push("## Project Context");
if (context.packageDescription) {
parts.push(`Description: ${context.packageDescription}`);
}
if (context.structure) {
parts.push(`Structure: ${context.structure}`);
}
if (context.readme) {
parts.push(`README:\n${context.readme}`);
}
parts.push("");
}
parts.push("## Branch Info");
parts.push(`Branch: ${context.branchName}`);
parts.push(`Target base: ${context.baseBranch}`);
parts.push("");
if (context.branchCommits.length > 0) {
parts.push("## Commits on This Branch");
for (const c of context.branchCommits) {
parts.push(c);
}
parts.push("");
}
parts.push("## Changes (diff from base)");
parts.push("```diff");
parts.push(context.diff);
parts.push("```");
parts.push("");
parts.push(
"Generate a pull request title and brief body for the above changes.",
);
return parts.join("\n");
}
+10
View File
@@ -19,3 +19,13 @@ export interface ProjectContext {
recentCommits: string[];
diff: string;
}
export interface PRContext {
readme: string | null;
packageDescription: string | null;
structure: string | null;
branchName: string;
baseBranch: string;
branchCommits: string[];
diff: string;
}