feat: 重构为 Monorepo 架构并实现 HTTP Server

架构变更:
- 采用 pnpm workspaces 实现 Monorepo 结构
- 将现有代码迁移到 packages/core
- 新增 packages/server HTTP 服务层

Server 功能:
- REST API: 会话管理、工具管理、配置管理
- WebSocket: 实时双向通信支持
- SSE: 服务端事件推送
- Hono + Bun 作为运行时

API 端点:
- GET/POST /api/sessions - 会话 CRUD
- GET/POST /api/sessions/:id/messages - 消息管理
- GET /api/sessions/:id/events - SSE 事件流
- WS /api/ws/:sessionId - WebSocket 连接
- GET/POST /api/tools - 工具管理
- GET/PUT /api/config - 配置管理
This commit is contained in:
2025-12-12 10:42:20 +08:00
parent 59dbed926e
commit 5e32375f0e
301 changed files with 3281 additions and 43 deletions
+169
View File
@@ -0,0 +1,169 @@
import * as fs from 'fs';
import * as path from 'path';
import type { AgentConfigFile } from './types.js';
/**
* 配置文件搜索路径
*/
const CONFIG_PATHS = [
'.ai-assist/agents.yaml',
'.ai-assist/agents.yml',
'.ai-assist/agents.json',
'.ai-assist.yaml',
'.ai-assist.yml',
];
/**
* 解析 YAML 内容
*/
async function parseYaml(content: string): Promise<unknown> {
try {
const yaml = await import('js-yaml');
return yaml.load(content);
} catch {
console.warn('解析 YAML 配置文件失败');
return null;
}
}
/**
* 加载用户自定义 Agent 配置
* @param workdir 工作目录
* @returns 配置对象或 null
*/
export async function loadAgentConfig(workdir: string): Promise<AgentConfigFile | null> {
for (const configPath of CONFIG_PATHS) {
const fullPath = path.join(workdir, configPath);
try {
if (!fs.existsSync(fullPath)) {
continue;
}
const content = await fs.promises.readFile(fullPath, 'utf-8');
let config: unknown;
if (configPath.endsWith('.json')) {
config = JSON.parse(content);
} else if (configPath.endsWith('.yaml') || configPath.endsWith('.yml')) {
config = await parseYaml(content);
if (!config) continue;
} else {
continue;
}
// 验证配置格式
if (isValidAgentConfig(config)) {
return config;
} else {
console.warn(`Agent 配置格式无效: ${fullPath}`);
}
} catch (error) {
console.warn(`加载 Agent 配置失败: ${fullPath}`, error);
}
}
return null;
}
/**
* 验证配置格式是否有效
*/
function isValidAgentConfig(config: unknown): config is AgentConfigFile {
if (typeof config !== 'object' || config === null) {
return false;
}
const obj = config as Record<string, unknown>;
// defaults 是可选的
if (obj.defaults !== undefined && typeof obj.defaults !== 'object') {
return false;
}
// agents 是可选的
if (obj.agents !== undefined && typeof obj.agents !== 'object') {
return false;
}
return true;
}
/**
* 保存 Agent 配置到文件
* @param workdir 工作目录
* @param config 配置对象
* @param format 文件格式
*/
export async function saveAgentConfig(
workdir: string,
config: AgentConfigFile,
format: 'json' | 'yaml' = 'json'
): Promise<void> {
const dir = path.join(workdir, '.ai-assist');
// 确保目录存在
if (!fs.existsSync(dir)) {
await fs.promises.mkdir(dir, { recursive: true });
}
const filename = format === 'json' ? 'agents.json' : 'agents.yaml';
const fullPath = path.join(dir, filename);
let content: string;
if (format === 'json') {
content = JSON.stringify(config, null, 2);
} else {
try {
const yaml = await import('js-yaml');
content = yaml.dump(config, { indent: 2, lineWidth: 120 });
} catch {
// 回退到 JSON
content = JSON.stringify(config, null, 2);
console.warn('保存 YAML 失败,已保存为 JSON 格式');
}
}
await fs.promises.writeFile(fullPath, content, 'utf-8');
}
/**
* 获取配置文件模板
*/
export function getConfigTemplate(): AgentConfigFile {
return {
defaults: {
maxSteps: 15,
model: {
temperature: 0.7,
},
permission: {
bash: {
rules: [
{ pattern: 'rm -rf *', action: 'deny' },
{ pattern: 'git push --force*', action: 'deny' },
],
},
},
},
agents: {
'custom-agent': {
description: '自定义 Agent 示例',
mode: 'subagent',
prompt: '你是一个自定义助手。',
tools: {
disabled: ['bash'],
},
permission: {
file: {
read: 'allow',
write: 'ask',
},
},
maxSteps: 10,
},
},
};
}
+322
View File
@@ -0,0 +1,322 @@
import {
generateText,
streamText,
stepCountIs,
type ModelMessage,
type Tool as AITool,
type LanguageModel,
} from 'ai';
import type { Tool, ToolResult, AgentConfig, ContentBlock } from '../types/index.js';
import { buildZodSchema } from '../types/index.js';
import { ToolRegistry } from '../tools/registry.js';
import type {
AgentInfo,
AgentExecutionContext,
AgentExecutionResult,
ImageData,
} from './types.js';
import { checkBashPermission } from './permission-merger.js';
import { getModelFactory } from '../core/providers.js';
/**
* Agent 执行器
* 根据 Agent 配置执行任务,支持工具过滤和权限控制
*/
export class AgentExecutor {
private agentInfo: AgentInfo;
private baseConfig: AgentConfig;
private toolRegistry: ToolRegistry;
private getModel: (model: string) => LanguageModel;
constructor(
agentInfo: AgentInfo,
baseConfig: AgentConfig,
toolRegistry: ToolRegistry
) {
this.agentInfo = agentInfo;
this.baseConfig = baseConfig;
this.toolRegistry = toolRegistry;
// 获取模型工厂
const provider = agentInfo.model?.provider ?? baseConfig.provider;
this.getModel = getModelFactory(provider, {
apiKey: baseConfig.apiKey,
baseUrl: baseConfig.baseUrl,
});
}
/**
* 执行任务
*/
async execute(
prompt: string,
context: AgentExecutionContext
): Promise<AgentExecutionResult> {
const { onStream, onToolCall, onToolResult, images } = context;
// 获取过滤后的工具
const tools = this.getFilteredTools();
const vercelTools = this.buildVercelTools(tools);
// 构建系统提示词
const systemPrompt = this.buildSystemPrompt();
// 获取模型配置
const modelName = this.agentInfo.model?.model ?? this.baseConfig.model;
const maxSteps = this.agentInfo.maxSteps ?? 10;
const maxTokens = this.agentInfo.model?.maxTokens ?? this.baseConfig.maxTokens;
// 构建消息内容(支持图片)
const messageContent = this.buildMessageContent(prompt, images);
// 构建初始消息
const messages: ModelMessage[] = [
{
role: 'user',
content: messageContent,
},
];
let fullResponse = '';
let steps = 0;
try {
if (onStream) {
// 流式模式
const result = streamText({
model: this.getModel(modelName),
system: systemPrompt,
messages,
tools: vercelTools,
maxOutputTokens: maxTokens,
stopWhen: stepCountIs(maxSteps),
onChunk: ({ chunk }) => {
if (chunk.type === 'tool-call') {
steps++;
const toolArgs = 'input' in chunk ? chunk.input : {};
onToolCall?.(chunk.toolName, toolArgs as Record<string, unknown>);
onStream(`\n[调用工具: ${chunk.toolName}]\n`);
} else if (chunk.type === 'tool-result') {
const output = (chunk as { output?: ToolResult }).output;
onToolResult?.(
(chunk as { toolName?: string }).toolName ?? 'unknown',
output
);
if (output && typeof output === 'object') {
if (output.success) {
const displayOutput =
output.output.length > 500
? output.output.substring(0, 500) + '...(截断)'
: output.output;
onStream(`[结果: ${displayOutput}]\n`);
} else {
onStream(`[错误: ${output.error}]\n`);
}
}
}
},
});
for await (const chunk of result.textStream) {
fullResponse += chunk;
onStream(chunk);
}
await result.response;
} else {
// 非流式模式
const result = await generateText({
model: this.getModel(modelName),
system: systemPrompt,
messages,
tools: vercelTools,
maxOutputTokens: maxTokens,
stopWhen: stepCountIs(maxSteps),
});
fullResponse = result.text;
steps = result.steps.length;
}
return {
success: true,
text: fullResponse,
steps,
sessionId: context.parentSessionId ?? 'standalone',
};
} catch (error) {
const errorMessage = error instanceof Error ? error.message : String(error);
return {
success: false,
text: '',
steps,
sessionId: context.parentSessionId ?? 'standalone',
error: errorMessage,
};
}
}
/**
* 获取过滤后的工具列表
*/
private getFilteredTools(): Tool[] {
const allTools = this.toolRegistry.getAllTools();
const toolConfig = this.agentInfo.tools;
// 如果没有工具配置,返回所有工具
if (!toolConfig) {
return allTools;
}
let filteredTools = allTools;
// 如果指定了 enabled,只保留这些工具
if (toolConfig.enabled && toolConfig.enabled.length > 0) {
const enabledSet = new Set(toolConfig.enabled);
filteredTools = filteredTools.filter((t) => enabledSet.has(t.name));
}
// 移除 disabled 的工具
if (toolConfig.disabled && toolConfig.disabled.length > 0) {
const disabledSet = new Set(toolConfig.disabled);
filteredTools = filteredTools.filter((t) => !disabledSet.has(t.name));
}
// 如果禁止嵌套 Task,移除 task 工具
if (toolConfig.noTask) {
filteredTools = filteredTools.filter((t) => t.name !== 'task');
}
return filteredTools;
}
/**
* 构建 Vercel AI SDK 工具格式
*/
private buildVercelTools(tools: Tool[]): Record<string, AITool> {
const vercelTools: Record<string, AITool> = {};
for (const tool of tools) {
const schema = buildZodSchema(tool.parameters);
vercelTools[tool.name] = {
description: tool.description,
inputSchema: schema,
execute: async (params) => {
// 权限检查
const permissionResult = await this.checkToolPermission(
tool.name,
params as Record<string, unknown>
);
if (!permissionResult.allowed) {
return {
success: false,
output: '',
error: `权限拒绝: ${permissionResult.reason}`,
};
}
return tool.execute(params as Record<string, unknown>);
},
} as AITool;
}
return vercelTools;
}
/**
* 检查工具调用权限
*/
private async checkToolPermission(
toolName: string,
params: Record<string, unknown>
): Promise<{ allowed: boolean; reason?: string }> {
const permission = this.agentInfo.permission;
if (!permission) {
return { allowed: true };
}
// Bash 权限检查
if (toolName === 'bash' && permission.bash) {
const command = params.command as string;
if (!command) {
return { allowed: true };
}
const action = checkBashPermission(command, permission.bash);
if (action === 'deny') {
return { allowed: false, reason: `命令被禁止: ${command}` };
}
// ask 在这里视为允许(实际的 ask 逻辑在权限管理器中处理)
}
// 文件写入权限检查
if (['write_file', 'edit_file', 'delete_file'].includes(toolName)) {
const filePermission = permission.file;
if (filePermission) {
const operation = toolName === 'write_file' ? 'write' :
toolName === 'edit_file' ? 'edit' : 'delete';
const action = filePermission[operation];
if (action === 'deny') {
return { allowed: false, reason: `${operation} 操作被禁止` };
}
}
}
// Git 写操作权限检查
const gitWriteTools = ['git_add', 'git_commit', 'git_push', 'git_checkout', 'git_stash'];
if (gitWriteTools.includes(toolName) && permission.git?.write === 'deny') {
return { allowed: false, reason: 'Git 写操作被禁止' };
}
return { allowed: true };
}
/**
* 构建系统提示词
*/
private buildSystemPrompt(): string {
// 如果 Agent 有自定义 prompt,使用它
if (this.agentInfo.prompt) {
return this.agentInfo.prompt;
}
// 否则使用基础配置的 systemPrompt
return this.baseConfig.systemPrompt;
}
/**
* 构建消息内容(支持图片)
*/
private buildMessageContent(
prompt: string,
images?: ImageData[]
): string | ContentBlock[] {
// 如果没有图片,直接返回文本
if (!images || images.length === 0) {
return prompt;
}
// 构建多模态内容
const blocks: ContentBlock[] = [];
// 先添加图片
for (const img of images) {
blocks.push({
type: 'image',
image: img.data,
mimeType: img.mimeType,
});
}
// 再添加文本
if (prompt) {
blocks.push({
type: 'text',
text: prompt,
});
}
return blocks;
}
}
+59
View File
@@ -0,0 +1,59 @@
// Types
export type {
AgentMode,
PermissionAction,
PermissionRule,
AgentBashPermission,
AgentFilePermission,
AgentGitPermission,
AgentPermission,
AgentModelConfig,
AgentToolConfig,
AgentInfo,
AgentConfigFile,
AgentExecutionContext,
AgentExecutionResult,
} from './types.js';
// Registry
export { AgentRegistry, agentRegistry } from './registry.js';
// Executor
export { AgentExecutor } from './executor.js';
// Manager
export {
AgentManager,
getAgentManager,
resetAgentManager,
type BackgroundAgent,
type BackgroundAgentStatus,
} from './manager.js';
// Permission Merger
export {
SYSTEM_DEFAULT_PERMISSION,
mergePermissions,
matchRule,
checkBashPermission,
checkFilePathPermission,
} from './permission-merger.js';
// Config Loader
export {
loadAgentConfig,
saveAgentConfig,
getConfigTemplate,
} from './config-loader.js';
// Presets
export {
presetAgents,
getPresetAgentNames,
isPresetAgent,
generalAgent,
exploreAgent,
codeReviewerAgent,
buildAgent,
planAgent,
} from './presets/index.js';
+249
View File
@@ -0,0 +1,249 @@
import { v4 as uuidv4 } from 'uuid';
import type { AgentConfig } from '../types/index.js';
import type { AgentExecutionContext, AgentInfo } from './types.js';
import { AgentExecutor } from './executor.js';
import { ToolRegistry } from '../tools/registry.js';
/**
* 后台 Agent 状态
*/
export type BackgroundAgentStatus = 'running' | 'completed' | 'failed';
/**
* 后台 Agent 信息
*/
export interface BackgroundAgent {
/** Agent 唯一 ID */
id: string;
/** Agent 类型名称 */
agentName: string;
/** 任务描述 */
description: string;
/** 执行状态 */
status: BackgroundAgentStatus;
/** 任务提示词 */
prompt: string;
/** 开始时间 */
startedAt: Date;
/** 完成时间 */
completedAt?: Date;
/** 执行结果 */
result?: string;
/** 错误信息 */
error?: string;
/** 执行步数 */
steps?: number;
}
/**
* Agent 管理器
* 负责管理后台 Agent 的生命周期
*/
export class AgentManager {
private backgroundAgents = new Map<string, BackgroundAgent>();
private completionCallbacks = new Map<string, Array<() => void>>();
/**
* 启动后台 Agent
*/
async runInBackground(
agentInfo: AgentInfo,
description: string,
prompt: string,
baseConfig: AgentConfig,
toolRegistry: ToolRegistry,
context: AgentExecutionContext
): Promise<string> {
const agentId = uuidv4().substring(0, 8); // 短 ID
// 创建后台 Agent 记录
this.backgroundAgents.set(agentId, {
id: agentId,
agentName: agentInfo.name,
description,
status: 'running',
prompt,
startedAt: new Date(),
});
// 异步执行,不等待结果
this.executeAsync(agentId, agentInfo, prompt, baseConfig, toolRegistry, context);
return agentId;
}
/**
* 异步执行 Agent 任务
*/
private async executeAsync(
agentId: string,
agentInfo: AgentInfo,
prompt: string,
baseConfig: AgentConfig,
toolRegistry: ToolRegistry,
context: AgentExecutionContext
): Promise<void> {
try {
const executor = new AgentExecutor(agentInfo, baseConfig, toolRegistry);
const result = await executor.execute(prompt, {
...context,
onStream: undefined, // 后台运行不使用流式输出
});
// 更新状态为完成
const agent = this.backgroundAgents.get(agentId);
if (agent) {
agent.status = result.success ? 'completed' : 'failed';
agent.completedAt = new Date();
agent.result = result.text;
agent.steps = result.steps;
if (!result.success) {
agent.error = result.error;
}
}
} catch (error) {
// 更新状态为失败
const agent = this.backgroundAgents.get(agentId);
if (agent) {
agent.status = 'failed';
agent.completedAt = new Date();
agent.error = error instanceof Error ? error.message : String(error);
}
}
// 触发等待回调
this.notifyCompletion(agentId);
}
/**
* 获取后台 Agent 状态
*/
getAgent(agentId: string): BackgroundAgent | null {
return this.backgroundAgents.get(agentId) || null;
}
/**
* 获取 Agent 输出(支持阻塞等待)
*/
async getAgentOutput(
agentId: string,
block: boolean = true,
timeoutSeconds: number = 150
): Promise<BackgroundAgent | null> {
const agent = this.backgroundAgents.get(agentId);
if (!agent) {
return null;
}
// 如果已完成或不需要阻塞,直接返回
if (agent.status !== 'running' || !block) {
return agent;
}
// 阻塞等待完成
return this.waitForCompletion(agentId, timeoutSeconds);
}
/**
* 等待 Agent 完成
*/
private waitForCompletion(
agentId: string,
timeoutSeconds: number
): Promise<BackgroundAgent | null> {
return new Promise((resolve) => {
const agent = this.backgroundAgents.get(agentId);
if (!agent || agent.status !== 'running') {
resolve(agent || null);
return;
}
// 设置超时
const timeoutId = setTimeout(() => {
// 移除回调
const callbacks = this.completionCallbacks.get(agentId);
if (callbacks) {
const index = callbacks.indexOf(callback);
if (index > -1) {
callbacks.splice(index, 1);
}
}
resolve(this.backgroundAgents.get(agentId) || null);
}, timeoutSeconds * 1000);
// 注册完成回调
const callback = () => {
clearTimeout(timeoutId);
resolve(this.backgroundAgents.get(agentId) || null);
};
if (!this.completionCallbacks.has(agentId)) {
this.completionCallbacks.set(agentId, []);
}
this.completionCallbacks.get(agentId)!.push(callback);
});
}
/**
* 通知等待者任务已完成
*/
private notifyCompletion(agentId: string): void {
const callbacks = this.completionCallbacks.get(agentId);
if (callbacks) {
for (const callback of callbacks) {
callback();
}
this.completionCallbacks.delete(agentId);
}
}
/**
* 列出所有后台 Agent
*/
listAgents(): BackgroundAgent[] {
return Array.from(this.backgroundAgents.values());
}
/**
* 列出运行中的 Agent
*/
listRunningAgents(): BackgroundAgent[] {
return this.listAgents().filter((a) => a.status === 'running');
}
/**
* 清理已完成的 Agent 记录
*/
cleanup(maxAge: number = 3600000): void {
const now = Date.now();
for (const [id, agent] of this.backgroundAgents) {
if (
agent.status !== 'running' &&
agent.completedAt &&
now - agent.completedAt.getTime() > maxAge
) {
this.backgroundAgents.delete(id);
}
}
}
}
// 单例实例
let agentManager: AgentManager | null = null;
/**
* 获取 AgentManager 单例
*/
export function getAgentManager(): AgentManager {
if (!agentManager) {
agentManager = new AgentManager();
}
return agentManager;
}
/**
* 重置 AgentManager(测试用)
*/
export function resetAgentManager(): void {
agentManager = null;
}
@@ -0,0 +1,221 @@
import type {
AgentPermission,
AgentFilePermission,
AgentBashPermission,
AgentGitPermission,
PermissionAction,
PermissionRule,
} from './types.js';
/**
* 系统默认权限配置
*/
export const SYSTEM_DEFAULT_PERMISSION: AgentPermission = {
file: {
read: 'allow',
write: 'ask',
edit: 'ask',
delete: 'ask',
},
bash: {
enabled: true,
default: 'ask',
rules: [
// 安全命令
{ pattern: 'ls *', action: 'allow' },
{ pattern: 'pwd', action: 'allow' },
{ pattern: 'cat *', action: 'allow' },
{ pattern: 'head *', action: 'allow' },
{ pattern: 'tail *', action: 'allow' },
{ pattern: 'wc *', action: 'allow' },
{ pattern: 'echo *', action: 'allow' },
{ pattern: 'which *', action: 'allow' },
{ pattern: 'type *', action: 'allow' },
// 危险命令
{ pattern: 'rm -rf *', action: 'deny' },
{ pattern: 'rm -fr *', action: 'deny' },
{ pattern: 'sudo *', action: 'deny' },
{ pattern: 'chmod 777 *', action: 'deny' },
],
},
web: 'ask',
git: {
read: 'allow',
write: 'ask',
dangerous: 'deny',
},
};
/**
* 合并单个权限值
* 优先级: agent > global > system
*/
function mergeAction(
system: PermissionAction | undefined,
global: PermissionAction | undefined,
agent: PermissionAction | undefined
): PermissionAction {
return agent ?? global ?? system ?? 'ask';
}
/**
* 合并规则数组
* Agent 规则优先,然后是 global,最后是 system
*/
function mergeRules(
system: PermissionRule[] | undefined,
global: PermissionRule[] | undefined,
agent: PermissionRule[] | undefined
): PermissionRule[] {
return [
...(agent ?? []),
...(global ?? []),
...(system ?? []),
];
}
/**
* 合并文件权限配置
*/
function mergeFilePermission(
system: AgentFilePermission | undefined,
global: AgentFilePermission | undefined,
agent: AgentFilePermission | undefined
): AgentFilePermission {
return {
read: mergeAction(system?.read, global?.read, agent?.read),
write: mergeAction(system?.write, global?.write, agent?.write),
edit: mergeAction(system?.edit, global?.edit, agent?.edit),
delete: mergeAction(system?.delete, global?.delete, agent?.delete),
sensitivePaths: mergeRules(
system?.sensitivePaths,
global?.sensitivePaths,
agent?.sensitivePaths
),
};
}
/**
* 合并 Bash 权限配置
*/
function mergeBashPermission(
system: AgentBashPermission | undefined,
global: AgentBashPermission | undefined,
agent: AgentBashPermission | undefined
): AgentBashPermission {
// 如果 agent 显式禁用,直接返回
if (agent?.enabled === false) {
return { enabled: false };
}
// 如果 global 禁用且 agent 没有覆盖
if (global?.enabled === false && agent?.enabled === undefined) {
return { enabled: false };
}
return {
enabled: agent?.enabled ?? global?.enabled ?? system?.enabled ?? true,
rules: mergeRules(system?.rules, global?.rules, agent?.rules),
default: mergeAction(system?.default, global?.default, agent?.default),
};
}
/**
* 合并 Git 权限配置
*/
function mergeGitPermission(
system: AgentGitPermission | undefined,
global: AgentGitPermission | undefined,
agent: AgentGitPermission | undefined
): AgentGitPermission {
return {
read: mergeAction(system?.read, global?.read, agent?.read),
write: mergeAction(system?.write, global?.write, agent?.write),
dangerous: mergeAction(system?.dangerous, global?.dangerous, agent?.dangerous),
};
}
/**
* 合并完整的权限配置
* 优先级: Agent 配置 > 全局默认 > 系统默认
*/
export function mergePermissions(
systemDefault: AgentPermission,
globalConfig: AgentPermission | undefined,
agentConfig: AgentPermission | undefined
): AgentPermission {
return {
file: mergeFilePermission(
systemDefault.file,
globalConfig?.file,
agentConfig?.file
),
bash: mergeBashPermission(
systemDefault.bash,
globalConfig?.bash,
agentConfig?.bash
),
web: mergeAction(systemDefault.web, globalConfig?.web, agentConfig?.web),
git: mergeGitPermission(
systemDefault.git,
globalConfig?.git,
agentConfig?.git
),
};
}
/**
* 检查命令是否匹配规则
* 支持通配符 * 匹配任意字符
*/
export function matchRule(command: string, pattern: string): boolean {
// 将通配符模式转换为正则表达式
const regexPattern = pattern
.replace(/[.+^${}()|[\]\\]/g, '\\$&') // 转义特殊字符
.replace(/\*/g, '.*') // * 转换为 .*
.replace(/\?/g, '.'); // ? 转换为 .
const regex = new RegExp(`^${regexPattern}$`, 'i');
return regex.test(command);
}
/**
* 根据规则列表检查命令权限
*/
export function checkBashPermission(
command: string,
permission: AgentBashPermission
): PermissionAction {
// 如果禁用,直接拒绝
if (permission.enabled === false) {
return 'deny';
}
// 按顺序检查规则
for (const rule of permission.rules ?? []) {
if (matchRule(command, rule.pattern)) {
return rule.action;
}
}
// 返回默认策略
return permission.default ?? 'ask';
}
/**
* 根据规则列表检查文件路径权限
*/
export function checkFilePathPermission(
filePath: string,
sensitivePaths: PermissionRule[] | undefined
): PermissionAction | null {
if (!sensitivePaths) return null;
for (const rule of sensitivePaths) {
if (matchRule(filePath, rule.pattern)) {
return rule.action;
}
}
return null;
}
+25
View File
@@ -0,0 +1,25 @@
import type { AgentInfo } from '../types.js';
/**
* 构建 Agent
* 主模式,拥有完整权限执行编码任务
*/
export const buildAgent: Omit<AgentInfo, 'name'> = {
description: '构建模式,拥有完整权限执行编码任务',
mode: 'primary',
prompt: `你是一个高效的软件工程师。你的任务是直接执行编码任务。
工作原则:
1. 先理解需求,必要时阅读相关代码
2. 直接执行修改,不要等待确认
3. 修改完成后进行测试验证
4. 保持代码质量,遵循项目规范
你拥有完整的文件系统和 bash 权限,可以:
- 创建、修改、删除文件
- 执行 bash 命令(构建、测试等)
- 使用 git 进行版本控制
始终保持专注和高效,直接解决问题。`,
maxSteps: 30,
};
@@ -0,0 +1,86 @@
import type { AgentInfo } from '../types.js';
/**
* 代码审查 Agent
* 检查代码质量、安全性和最佳实践
*/
export const codeReviewerAgent: Omit<AgentInfo, 'name'> = {
description: '代码审查专家,检查代码质量、安全性和最佳实践',
mode: 'subagent',
prompt: `你是一个资深代码审查专家。你的任务是全面分析代码质量。
审查重点:
1. **代码质量**
- 可读性和命名规范
- 代码结构和组织
- 重复代码和可复用性
- 注释和文档
2. **潜在 Bug**
- 边界条件处理
- 空值/undefined 检查
- 类型安全问题
- 异步错误处理
3. **安全漏洞**
- SQL 注入
- XSS 跨站脚本
- 敏感信息泄露
- 权限控制问题
4. **性能问题**
- 不必要的计算
- 内存泄漏风险
- N+1 查询问题
- 大数据处理
5. **最佳实践**
- 设计模式应用
- SOLID 原则
- 错误处理策略
- 测试覆盖建议
输出格式:
- 问题分为 Critical(严重)、Warning(警告)、Info(建议)三个等级
- 每个问题说明:位置、描述、建议修复方案
- 最后给出总体评分和改进建议`,
tools: {
enabled: [
'read_file',
'list_directory',
'search_files',
'grep_content',
'git_status',
'git_diff',
],
noTask: true,
},
permission: {
file: {
read: 'allow',
write: 'deny',
edit: 'deny',
delete: 'deny',
},
bash: {
enabled: true,
rules: [
{ pattern: 'npm run lint*', action: 'allow' },
{ pattern: 'npm run test*', action: 'allow' },
{ pattern: 'npx eslint *', action: 'allow' },
{ pattern: 'npx tsc --noEmit*', action: 'allow' },
{ pattern: 'tsc --noEmit*', action: 'allow' },
],
default: 'deny',
},
git: {
read: 'allow',
write: 'deny',
dangerous: 'deny',
},
},
model: {
temperature: 0.3, // 低温度,更精确的分析
},
maxSteps: 10,
};
@@ -0,0 +1,49 @@
import type { AgentInfo } from '../types.js';
/**
* 探索 Agent
* 快速探索代码库,搜索文件和代码结构(只读)
*/
export const exploreAgent: Omit<AgentInfo, 'name'> = {
description: '快速探索代码库,搜索文件和代码结构(只读)',
mode: 'subagent',
prompt: `你是一个代码探索专家。你的任务是快速搜索和理解代码库结构。
规则:
- 只做搜索和读取操作,禁止修改任何文件
- 使用 search_files、grep_content、read_file、list_directory 工具来探索代码
- 提供清晰、结构化的分析结果
- 关注代码结构、依赖关系和关键实现
输出格式:
- 使用 Markdown 格式组织信息
- 列出关键文件和它们的作用
- 总结代码结构和设计模式`,
tools: {
enabled: [
'read_file',
'list_directory',
'search_files',
'grep_content',
'get_file_info',
],
noTask: true,
},
permission: {
file: {
read: 'allow',
write: 'deny',
edit: 'deny',
delete: 'deny',
},
bash: {
enabled: false,
},
git: {
read: 'allow',
write: 'deny',
dangerous: 'deny',
},
},
maxSteps: 20,
};
@@ -0,0 +1,14 @@
import type { AgentInfo } from '../types.js';
/**
* 通用 Agent
* 适合复杂的多步骤任务、代码搜索和问题研究
*/
export const generalAgent: Omit<AgentInfo, 'name'> = {
description: '通用 Agent,适合复杂的多步骤任务、代码搜索和问题研究',
mode: 'subagent',
tools: {
noTask: true, // 禁止嵌套调用 Task
},
maxSteps: 15,
};
+35
View File
@@ -0,0 +1,35 @@
import type { AgentInfo } from '../types.js';
import { generalAgent } from './general.js';
import { exploreAgent } from './explore.js';
import { codeReviewerAgent } from './code-reviewer.js';
import { buildAgent } from './build.js';
import { planAgent } from './plan.js';
import { visionAgent } from './vision.js';
/**
* 预设 Agent 集合
*/
export const presetAgents: Record<string, Omit<AgentInfo, 'name'>> = {
general: generalAgent,
explore: exploreAgent,
'code-reviewer': codeReviewerAgent,
build: buildAgent,
plan: planAgent,
vision: visionAgent,
};
/**
* 获取所有预设 Agent 名称
*/
export function getPresetAgentNames(): string[] {
return Object.keys(presetAgents);
}
/**
* 检查是否为预设 Agent
*/
export function isPresetAgent(name: string): boolean {
return name in presetAgents;
}
export { generalAgent, exploreAgent, codeReviewerAgent, buildAgent, planAgent, visionAgent };
+82
View File
@@ -0,0 +1,82 @@
import type { AgentInfo } from '../types.js';
/**
* 计划 Agent
* 主模式,设计实现方案(不执行修改)
*/
export const planAgent: Omit<AgentInfo, 'name'> = {
description: '计划模式,设计实现方案(不执行修改)',
mode: 'primary',
prompt: `你是一个软件架构师。你的任务是设计实现方案,而不是直接执行。
工作流程:
1. 理解需求:分析用户的需求和目标
2. 调研现状:阅读相关代码,了解现有架构
3. 设计方案:提出详细的实现计划
4. 评估风险:识别潜在问题和挑战
规则:
- 可以读取代码来了解现状
- 只输出计划,不要实际执行修改
- 提供具体、可操作的步骤
输出格式:
## 需求分析
[对需求的理解]
## 现状分析
[相关代码的结构和设计]
## 实现方案
### 步骤 1: [标题]
- 目标: ...
- 涉及文件: ...
- 具体修改: ...
### 步骤 2: [标题]
...
## 风险评估
- [风险1]: [应对方案]
- [风险2]: [应对方案]
## 测试计划
- [测试项1]
- [测试项2]`,
tools: {
disabled: [
'write_file',
'edit_file',
'delete_file',
'move_file',
'copy_file',
'create_directory',
'bash',
'git_add',
'git_commit',
'git_push',
'git_checkout',
'git_stash',
],
},
permission: {
file: {
read: 'allow',
write: 'deny',
edit: 'deny',
delete: 'deny',
},
bash: {
enabled: false,
},
git: {
read: 'allow',
write: 'deny',
dangerous: 'deny',
},
},
model: {
temperature: 0.5,
},
maxSteps: 15,
};
+52
View File
@@ -0,0 +1,52 @@
import type { AgentInfo } from '../types.js';
/**
* Vision Agent
* 图片理解专家,使用多模态模型分析图片内容
*/
export const visionAgent: Omit<AgentInfo, 'name'> = {
description: '图片理解专家,分析截图、设计稿、架构图等',
mode: 'subagent',
prompt: `你是一个专业的图片分析专家。你的任务是详细描述和分析用户提供的图片内容。
分析要点:
1. **整体概述**:图片的类型(截图、设计稿、图表、照片等)和主要内容
2. **布局结构**:页面/图片的整体布局、区域划分
3. **文字内容**:提取图片中的所有可见文字(完整、准确)
4. **UI 元素**:按钮、输入框、菜单、图标等元素及其状态
5. **视觉细节**:颜色、字体、间距、对齐等设计细节
6. **交互状态**:hover、选中、禁用等状态指示
7. **潜在问题**:如果用户询问问题,指出可能的问题或改进点
输出格式:
- 使用清晰的 Markdown 格式
- 先给出整体概述,再逐一分析细节
- 如果是 UI 截图,按区域从上到下、从左到右描述
- 提取的文字用引号标注
注意事项:
- 描述要准确、具体,避免模糊表述
- 如果某些内容不清晰,明确说明
- 根据用户的问题重点分析相关部分`,
tools: {
enabled: [],
noTask: true,
},
permission: {
file: {
read: 'deny',
write: 'deny',
edit: 'deny',
delete: 'deny',
},
bash: {
enabled: false,
},
git: {
read: 'deny',
write: 'deny',
dangerous: 'deny',
},
},
maxSteps: 1,
};
+161
View File
@@ -0,0 +1,161 @@
import type { AgentInfo, AgentConfigFile, AgentMode } from './types.js';
import { presetAgents } from './presets/index.js';
import { loadAgentConfig } from './config-loader.js';
import { mergePermissions, SYSTEM_DEFAULT_PERMISSION } from './permission-merger.js';
/**
* Agent 注册表
* 管理所有 Agent 的注册、查询和配置合并
*/
export class AgentRegistry {
private agents: Map<string, AgentInfo> = new Map();
private globalConfig: AgentConfigFile['defaults'] | null = null;
private userConfig: AgentConfigFile | null = null;
private initialized = false;
/**
* 初始化 - 加载预设和用户配置
*/
async init(workdir: string): Promise<void> {
if (this.initialized) return;
// 1. 注册预设 Agent
for (const [name, agentConfig] of Object.entries(presetAgents)) {
this.agents.set(name, { ...agentConfig, name });
}
// 2. 加载用户配置
this.userConfig = await loadAgentConfig(workdir);
if (this.userConfig) {
this.globalConfig = this.userConfig.defaults ?? null;
// 注册用户自定义 Agent
if (this.userConfig.agents) {
for (const [name, config] of Object.entries(this.userConfig.agents)) {
this.agents.set(name, { ...config, name });
}
}
}
this.initialized = true;
}
/**
* 获取 Agent(应用权限合并)
*/
get(name: string): AgentInfo | undefined {
const agent = this.agents.get(name);
if (!agent) return undefined;
return this.applyGlobalConfig(agent);
}
/**
* 列出所有 Agent
*/
list(mode?: AgentMode): AgentInfo[] {
return [...this.agents.values()]
.filter((a) => !mode || a.mode === mode || a.mode === 'all')
.map((a) => this.applyGlobalConfig(a));
}
/**
* 列出可作为子 Agent 的 Agent(供 Task 工具使用)
*/
listSubagents(): AgentInfo[] {
return this.list().filter((a) => a.mode !== 'primary');
}
/**
* 列出可作为主交互 Agent 的 Agent(供 /agent 命令使用)
*/
listPrimaryAgents(): AgentInfo[] {
return this.list().filter((a) => a.mode !== 'subagent');
}
/**
* 动态注册 Agent(运行时)
*/
register(agent: AgentInfo): void {
this.agents.set(agent.name, agent);
}
/**
* 移除 Agent
*/
remove(name: string): boolean {
return this.agents.delete(name);
}
/**
* 检查 Agent 是否存在
*/
has(name: string): boolean {
return this.agents.has(name);
}
/**
* 获取 Agent 数量
*/
get size(): number {
return this.agents.size;
}
/**
* 获取所有 Agent 名称
*/
getNames(): string[] {
return [...this.agents.keys()];
}
/**
* 获取全局配置
*/
getGlobalConfig(): AgentConfigFile['defaults'] | null {
return this.globalConfig;
}
/**
* 应用全局配置到 Agent
*/
private applyGlobalConfig(agent: AgentInfo): AgentInfo {
// 合并 maxSteps
const maxSteps = agent.maxSteps ?? this.globalConfig?.maxSteps ?? 10;
// 合并模型配置
const model = {
...this.globalConfig?.model,
...agent.model,
};
// 合并权限配置
const permission = mergePermissions(
SYSTEM_DEFAULT_PERMISSION,
this.globalConfig?.permission,
agent.permission
);
return {
...agent,
maxSteps,
model: Object.keys(model).length > 0 ? model : undefined,
permission,
};
}
/**
* 生成 Task 工具的 Agent 描述(用于工具 description)
*/
generateSubagentDescription(): string {
const subagents = this.listSubagents();
if (subagents.length === 0) {
return '当前没有可用的子 Agent';
}
const descriptions = subagents.map((a) => `- ${a.name}: ${a.description}`);
return `可用的 Agent:\n${descriptions.join('\n')}`;
}
}
// 导出单例
export const agentRegistry = new AgentRegistry();
+169
View File
@@ -0,0 +1,169 @@
import type { ProviderType } from '../types/index.js';
import type { PermissionAction, PermissionRule } from '../permission/types.js';
// 重新导出权限类型,方便外部使用
export type { PermissionAction, PermissionRule };
/**
* Agent 模式
* - primary: 主 Agent,用户直接使用
* - subagent: 子 Agent,由 Task 工具调用
* - all: 两种方式都可以
*/
export type AgentMode = 'primary' | 'subagent' | 'all';
/**
* Agent Bash 权限配置
*/
export interface AgentBashPermission {
/** 是否启用 bash */
enabled?: boolean;
/** 命令规则,按顺序匹配(支持通配符如 "git diff*", "rm -rf*") */
rules?: PermissionRule[];
/** 默认策略 */
default?: PermissionAction;
}
/**
* Agent 文件权限配置
*/
export interface AgentFilePermission {
read?: PermissionAction;
write?: PermissionAction;
edit?: PermissionAction;
delete?: PermissionAction;
/** 敏感路径规则 */
sensitivePaths?: PermissionRule[];
}
/**
* Agent Git 权限配置
*/
export interface AgentGitPermission {
/** 读操作(status, diff, log 等) */
read?: PermissionAction;
/** 写操作(add, commit, push 等) */
write?: PermissionAction;
/** 危险操作(force push, reset --hard 等) */
dangerous?: PermissionAction;
}
/**
* Agent 完整权限配置
*/
export interface AgentPermission {
file?: AgentFilePermission;
bash?: AgentBashPermission;
web?: PermissionAction;
git?: AgentGitPermission;
}
/**
* Agent 模型配置
*/
export interface AgentModelConfig {
/** Provider 类型 */
provider?: ProviderType;
/** 模型名称 */
model?: string;
/** 温度参数 */
temperature?: number;
/** Top P 参数 */
topP?: number;
/** 最大输出 tokens */
maxTokens?: number;
}
/**
* Agent 工具配置
*/
export interface AgentToolConfig {
/** 禁用的工具列表 */
disabled?: string[];
/** 启用的工具列表(如果设置,则只启用这些) */
enabled?: string[];
/** 禁止嵌套 Task */
noTask?: boolean;
}
/**
* Agent 定义
*/
export interface AgentInfo {
/** Agent 名称 */
name: string;
/** Agent 描述 */
description: string;
/** Agent 模式 */
mode: AgentMode;
/** 自定义 System Prompt */
prompt?: string;
/** 模型配置 */
model?: AgentModelConfig;
/** 工具配置 */
tools?: AgentToolConfig;
/** 权限配置 */
permission?: AgentPermission;
/** 最大执行步数 */
maxSteps?: number;
}
/**
* Agent 配置文件格式(用户自定义)
*/
export interface AgentConfigFile {
/** 全局默认配置 */
defaults?: {
maxSteps?: number;
model?: AgentModelConfig;
permission?: AgentPermission;
};
/** Agent 定义 */
agents?: Record<string, Omit<AgentInfo, 'name'>>;
}
/**
* 图片数据(用于 Agent 执行上下文)
*/
export interface ImageData {
/** base64 编码的图片数据 */
data: string;
/** MIME 类型 */
mimeType: string;
/** 文件名(可选) */
filename?: string;
}
/**
* Agent 执行上下文
*/
export interface AgentExecutionContext {
/** 父会话 ID */
parentSessionId?: string;
/** 工作目录 */
workdir: string;
/** 图片数据(用于支持多模态输入) */
images?: ImageData[];
/** 回调:输出流 */
onStream?: (text: string) => void;
/** 回调:工具调用 */
onToolCall?: (toolName: string, params: Record<string, unknown>) => void;
/** 回调:工具结果 */
onToolResult?: (toolName: string, result: unknown) => void;
}
/**
* Agent 执行结果
*/
export interface AgentExecutionResult {
/** 是否成功 */
success: boolean;
/** 输出文本 */
text: string;
/** 执行步数 */
steps: number;
/** 会话 ID */
sessionId: string;
/** 错误信息 */
error?: string;
}
+35
View File
@@ -0,0 +1,35 @@
/**
* 检查点系统模块
*
* 提供工作区快照和回滚功能,使用 Shadow Git 架构
* 参考 Cline 的实现
*/
// 检查点管理器
export {
CheckpointManager,
getCheckpointManager,
initCheckpointManager,
resetCheckpointManager,
} from './manager.js';
// Shadow Git
export { ShadowGit, createShadowGit, hashWorkingDir } from './shadow-git.js';
// 类型
export type {
CheckpointMetadata,
CheckpointConfig,
CheckpointTrigger,
FileChange,
FileChangeType,
DiffInfo,
FileDiff,
RollbackOptions,
RollbackResult,
CheckpointEvent,
CheckpointEventType,
CheckpointEventListener,
} from './types.js';
export { DEFAULT_CHECKPOINT_CONFIG } from './types.js';
+613
View File
@@ -0,0 +1,613 @@
/**
* 检查点管理器
* 管理检查点的创建、回滚、清理等操作
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import * as os from 'os';
import { nanoid } from 'nanoid';
import { ShadowGit, createShadowGit } from './shadow-git.js';
import type {
CheckpointMetadata,
CheckpointConfig,
CheckpointTrigger,
RollbackOptions,
RollbackResult,
DiffInfo,
FileDiff,
CheckpointEvent,
CheckpointEventListener,
DEFAULT_CHECKPOINT_CONFIG,
} from './types.js';
/**
* 检查点提交消息前缀
*/
const CHECKPOINT_PREFIX = 'checkpoint:';
/**
* 检查点管理器
*/
export class CheckpointManager {
private shadowGit: ShadowGit;
private config: CheckpointConfig;
private workDir: string;
private checkpointsIndex: Map<string, CheckpointMetadata> = new Map();
private initialized = false;
private lastCheckpointTime = 0;
private eventListeners: Set<CheckpointEventListener> = new Set();
// 防止重复创建检查点的最小间隔 (毫秒)
private static readonly MIN_CHECKPOINT_INTERVAL = 1000;
constructor(workDir: string, config: Partial<CheckpointConfig> = {}) {
this.workDir = path.resolve(workDir);
this.config = {
enabled: true,
autoCheckpoint: {
beforeWrite: true,
beforeEdit: true,
beforeDelete: true,
beforeMove: true,
beforeBash: false,
},
maxCheckpoints: 100,
maxAge: 7 * 24 * 60 * 60 * 1000,
storageDir: path.join(os.homedir(), '.ai-assist', 'checkpoints'),
...config,
};
this.shadowGit = createShadowGit(this.workDir, this.config.storageDir);
}
/**
* 初始化检查点管理器
*/
async initialize(): Promise<void> {
if (this.initialized) return;
if (!this.config.enabled) {
this.initialized = true;
return;
}
// 初始化 Shadow Git
await this.shadowGit.initialize();
// 加载检查点索引
await this.loadCheckpointsIndex();
this.initialized = true;
}
/**
* 加载检查点索引
*/
private async loadCheckpointsIndex(): Promise<void> {
try {
const commits = await this.shadowGit.getCommits(this.config.maxCheckpoints);
for (const commit of commits) {
if (commit.message.startsWith(CHECKPOINT_PREFIX)) {
try {
const jsonStr = commit.message.slice(CHECKPOINT_PREFIX.length);
const metadata = JSON.parse(jsonStr) as CheckpointMetadata;
metadata.commitHash = commit.hash;
this.checkpointsIndex.set(metadata.id, metadata);
} catch {
// 解析失败,跳过
}
}
}
} catch {
// 仓库可能是空的
}
}
/**
* 判断是否应该为指定工具创建检查点
*/
shouldCreateCheckpoint(tool: string): boolean {
if (!this.config.enabled) return false;
const { autoCheckpoint } = this.config;
switch (tool) {
case 'write_file':
return autoCheckpoint.beforeWrite;
case 'edit_file':
return autoCheckpoint.beforeEdit;
case 'delete_file':
return autoCheckpoint.beforeDelete;
case 'move_file':
case 'copy_file':
return autoCheckpoint.beforeMove;
case 'bash':
return autoCheckpoint.beforeBash;
default:
return false;
}
}
/**
* 在工具执行前创建检查点
*/
async beforeToolExecution(
tool: string,
params: Record<string, unknown>
): Promise<string | null> {
if (!this.shouldCreateCheckpoint(tool)) {
return null;
}
// 防止过于频繁的检查点创建
const now = Date.now();
if (now - this.lastCheckpointTime < CheckpointManager.MIN_CHECKPOINT_INTERVAL) {
return null;
}
try {
const checkpoint = await this.createCheckpoint({
trigger: `tool:${tool}` as CheckpointTrigger,
toolCall: { tool, params },
description: this.generateDescription(tool, params),
});
this.lastCheckpointTime = now;
return checkpoint.id;
} catch (error) {
console.warn('Failed to create checkpoint:', error);
return null;
}
}
/**
* 生成检查点描述
*/
private generateDescription(
tool: string,
params: Record<string, unknown>
): string {
switch (tool) {
case 'write_file':
return `Write file: ${params.file_path || params.path}`;
case 'edit_file':
return `Edit file: ${params.file_path || params.path}`;
case 'delete_file':
return `Delete file: ${params.file_path || params.path}`;
case 'move_file':
return `Move: ${params.source} -> ${params.destination}`;
case 'copy_file':
return `Copy: ${params.source} -> ${params.destination}`;
case 'bash':
return `Bash: ${String(params.command).slice(0, 50)}`;
default:
return `Tool: ${tool}`;
}
}
/**
* 创建检查点
*/
async createCheckpoint(options: {
name?: string;
description?: string;
trigger?: CheckpointTrigger;
toolCall?: { tool: string; params: Record<string, unknown> };
}): Promise<CheckpointMetadata> {
await this.initialize();
if (!this.config.enabled) {
throw new Error('Checkpoint system is disabled');
}
const id = nanoid(10);
const timestamp = Date.now();
// 创建元数据
const metadata: CheckpointMetadata = {
id,
name: options.name,
description: options.description,
timestamp,
trigger: options.trigger || 'manual',
toolCall: options.toolCall,
commitHash: '', // 待填充
filesChanged: 0, // 待填充
};
// 获取变更文件数
try {
const diff = await this.shadowGit.getWorkingDirDiff();
metadata.filesChanged = diff.files.length;
} catch {
// 忽略
}
// 创建 commit
const commitMessage = CHECKPOINT_PREFIX + JSON.stringify(metadata);
const commitHash = await this.shadowGit.createCommit(commitMessage);
metadata.commitHash = commitHash;
// 更新索引
this.checkpointsIndex.set(id, metadata);
// 触发事件
this.emitEvent({
type: 'created',
checkpoint: metadata,
timestamp,
});
// 异步清理
this.cleanupAsync();
return metadata;
}
/**
* 创建命名检查点
*/
async createNamedCheckpoint(name: string, description?: string): Promise<CheckpointMetadata> {
return this.createCheckpoint({
name,
description,
trigger: 'manual',
});
}
/**
* 获取所有检查点
*/
async listCheckpoints(): Promise<CheckpointMetadata[]> {
await this.initialize();
return Array.from(this.checkpointsIndex.values()).sort(
(a, b) => b.timestamp - a.timestamp
);
}
/**
* 获取指定检查点
*/
async getCheckpoint(idOrHash: string): Promise<CheckpointMetadata | null> {
await this.initialize();
// 先按 ID 查找
if (this.checkpointsIndex.has(idOrHash)) {
return this.checkpointsIndex.get(idOrHash)!;
}
// 再按 commit hash 查找
for (const checkpoint of this.checkpointsIndex.values()) {
if (checkpoint.commitHash.startsWith(idOrHash)) {
return checkpoint;
}
}
return null;
}
/**
* 获取最近的检查点
*/
async getLatestCheckpoint(): Promise<CheckpointMetadata | null> {
const checkpoints = await this.listCheckpoints();
return checkpoints[0] || null;
}
/**
* 获取检查点与当前工作区的差异
*/
async getDiff(checkpointId: string): Promise<DiffInfo> {
await this.initialize();
const checkpoint = await this.getCheckpoint(checkpointId);
if (!checkpoint) {
throw new Error(`Checkpoint not found: ${checkpointId}`);
}
return this.shadowGit.getDiffSummary(checkpoint.commitHash, 'HEAD');
}
/**
* 获取两个检查点之间的差异
*/
async getDiffBetween(fromId: string, toId: string): Promise<DiffInfo> {
await this.initialize();
const fromCheckpoint = await this.getCheckpoint(fromId);
const toCheckpoint = await this.getCheckpoint(toId);
if (!fromCheckpoint) {
throw new Error(`Checkpoint not found: ${fromId}`);
}
if (!toCheckpoint) {
throw new Error(`Checkpoint not found: ${toId}`);
}
return this.shadowGit.getDiffSummary(
fromCheckpoint.commitHash,
toCheckpoint.commitHash
);
}
/**
* 获取文件差异详情
*/
async getFileDiff(checkpointId: string, filePath: string): Promise<FileDiff> {
await this.initialize();
const checkpoint = await this.getCheckpoint(checkpointId);
if (!checkpoint) {
throw new Error(`Checkpoint not found: ${checkpointId}`);
}
const head = await this.shadowGit.getHead();
return this.shadowGit.getFileDiff(checkpoint.commitHash, head, filePath);
}
/**
* 回滚到检查点
*/
async rollback(options: RollbackOptions): Promise<RollbackResult> {
await this.initialize();
const checkpoint = await this.getCheckpoint(options.target);
if (!checkpoint) {
throw new Error(`Checkpoint not found: ${options.target}`);
}
// 获取当前 HEAD 用于可能的撤销
const previousCommit = await this.shadowGit.getHead();
// 预览模式
if (options.dryRun) {
const diff = await this.getDiff(checkpoint.id);
return {
success: true,
restoredFiles: diff.files.map((f) => f.path),
errors: [],
previousCommit,
};
}
const result: RollbackResult = {
success: true,
restoredFiles: [],
errors: [],
previousCommit,
};
try {
if (options.files && options.files.length > 0) {
// 选择性回滚
await this.shadowGit.checkoutFiles(checkpoint.commitHash, options.files);
result.restoredFiles = options.files;
} else {
// 完整回滚
await this.shadowGit.resetHard(checkpoint.commitHash);
// 获取恢复的文件列表
const diff = await this.shadowGit.getDiffSummary(
previousCommit,
checkpoint.commitHash
);
result.restoredFiles = diff.files.map((f) => f.path);
}
// 触发事件
this.emitEvent({
type: 'restored',
checkpoint,
timestamp: Date.now(),
details: {
files: result.restoredFiles,
previousCommit,
},
});
} catch (error) {
result.success = false;
result.errors.push({
file: '*',
error: error instanceof Error ? error.message : String(error),
});
}
return result;
}
/**
* 撤销操作 (回滚到上一个检查点)
*/
async undo(): Promise<RollbackResult> {
const latest = await this.getLatestCheckpoint();
if (!latest) {
throw new Error('No checkpoints available');
}
// 找到倒数第二个检查点
const checkpoints = await this.listCheckpoints();
if (checkpoints.length < 2) {
// 只有一个检查点,回滚到它
return this.rollback({ target: latest.id });
}
// 回滚到倒数第二个检查点
return this.rollback({ target: checkpoints[1].id });
}
/**
* 删除检查点
*/
async deleteCheckpoint(checkpointId: string): Promise<boolean> {
await this.initialize();
if (!this.checkpointsIndex.has(checkpointId)) {
return false;
}
const checkpoint = this.checkpointsIndex.get(checkpointId)!;
this.checkpointsIndex.delete(checkpointId);
// 触发事件
this.emitEvent({
type: 'deleted',
checkpoint,
timestamp: Date.now(),
});
return true;
}
/**
* 异步清理过期检查点
*/
private async cleanupAsync(): Promise<void> {
setTimeout(async () => {
try {
await this.cleanup();
} catch (error) {
console.warn('Checkpoint cleanup failed:', error);
}
}, 100);
}
/**
* 清理过期检查点
*/
async cleanup(): Promise<number> {
await this.initialize();
const checkpoints = await this.listCheckpoints();
const now = Date.now();
let deletedCount = 0;
// 按时间过期清理
for (const checkpoint of checkpoints) {
if (now - checkpoint.timestamp > this.config.maxAge) {
await this.deleteCheckpoint(checkpoint.id);
deletedCount++;
}
}
// 按数量限制清理
const remaining = checkpoints.length - deletedCount;
if (remaining > this.config.maxCheckpoints) {
const toDelete = checkpoints.slice(this.config.maxCheckpoints);
for (const checkpoint of toDelete) {
if (this.checkpointsIndex.has(checkpoint.id)) {
await this.deleteCheckpoint(checkpoint.id);
deletedCount++;
}
}
}
if (deletedCount > 0) {
// 触发清理事件
this.emitEvent({
type: 'cleanup',
timestamp: now,
details: { deletedCount },
});
// 运行 git gc
await this.shadowGit.cleanup(this.config.maxCheckpoints);
}
return deletedCount;
}
/**
* 获取检查点存储统计
*/
async getStats(): Promise<{
count: number;
oldestTimestamp: number | null;
newestTimestamp: number | null;
}> {
const checkpoints = await this.listCheckpoints();
return {
count: checkpoints.length,
oldestTimestamp: checkpoints.length > 0
? checkpoints[checkpoints.length - 1].timestamp
: null,
newestTimestamp: checkpoints.length > 0 ? checkpoints[0].timestamp : null,
};
}
/**
* 添加事件监听器
*/
addEventListener(listener: CheckpointEventListener): void {
this.eventListeners.add(listener);
}
/**
* 移除事件监听器
*/
removeEventListener(listener: CheckpointEventListener): void {
this.eventListeners.delete(listener);
}
/**
* 触发事件
*/
private emitEvent(event: CheckpointEvent): void {
for (const listener of this.eventListeners) {
try {
listener(event);
} catch (error) {
console.warn('Checkpoint event listener error:', error);
}
}
}
/**
* 检查是否启用
*/
isEnabled(): boolean {
return this.config.enabled;
}
/**
* 获取配置
*/
getConfig(): CheckpointConfig {
return { ...this.config };
}
}
// 全局检查点管理器实例
let globalCheckpointManager: CheckpointManager | null = null;
/**
* 获取全局检查点管理器实例
*/
export function getCheckpointManager(): CheckpointManager {
if (!globalCheckpointManager) {
globalCheckpointManager = new CheckpointManager(process.cwd());
}
return globalCheckpointManager;
}
/**
* 初始化全局检查点管理器
*/
export async function initCheckpointManager(
workDir: string,
config?: Partial<CheckpointConfig>
): Promise<CheckpointManager> {
globalCheckpointManager = new CheckpointManager(workDir, config);
await globalCheckpointManager.initialize();
return globalCheckpointManager;
}
/**
* 重置全局检查点管理器 (用于测试)
*/
export function resetCheckpointManager(): void {
globalCheckpointManager = null;
}
+576
View File
@@ -0,0 +1,576 @@
/**
* Shadow Git 存储实现
* 使用隔离的 Git 仓库存储检查点,不影响用户的主仓库
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import { execFile } from 'child_process';
import { promisify } from 'util';
import type { FileChange, DiffInfo, FileDiff } from './types.js';
const execFileAsync = promisify(execFile);
/**
* 计算工作目录哈希
* 使用 31 进制哈希算法,与 Cline 保持一致
*/
export function hashWorkingDir(workingDir: string): string {
let hash = 0;
for (let i = 0; i < workingDir.length; i++) {
hash = ((hash * 31 + workingDir.charCodeAt(i)) >>> 0) % 2147483647;
}
return hash.toString().slice(0, 13).padStart(13, '0');
}
/**
* 需要排除的目录列表
*/
const EXCLUDED_DIRS = [
'node_modules',
'.git',
'dist',
'build',
'.next',
'__pycache__',
'.pytest_cache',
'coverage',
'.nyc_output',
'.ai-assist',
];
/**
* Shadow Git 管理器
*/
export class ShadowGit {
private workDir: string;
private shadowGitDir: string;
private initialized = false;
private cwdHash: string;
constructor(workDir: string, storageBaseDir: string) {
this.workDir = path.resolve(workDir);
this.cwdHash = hashWorkingDir(this.workDir);
this.shadowGitDir = path.join(storageBaseDir, this.cwdHash);
}
/**
* 获取工作目录哈希
*/
getCwdHash(): string {
return this.cwdHash;
}
/**
* 获取 Shadow Git 目录
*/
getShadowGitDir(): string {
return this.shadowGitDir;
}
/**
* 初始化 Shadow Git 仓库
*/
async initialize(): Promise<void> {
if (this.initialized) return;
const gitDir = path.join(this.shadowGitDir, '.git');
try {
await fs.access(gitDir);
// 已存在,验证配置
await this.verifyConfig();
} catch {
// 不存在,创建新仓库
await this.createRepository();
}
this.initialized = true;
}
/**
* 创建新的 Shadow Git 仓库
*/
private async createRepository(): Promise<void> {
// 创建目录
await fs.mkdir(this.shadowGitDir, { recursive: true });
// 初始化 Git 仓库
await this.git(['init']);
// 配置用户信息
await this.git(['config', 'user.name', 'AI Assistant Checkpoint']);
await this.git(['config', 'user.email', 'checkpoint@ai-assist.local']);
// 配置工作目录
await this.git(['config', 'core.worktree', this.workDir]);
// 禁用 GPG 签名
await this.git(['config', 'commit.gpgsign', 'false']);
// 创建 .gitignore
const gitignoreContent = EXCLUDED_DIRS.map((d) => `${d}/`).join('\n') + '\n';
await fs.writeFile(
path.join(this.shadowGitDir, '.gitignore'),
gitignoreContent
);
// 创建初始提交
await this.git(['add', '.gitignore']);
await this.git([
'commit',
'--allow-empty',
'-m',
'Initial checkpoint repository',
]);
}
/**
* 验证现有配置
*/
private async verifyConfig(): Promise<void> {
try {
const { stdout } = await this.git(['config', 'core.worktree']);
const configuredWorkDir = stdout.trim();
if (configuredWorkDir !== this.workDir) {
// 更新工作目录配置
await this.git(['config', 'core.worktree', this.workDir]);
}
} catch {
// 配置不存在,添加
await this.git(['config', 'core.worktree', this.workDir]);
}
}
/**
* 执行 Git 命令
*/
private async git(
args: string[],
options: { cwd?: string } = {}
): Promise<{ stdout: string; stderr: string }> {
const cwd = options.cwd || this.shadowGitDir;
const gitDir = path.join(this.shadowGitDir, '.git');
try {
const result = await execFileAsync(
'git',
['--git-dir', gitDir, '--work-tree', this.workDir, ...args],
{
cwd,
maxBuffer: 50 * 1024 * 1024, // 50MB
env: {
...process.env,
GIT_TERMINAL_PROMPT: '0',
},
}
);
return result;
} catch (error: any) {
// 某些 git 命令失败是正常的 (如空提交)
if (error.stdout !== undefined) {
return { stdout: error.stdout || '', stderr: error.stderr || '' };
}
throw error;
}
}
/**
* 创建检查点提交
*/
async createCommit(message: string): Promise<string> {
await this.initialize();
// 暂时禁用嵌套 .git 目录
const nestedGitDirs = await this.findNestedGitDirs();
await this.renameNestedGitDirs(nestedGitDirs, true);
try {
// 添加所有文件
await this.git(['add', '.', '--ignore-errors']);
// 创建提交
await this.git([
'commit',
'--allow-empty',
'--no-verify',
'-m',
message,
]);
// 获取 commit hash
const { stdout } = await this.git(['rev-parse', 'HEAD']);
return stdout.trim();
} finally {
// 恢复嵌套 .git 目录
await this.renameNestedGitDirs(nestedGitDirs, false);
}
}
/**
* 查找嵌套的 .git 目录
*/
private async findNestedGitDirs(): Promise<string[]> {
const nestedDirs: string[] = [];
const walk = async (dir: string, depth = 0): Promise<void> => {
if (depth > 5) return; // 限制深度
try {
const entries = await fs.readdir(dir, { withFileTypes: true });
for (const entry of entries) {
if (!entry.isDirectory()) continue;
const fullPath = path.join(dir, entry.name);
// 跳过排除的目录
if (EXCLUDED_DIRS.includes(entry.name)) continue;
if (entry.name === '.git') {
nestedDirs.push(fullPath);
} else if (!entry.name.startsWith('.')) {
await walk(fullPath, depth + 1);
}
}
} catch {
// 忽略无法访问的目录
}
};
await walk(this.workDir);
return nestedDirs;
}
/**
* 重命名嵌套 .git 目录
*/
private async renameNestedGitDirs(
dirs: string[],
disable: boolean
): Promise<void> {
for (const dir of dirs) {
const disabledName = dir + '_disabled';
try {
if (disable) {
await fs.rename(dir, disabledName);
} else {
await fs.rename(disabledName, dir);
}
} catch {
// 忽略错误
}
}
}
/**
* 获取当前 HEAD commit hash
*/
async getHead(): Promise<string> {
await this.initialize();
const { stdout } = await this.git(['rev-parse', 'HEAD']);
return stdout.trim();
}
/**
* 重置到指定 commit
*/
async resetHard(commitHash: string): Promise<void> {
await this.initialize();
// 暂时禁用嵌套 .git 目录
const nestedGitDirs = await this.findNestedGitDirs();
await this.renameNestedGitDirs(nestedGitDirs, true);
try {
await this.git(['reset', '--hard', commitHash]);
} finally {
await this.renameNestedGitDirs(nestedGitDirs, false);
}
}
/**
* 获取 commit 列表
*/
async getCommits(limit = 100): Promise<
Array<{
hash: string;
message: string;
timestamp: number;
}>
> {
await this.initialize();
const { stdout } = await this.git([
'log',
`--max-count=${limit}`,
'--format=%H|%s|%ct',
]);
if (!stdout.trim()) return [];
return stdout
.trim()
.split('\n')
.map((line) => {
const [hash, message, timestamp] = line.split('|');
return {
hash,
message,
timestamp: parseInt(timestamp, 10) * 1000,
};
});
}
/**
* 获取两个 commit 之间的差异摘要
*/
async getDiffSummary(fromCommit: string, toCommit = 'HEAD'): Promise<DiffInfo> {
await this.initialize();
const { stdout } = await this.git([
'diff',
'--stat',
'--numstat',
fromCommit,
toCommit,
]);
const files: FileChange[] = [];
let totalInsertions = 0;
let totalDeletions = 0;
// 解析 numstat 输出
const lines = stdout.trim().split('\n');
for (const line of lines) {
const match = line.match(/^(\d+|-)\t(\d+|-)\t(.+)$/);
if (match) {
const insertions = match[1] === '-' ? 0 : parseInt(match[1], 10);
const deletions = match[2] === '-' ? 0 : parseInt(match[2], 10);
const filePath = match[3];
// 检测重命名
const renameMatch = filePath.match(/^(.+)\{(.+) => (.+)\}(.*)$/);
if (renameMatch) {
const prefix = renameMatch[1];
const oldName = renameMatch[2];
const newName = renameMatch[3];
const suffix = renameMatch[4];
files.push({
path: prefix + newName + suffix,
oldPath: prefix + oldName + suffix,
type: 'renamed',
insertions,
deletions,
});
} else {
// 获取文件状态
const type = await this.getFileChangeType(fromCommit, toCommit, filePath);
files.push({
path: filePath,
type,
insertions,
deletions,
});
}
totalInsertions += insertions;
totalDeletions += deletions;
}
}
return {
from: fromCommit,
to: toCommit,
files,
totalInsertions,
totalDeletions,
};
}
/**
* 获取文件变更类型
*/
private async getFileChangeType(
fromCommit: string,
toCommit: string,
filePath: string
): Promise<FileChange['type']> {
try {
const { stdout } = await this.git([
'diff',
'--name-status',
fromCommit,
toCommit,
'--',
filePath,
]);
const status = stdout.trim().charAt(0);
switch (status) {
case 'A':
return 'added';
case 'D':
return 'deleted';
case 'R':
return 'renamed';
default:
return 'modified';
}
} catch {
return 'modified';
}
}
/**
* 获取文件内容差异
*/
async getFileDiff(
fromCommit: string,
toCommit: string,
filePath: string
): Promise<FileDiff> {
await this.initialize();
const type = await this.getFileChangeType(fromCommit, toCommit, filePath);
let oldContent: string | undefined;
let newContent: string | undefined;
let patch: string | undefined;
try {
if (type !== 'added') {
const { stdout } = await this.git(['show', `${fromCommit}:${filePath}`]);
oldContent = stdout;
}
} catch {
// 文件在旧 commit 中不存在
}
try {
if (type !== 'deleted') {
const { stdout } = await this.git(['show', `${toCommit}:${filePath}`]);
newContent = stdout;
}
} catch {
// 文件在新 commit 中不存在
}
try {
const { stdout } = await this.git([
'diff',
fromCommit,
toCommit,
'--',
filePath,
]);
patch = stdout;
} catch {
// 无法生成 diff
}
return {
path: filePath,
type,
oldContent,
newContent,
patch,
};
}
/**
* 检出指定 commit 的特定文件
*/
async checkoutFiles(commitHash: string, files: string[]): Promise<void> {
await this.initialize();
for (const file of files) {
try {
await this.git(['checkout', commitHash, '--', file]);
} catch (error) {
// 文件可能不存在于该 commit
console.warn(`Failed to checkout ${file} from ${commitHash}:`, error);
}
}
}
/**
* 获取指定 commit 中文件的内容
*/
async getFileContent(commitHash: string, filePath: string): Promise<string | null> {
await this.initialize();
try {
const { stdout } = await this.git(['show', `${commitHash}:${filePath}`]);
return stdout;
} catch {
return null;
}
}
/**
* 清理旧的 commit(保留最近 N 个)
*/
async cleanup(keepCount: number): Promise<number> {
await this.initialize();
const commits = await this.getCommits(keepCount + 100);
if (commits.length <= keepCount) {
return 0;
}
// 使用 git gc 清理
try {
await this.git(['gc', '--aggressive', '--prune=now']);
} catch {
// gc 可能失败,忽略
}
return commits.length - keepCount;
}
/**
* 检查是否有未提交的变更
*/
async hasChanges(): Promise<boolean> {
await this.initialize();
const { stdout } = await this.git(['status', '--porcelain']);
return stdout.trim().length > 0;
}
/**
* 获取工作目录与 HEAD 的差异
*/
async getWorkingDirDiff(): Promise<DiffInfo> {
await this.initialize();
// 先添加所有文件到暂存区以检测新文件
const nestedGitDirs = await this.findNestedGitDirs();
await this.renameNestedGitDirs(nestedGitDirs, true);
try {
await this.git(['add', '.', '--ignore-errors']);
const result = await this.getDiffSummary('HEAD', '--staged');
// 重置暂存区
await this.git(['reset', 'HEAD']);
return result;
} finally {
await this.renameNestedGitDirs(nestedGitDirs, false);
}
}
}
/**
* 创建 Shadow Git 实例
*/
export function createShadowGit(
workDir: string,
storageBaseDir: string
): ShadowGit {
return new ShadowGit(workDir, storageBaseDir);
}
+191
View File
@@ -0,0 +1,191 @@
/**
* 检查点系统类型定义
* 基于 Cline 的 Shadow Git 架构
*/
/**
* 检查点触发类型
*/
export type CheckpointTrigger =
| 'auto' // 自动创建
| 'manual' // 用户手动
| 'tool:write_file' // 写文件前
| 'tool:edit_file' // 编辑文件前
| 'tool:delete_file' // 删除文件前
| 'tool:move_file' // 移动文件前
| 'tool:copy_file' // 复制文件前
| 'tool:bash' // bash 命令前
| 'task_start' // 任务开始
| 'task_complete'; // 任务完成
/**
* 检查点元数据
*/
export interface CheckpointMetadata {
/** 唯一标识 */
id: string;
/** 用户可读名称 */
name?: string;
/** 描述信息 */
description?: string;
/** 创建时间戳 */
timestamp: number;
/** 触发类型 */
trigger: CheckpointTrigger;
/** 关联的工具调用 */
toolCall?: {
tool: string;
params: Record<string, unknown>;
};
/** Git commit hash */
commitHash: string;
/** 受影响的文件数 */
filesChanged: number;
}
/**
* 检查点配置
*/
export interface CheckpointConfig {
/** 是否启用检查点系统 */
enabled: boolean;
/** 自动检查点配置 */
autoCheckpoint: {
/** 写文件前创建检查点 */
beforeWrite: boolean;
/** 编辑文件前创建检查点 */
beforeEdit: boolean;
/** 删除文件前创建检查点 */
beforeDelete: boolean;
/** 移动/复制文件前创建检查点 */
beforeMove: boolean;
/** bash 命令前创建检查点 */
beforeBash: boolean;
};
/** 最大保留检查点数量 */
maxCheckpoints: number;
/** 检查点最大保留时间 (毫秒) */
maxAge: number;
/** Shadow Git 存储目录 */
storageDir: string;
}
/**
* 默认配置
*/
export const DEFAULT_CHECKPOINT_CONFIG: CheckpointConfig = {
enabled: true,
autoCheckpoint: {
beforeWrite: true,
beforeEdit: true,
beforeDelete: true,
beforeMove: true,
beforeBash: false,
},
maxCheckpoints: 100,
maxAge: 7 * 24 * 60 * 60 * 1000, // 7 天
storageDir: '.ai-assist/checkpoints',
};
/**
* 文件变更类型
*/
export type FileChangeType = 'added' | 'modified' | 'deleted' | 'renamed';
/**
* 文件变更信息
*/
export interface FileChange {
/** 文件路径 */
path: string;
/** 变更类型 */
type: FileChangeType;
/** 旧路径 (重命名时) */
oldPath?: string;
/** 添加的行数 */
insertions?: number;
/** 删除的行数 */
deletions?: number;
}
/**
* 差异信息
*/
export interface DiffInfo {
/** 源检查点/commit */
from: string;
/** 目标检查点/commit (HEAD 表示当前工作区) */
to: string;
/** 变更的文件列表 */
files: FileChange[];
/** 总添加行数 */
totalInsertions: number;
/** 总删除行数 */
totalDeletions: number;
}
/**
* 文件内容差异
*/
export interface FileDiff {
/** 文件路径 */
path: string;
/** 变更类型 */
type: FileChangeType;
/** 旧内容 */
oldContent?: string;
/** 新内容 */
newContent?: string;
/** 差异补丁 (unified diff 格式) */
patch?: string;
}
/**
* 回滚选项
*/
export interface RollbackOptions {
/** 检查点 ID 或 commit hash */
target: string;
/** 只回滚指定文件 */
files?: string[];
/** 预览模式 (不实际执行) */
dryRun?: boolean;
}
/**
* 回滚结果
*/
export interface RollbackResult {
/** 是否成功 */
success: boolean;
/** 恢复的文件列表 */
restoredFiles: string[];
/** 错误列表 */
errors: Array<{ file: string; error: string }>;
/** 回滚前的 commit hash (用于撤销回滚) */
previousCommit?: string;
}
/**
* 检查点事件类型
*/
export type CheckpointEventType =
| 'created' // 检查点已创建
| 'restored' // 已回滚到检查点
| 'deleted' // 检查点已删除
| 'cleanup'; // 清理过期检查点
/**
* 检查点事件
*/
export interface CheckpointEvent {
type: CheckpointEventType;
checkpoint?: CheckpointMetadata;
timestamp: number;
details?: Record<string, unknown>;
}
/**
* 检查点事件监听器
*/
export type CheckpointEventListener = (event: CheckpointEvent) => void;
+197
View File
@@ -0,0 +1,197 @@
/**
* 内置 Commands
*
* 提供一些常用的预定义 Commands
*/
import type { Command } from '../types.js';
/**
* /init - 初始化项目配置
*/
export const initCommand: Command = {
name: 'init',
description: '分析代码库并创建 AGENTS.md 配置文件',
template: `Please analyze this codebase and create an AGENTS.md file containing:
1. **Build/lint/test commands** - Document how to build, lint, and test the project
2. **Code style guidelines** - Document the coding conventions used
3. **Project structure** - Overview of the directory structure
4. **Key dependencies** - Important libraries and frameworks used
Additional context: $ARGUMENTS
Start by exploring the project structure and package configuration files.`,
agent: 'explore',
subtask: false,
source: 'builtin',
};
/**
* /review - 代码审查
*/
export const reviewCommand: Command = {
name: 'review',
description: '审查代码变更',
template: `You are a code reviewer. Your job is to review code changes.
Input: $ARGUMENTS
## Determining What to Review
Based on the input, determine which type of review to perform:
1. **No arguments**: Review uncommitted changes using \`git diff\`
2. **Commit hash**: Review that specific commit using \`git show $ARGUMENTS\`
3. **Branch name**: Compare to specified branch using \`git diff $ARGUMENTS...HEAD\`
4. **PR URL**: Review the pull request (if gh CLI is available)
## What to Look For
- **Bugs**: Logic errors, edge cases, null/undefined handling, security issues
- **Structure**: Does it follow existing patterns? Is it maintainable?
- **Performance**: Only flag if obviously problematic
- **Tests**: Are there adequate tests for the changes?
## Output Format
Provide a structured review with:
1. Summary of changes
2. Issues found (categorized by severity)
3. Suggestions for improvement
4. Positive observations`,
agent: 'code-reviewer',
subtask: false,
source: 'builtin',
};
/**
* /test - 运行并修复测试
*/
export const testCommand: Command = {
name: 'test',
description: '运行测试并修复失败的测试',
template: `Run the test suite and analyze the results.
Focus on: $ARGUMENTS
## Instructions
1. First, run the test command for this project
2. If tests fail, analyze the failures
3. Identify the root cause of each failure
4. Fix the failing tests or the code they're testing
5. Re-run tests to verify fixes
If no specific focus is provided, run all tests.`,
agent: 'general',
subtask: false,
source: 'builtin',
};
/**
* /fix - 修复问题
*/
export const fixCommand: Command = {
name: 'fix',
description: '修复指定的问题或错误',
template: `Please fix the following issue:
$ARGUMENTS
## Instructions
1. Understand the problem described
2. Locate the relevant code
3. Analyze the root cause
4. Implement a fix
5. Verify the fix works
6. Check for any side effects`,
agent: 'general',
subtask: false,
source: 'builtin',
};
/**
* /explain - 解释代码
*/
export const explainCommand: Command = {
name: 'explain',
description: '解释代码或概念',
template: `Please explain the following:
$ARGUMENTS
Provide a clear, structured explanation that includes:
1. Overview - What it does at a high level
2. How it works - Step by step breakdown
3. Key concepts - Important patterns or techniques used
4. Examples - Practical usage examples if applicable`,
agent: 'general',
subtask: false,
source: 'builtin',
};
/**
* /commit - 生成 commit 消息
*/
export const commitCommand: Command = {
name: 'commit',
description: '根据变更生成 Git commit 消息',
template: `Generate a Git commit message for the current changes.
Additional context: $ARGUMENTS
## Instructions
1. Run \`git diff --staged\` to see staged changes (or \`git diff\` if nothing staged)
2. Analyze the changes
3. Generate a commit message following Conventional Commits format:
- feat: New feature
- fix: Bug fix
- docs: Documentation
- style: Formatting
- refactor: Code restructuring
- test: Adding tests
- chore: Maintenance
4. Format:
- First line: type(scope): short description (50 chars max)
- Blank line
- Body: Detailed explanation if needed
5. Present the commit message for review`,
agent: 'general',
subtask: false,
source: 'builtin',
};
/**
* /help - 显示帮助
*/
export const helpCommand: Command = {
name: 'help',
description: '显示可用的命令和帮助信息',
template: `Show help information about available commands.
Topic: $ARGUMENTS
If a specific command is mentioned, provide detailed help for that command.
Otherwise, list all available commands with their descriptions.`,
agent: 'general',
subtask: false,
source: 'builtin',
};
/**
* 所有内置 Commands
*/
export const builtinCommands: Command[] = [
initCommand,
reviewCommand,
testCommand,
fixCommand,
explainCommand,
commitCommand,
helpCommand,
];
+284
View File
@@ -0,0 +1,284 @@
/**
* Command 执行器
*
* 负责解析和执行 Command:
* - 参数替换($ARGUMENTS, $1, $2, ...)
* - 文件引用(@filepath)
* - Shell 命令执行(!`command`)
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import { exec } from 'child_process';
import { promisify } from 'util';
import type { Command, CommandInput, CommandExecutionResult } from './types.js';
import { getCommandRegistry } from './registry.js';
const execAsync = promisify(exec);
/**
* Command 执行器
*/
export class CommandExecutor {
private workdir: string;
constructor(workdir: string = process.cwd()) {
this.workdir = workdir;
}
/**
* 解析用户输入的命令字符串
* 例如: "/review main..feature" → { command: "review", arguments: "main..feature", args: ["main..feature"] }
*/
parseInput(input: string): CommandInput | null {
// 移除开头的 /
const trimmed = input.trim();
if (!trimmed.startsWith('/')) {
return null;
}
const withoutSlash = trimmed.slice(1);
const spaceIndex = withoutSlash.indexOf(' ');
let commandName: string;
let argumentsStr: string;
if (spaceIndex === -1) {
commandName = withoutSlash;
argumentsStr = '';
} else {
commandName = withoutSlash.slice(0, spaceIndex);
argumentsStr = withoutSlash.slice(spaceIndex + 1).trim();
}
// 解析参数数组
const args = argumentsStr ? this.parseArgs(argumentsStr) : [];
return {
command: commandName,
arguments: argumentsStr,
args,
workdir: this.workdir,
};
}
/**
* 解析参数字符串为数组
* 支持引号包裹的参数
*/
private parseArgs(argsStr: string): string[] {
const args: string[] = [];
let current = '';
let inQuote = false;
let quoteChar = '';
for (const char of argsStr) {
if ((char === '"' || char === "'") && !inQuote) {
inQuote = true;
quoteChar = char;
} else if (char === quoteChar && inQuote) {
inQuote = false;
quoteChar = '';
} else if (char === ' ' && !inQuote) {
if (current) {
args.push(current);
current = '';
}
} else {
current += char;
}
}
if (current) {
args.push(current);
}
return args;
}
/**
* 执行 Command
*/
async execute(input: CommandInput): Promise<CommandExecutionResult> {
const registry = getCommandRegistry();
const command = registry.get(input.command);
if (!command) {
// 尝试搜索相似的 Command
const suggestions = registry.search(input.command, 3);
let errorMsg = `Command 不存在: /${input.command}`;
if (suggestions.length > 0) {
errorMsg += '\n\n你可能想要的 Command:\n';
for (const { command: cmd } of suggestions) {
errorMsg += `- /${cmd.name}`;
if (cmd.description) {
errorMsg += `: ${cmd.description}`;
}
errorMsg += '\n';
}
}
return {
success: false,
error: errorMsg,
};
}
try {
// 渲染模板
const prompt = await this.renderTemplate(command.template, input);
return {
success: true,
prompt,
agent: command.agent,
model: command.model,
subtask: command.subtask,
};
} catch (error) {
return {
success: false,
error: `Command 执行失败: ${error instanceof Error ? error.message : String(error)}`,
};
}
}
/**
* 渲染模板
*/
async renderTemplate(
template: string,
input: CommandInput
): Promise<string> {
let result = template;
// 1. 替换位置参数 $1, $2, ...
result = this.replacePositionalArgs(result, input.args);
// 2. 替换 $ARGUMENTS
result = result.replace(/\$ARGUMENTS/g, input.arguments);
// 3. 处理文件引用 @filepath
result = await this.resolveFileReferences(result, input.workdir);
// 4. 执行 Shell 命令 !`command`
result = await this.executeShellCommands(result, input.workdir);
return result;
}
/**
* 替换位置参数
*/
private replacePositionalArgs(template: string, args: string[]): string {
// 找出模板中使用的最大参数索引
const paramRegex = /\$(\d+)/g;
let maxIndex = 0;
let match;
while ((match = paramRegex.exec(template)) !== null) {
const index = parseInt(match[1], 10);
if (index > maxIndex) {
maxIndex = index;
}
}
// 替换参数
let result = template;
for (let i = 1; i <= maxIndex; i++) {
const value = i === maxIndex
? args.slice(i - 1).join(' ') // 最后一个参数获取剩余所有
: args[i - 1] || '';
result = result.replace(new RegExp(`\\$${i}`, 'g'), value);
}
return result;
}
/**
* 解析文件引用
* @filepath → 文件内容
*/
private async resolveFileReferences(
template: string,
workdir: string
): Promise<string> {
const fileRefRegex = /@([^\s]+)/g;
const matches = [...template.matchAll(fileRefRegex)];
if (matches.length === 0) {
return template;
}
let result = template;
for (const match of matches) {
const [fullMatch, filePath] = match;
const absolutePath = path.isAbsolute(filePath)
? filePath
: path.join(workdir, filePath);
try {
const content = await fs.readFile(absolutePath, 'utf-8');
// 替换为带有文件路径标记的内容
const replacement = `\`\`\`${path.extname(filePath).slice(1) || 'txt'}\n// ${filePath}\n${content}\n\`\`\``;
result = result.replace(fullMatch, replacement);
} catch (error) {
// 文件不存在,保留原样或提示
console.warn(`无法读取文件: ${absolutePath}`);
result = result.replace(fullMatch, `[文件不存在: ${filePath}]`);
}
}
return result;
}
/**
* 执行 Shell 命令
* !`command` → 命令输出
*/
private async executeShellCommands(
template: string,
workdir: string
): Promise<string> {
const shellRegex = /!\`([^`]+)\`/g;
const matches = [...template.matchAll(shellRegex)];
if (matches.length === 0) {
return template;
}
let result = template;
for (const match of matches) {
const [fullMatch, command] = match;
try {
const { stdout, stderr } = await execAsync(command, {
cwd: workdir,
timeout: 30000, // 30 秒超时
});
const output = (stdout + stderr).trim();
result = result.replace(fullMatch, output);
} catch (error) {
// 命令执行失败,替换为错误信息
const errorMsg = error instanceof Error ? error.message : String(error);
result = result.replace(fullMatch, `[命令执行失败: ${command}]\n${errorMsg}`);
}
}
return result;
}
}
/**
* 创建 Command 执行器
*/
export function createCommandExecutor(
workdir: string = process.cwd()
): CommandExecutor {
return new CommandExecutor(workdir);
}
+30
View File
@@ -0,0 +1,30 @@
/**
* Commands 模块
*
* 提供 Command 系统的所有功能导出
*/
// 类型
export type {
Command,
CommandInput,
CommandExecutionResult,
CommandSearchResult,
CommandFrontmatter,
} from './types.js';
// 加载器
export { CommandLoader, commandLoader } from './loader.js';
// 注册表
export {
CommandRegistry,
getCommandRegistry,
resetCommandRegistry,
} from './registry.js';
// 执行器
export { CommandExecutor, createCommandExecutor } from './executor.js';
// 内置 Commands
export { builtinCommands } from './builtin/index.js';
+172
View File
@@ -0,0 +1,172 @@
/**
* Command 加载器
*
* 负责从文件系统加载 Markdown 格式的 Command 定义。
* 支持从以下位置加载:
* 1. 内置 Commands(代码中定义)
* 2. 用户 Commands(~/.config/ai-terminal/commands/)
* 3. 项目 Commands(./.ai-terminal/commands/)
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import * as yaml from 'yaml';
import type { Command, CommandFrontmatter } from './types.js';
/**
* Command 加载器
*/
export class CommandLoader {
/**
* 从目录加载所有 Commands
*/
async loadFromDirectory(
dir: string,
source: 'user' | 'project'
): Promise<Command[]> {
const commands: Command[] = [];
try {
const exists = await fs
.access(dir)
.then(() => true)
.catch(() => false);
if (!exists) {
return commands;
}
await this.scanDirectory(dir, dir, source, commands);
} catch (error) {
console.warn(`读取 Commands 目录失败: ${dir}`, error);
}
return commands;
}
/**
* 递归扫描目录
*/
private async scanDirectory(
baseDir: string,
currentDir: string,
source: 'user' | 'project',
commands: Command[]
): Promise<void> {
const entries = await fs.readdir(currentDir, { withFileTypes: true });
for (const entry of entries) {
const fullPath = path.join(currentDir, entry.name);
if (entry.isFile() && entry.name.endsWith('.md')) {
try {
const command = await this.loadFromFile(fullPath, baseDir, source);
if (command) {
commands.push(command);
}
} catch (error) {
console.warn(`加载 Command 文件失败: ${fullPath}`, error);
}
} else if (entry.isDirectory() && !entry.name.startsWith('.')) {
// 递归加载子目录(支持嵌套路径如 deploy/staging)
await this.scanDirectory(baseDir, fullPath, source, commands);
}
}
}
/**
* 从单个 Markdown 文件加载 Command
*/
async loadFromFile(
filePath: string,
baseDir: string,
source: 'user' | 'project'
): Promise<Command | null> {
const content = await fs.readFile(filePath, 'utf-8');
// 从文件路径推断命令名称
// baseDir/deploy/staging.md → deploy/staging
const relativePath = path.relative(baseDir, filePath);
const name = relativePath.slice(0, -3); // 移除 .md 后缀
return this.parseMarkdownCommand(content, name, source, filePath);
}
/**
* 解析 Markdown 格式的 Command
*
* 格式示例:
* ```markdown
* ---
* description: 代码审查
* agent: explore
* model: sonnet
* subtask: true
* ---
*
* You are a code reviewer...
*
* Input: $ARGUMENTS
* ```
*/
parseMarkdownCommand(
content: string,
name: string,
source: 'user' | 'project' | 'builtin',
sourcePath?: string
): Command | null {
// 解析 frontmatter
const frontmatterMatch = content.match(/^---\n([\s\S]*?)\n---\n([\s\S]*)$/);
let frontmatter: CommandFrontmatter = {};
let template: string;
if (frontmatterMatch) {
const [, frontmatterStr, bodyContent] = frontmatterMatch;
try {
frontmatter = yaml.parse(frontmatterStr) as CommandFrontmatter;
} catch (error) {
console.warn(`解析 Command frontmatter 失败: ${sourcePath}`, error);
}
template = bodyContent.trim();
} else {
// 没有 frontmatter,整个内容作为模板
template = content.trim();
}
if (!template) {
return null;
}
return {
name,
description: frontmatter.description,
template,
agent: frontmatter.agent,
model: frontmatter.model,
subtask: frontmatter.subtask,
source,
sourcePath,
};
}
/**
* 获取用户 Commands 目录
*/
getUserCommandsDir(): string {
const home = process.env.HOME || process.env.USERPROFILE || '';
return path.join(home, '.config', 'ai-terminal', 'commands');
}
/**
* 获取项目 Commands 目录
*/
getProjectCommandsDir(workdir: string = process.cwd()): string {
return path.join(workdir, '.ai-terminal', 'commands');
}
}
/**
* 全局 Command 加载器实例
*/
export const commandLoader = new CommandLoader();
+199
View File
@@ -0,0 +1,199 @@
/**
* Command 注册表
*
* 管理所有可用的 Commands,支持:
* - 注册/注销 Commands
* - 按名称查询
* - 搜索 Commands
*/
import type { Command, CommandSearchResult } from './types.js';
import { commandLoader } from './loader.js';
import { builtinCommands } from './builtin/index.js';
/**
* Command 注册表
*/
export class CommandRegistry {
private commands = new Map<string, Command>();
private initialized = false;
/**
* 初始化注册表
*/
async initialize(workdir: string = process.cwd()): Promise<void> {
if (this.initialized) {
return;
}
// 1. 注册内置 Commands
for (const command of builtinCommands) {
this.register(command);
}
// 2. 加载用户 Commands
const userDir = commandLoader.getUserCommandsDir();
const userCommands = await commandLoader.loadFromDirectory(userDir, 'user');
for (const command of userCommands) {
this.register(command);
}
// 3. 加载项目 Commands
const projectDir = commandLoader.getProjectCommandsDir(workdir);
const projectCommands = await commandLoader.loadFromDirectory(
projectDir,
'project'
);
for (const command of projectCommands) {
this.register(command);
}
this.initialized = true;
}
/**
* 注册 Command
*/
register(command: Command): void {
// 项目 Commands 优先级最高,可以覆盖同名的内置/用户 Commands
const existing = this.commands.get(command.name);
if (existing) {
// 优先级: project > user > builtin
const priority = { project: 3, user: 2, builtin: 1 };
if (priority[command.source] < priority[existing.source]) {
return; // 不覆盖更高优先级的 Command
}
}
this.commands.set(command.name, command);
}
/**
* 注销 Command
*/
unregister(name: string): boolean {
return this.commands.delete(name);
}
/**
* 获取 Command
*/
get(name: string): Command | undefined {
return this.commands.get(name);
}
/**
* 检查 Command 是否存在
*/
has(name: string): boolean {
return this.commands.has(name);
}
/**
* 获取所有 Commands
*/
getAll(): Command[] {
return Array.from(this.commands.values());
}
/**
* 搜索 Commands
*/
search(query: string, limit: number = 10): CommandSearchResult[] {
const queryLower = query.toLowerCase();
const results: CommandSearchResult[] = [];
for (const command of this.commands.values()) {
let score = 0;
// 精确名称匹配
if (command.name.toLowerCase() === queryLower) {
score = 100;
}
// 名称前缀匹配
else if (command.name.toLowerCase().startsWith(queryLower)) {
score = 80;
}
// 名称包含匹配
else if (command.name.toLowerCase().includes(queryLower)) {
score = 60;
}
// 描述匹配
else if (command.description?.toLowerCase().includes(queryLower)) {
score = 40;
}
if (score > 0) {
results.push({ command, score });
}
}
// 按分数降序排序
results.sort((a, b) => b.score - a.score);
return results.slice(0, limit);
}
/**
* 列出所有 Commands(用于帮助显示)
*/
list(): Array<{ name: string; description?: string; source: string }> {
return this.getAll()
.map((cmd) => ({
name: cmd.name,
description: cmd.description,
source: cmd.source,
}))
.sort((a, b) => a.name.localeCompare(b.name));
}
/**
* 重新加载 Commands
*/
async reload(workdir: string = process.cwd()): Promise<void> {
this.commands.clear();
this.initialized = false;
await this.initialize(workdir);
}
/**
* 获取统计信息
*/
getStats(): {
total: number;
bySource: Record<string, number>;
} {
const commands = this.getAll();
const bySource: Record<string, number> = {};
for (const command of commands) {
bySource[command.source] = (bySource[command.source] || 0) + 1;
}
return {
total: commands.length,
bySource,
};
}
}
/**
* 全局 Command 注册表实例
*/
let commandRegistryInstance: CommandRegistry | null = null;
/**
* 获取全局 Command 注册表
*/
export function getCommandRegistry(): CommandRegistry {
if (!commandRegistryInstance) {
commandRegistryInstance = new CommandRegistry();
}
return commandRegistryInstance;
}
/**
* 重置全局 Command 注册表(用于测试)
*/
export function resetCommandRegistry(): void {
commandRegistryInstance = null;
}
+80
View File
@@ -0,0 +1,80 @@
/**
* Command 系统类型定义
*
* Command 是用户可通过斜杠命令触发的可复用提示词模板。
* 与 Skill 不同,Command 面向用户,可控制完整执行流程。
*/
/**
* Command 定义
*/
export interface Command {
/** Command 名称(从文件路径推断,如 deploy/staging) */
name: string;
/** Command 描述 */
description?: string;
/** 提示词模板 */
template: string;
/** 指定使用的 Agent */
agent?: string;
/** 指定使用的模型 (sonnet/opus/haiku) */
model?: string;
/** 是否作为子任务执行 */
subtask?: boolean;
/** 来源 */
source: 'builtin' | 'user' | 'project';
/** 来源路径 */
sourcePath?: string;
}
/**
* Command 执行输入
*/
export interface CommandInput {
/** Command 名称 */
command: string;
/** 原始参数字符串 */
arguments: string;
/** 解析后的参数数组 */
args: string[];
/** 当前工作目录 */
workdir: string;
}
/**
* Command 执行结果
*/
export interface CommandExecutionResult {
/** 是否成功 */
success: boolean;
/** 渲染后的提示 */
prompt?: string;
/** 指定的 Agent */
agent?: string;
/** 指定的模型 */
model?: string;
/** 是否作为子任务 */
subtask?: boolean;
/** 错误信息 */
error?: string;
}
/**
* Command 搜索结果
*/
export interface CommandSearchResult {
/** Command */
command: Command;
/** 匹配分数 */
score: number;
}
/**
* Command Frontmatter(Markdown 头部配置)
*/
export interface CommandFrontmatter {
description?: string;
agent?: string;
model?: string;
subtask?: boolean;
}
+196
View File
@@ -0,0 +1,196 @@
import { generateText, type ModelMessage, type LanguageModel } from 'ai';
import { TokenCounter } from './token-counter.js';
import {
SUMMARY_MARKER,
type CompressionConfig,
DEFAULT_COMPRESSION_CONFIG,
} from './types.js';
/**
* 摘要生成系统提示词
*/
const COMPACTION_SYSTEM_PROMPT = `你是一个专门生成对话摘要的助手。你的任务是将对话历史压缩成一个简洁但信息完整的摘要。
摘要应该包含:
1. 已完成的工作和关键结果
2. 当前正在进行的任务
3. 涉及的重要文件和代码
4. 用户的关键需求和约束
5. 下一步需要做的事情
要求:
- 保留关键技术细节(文件路径、函数名、配置等)
- 使用简洁的列表格式
- 不要遗漏重要信息
- 使用中文回复`;
/**
* 摘要生成用户提示词
*/
const COMPACTION_USER_PROMPT = `请总结上面的对话。这个摘要将是对话继续时唯一可用的上下文,所以要保留关键信息,包括:完成了什么、正在进行的工作、涉及的文件、下一步计划、以及用户的关键需求或约束。要简洁但足够详细,以便工作可以无缝继续。`;
/**
* 检查消息是否为摘要消息
*/
export function isSummaryMessage(message: ModelMessage): boolean {
if (typeof message.content === 'string') {
return message.content.includes(SUMMARY_MARKER);
}
if (Array.isArray(message.content)) {
return message.content.some(
(part) =>
typeof part === 'object' &&
'text' in part &&
typeof part.text === 'string' &&
part.text.includes(SUMMARY_MARKER)
);
}
return false;
}
/**
* 创建摘要消息
*/
function createSummaryMessage(summary: string): ModelMessage {
return {
role: 'assistant',
content: `${SUMMARY_MARKER}\n## 对话摘要\n\n${summary}\n${SUMMARY_MARKER}`,
};
}
/**
* Compaction 策略:使用 AI 生成对话摘要
*
* 逻辑:
* 1. 将历史消息(排除最近保护的部分)发送给 AI
* 2. AI 生成摘要
* 3. 用摘要消息替换旧消息
* 4. 保留最近的消息不变
*
* @param messages 消息数组
* @param model 语言模型
* @param config 压缩配置
* @returns 压缩后的消息数组和释放的 tokens
*/
export async function compact(
messages: ModelMessage[],
model: LanguageModel,
config: CompressionConfig = DEFAULT_COMPRESSION_CONFIG
): Promise<{ messages: ModelMessage[]; freedTokens: number }> {
const { pruneProtect } = config;
// 计算需要保护的消息数量
let protectedTokens = 0;
let protectedCount = 0;
for (let i = messages.length - 1; i >= 0; i--) {
const tokens = TokenCounter.estimateMessage(messages[i]);
if (protectedTokens + tokens > pruneProtect) {
break;
}
protectedTokens += tokens;
protectedCount++;
}
// 确保至少保护最后 2 条消息(除非 pruneProtect 为 0,表示强制压缩模式)
if (pruneProtect > 0) {
protectedCount = Math.max(protectedCount, 2);
} else {
// 强制压缩模式:至少保护 1 条消息
protectedCount = Math.max(protectedCount, 1);
}
// 分割消息:需要压缩的部分 vs 保护的部分
const toCompact = messages.slice(0, messages.length - protectedCount);
const toKeep = messages.slice(messages.length - protectedCount);
// 如果没有需要压缩的消息,直接返回
if (toCompact.length === 0) {
return { messages, freedTokens: 0 };
}
// 检查是否已有摘要消息
const existingSummaryIndex = toCompact.findIndex(isSummaryMessage);
const messagesForSummary =
existingSummaryIndex >= 0 ? toCompact.slice(existingSummaryIndex) : toCompact;
// 计算压缩前的 tokens
const beforeTokens = TokenCounter.estimateMessages(toCompact);
try {
// 调用 AI 生成摘要
const result = await generateText({
model,
system: COMPACTION_SYSTEM_PROMPT,
messages: [
...messagesForSummary,
{
role: 'user',
content: COMPACTION_USER_PROMPT,
},
],
maxOutputTokens: 2000,
});
const summaryMessage = createSummaryMessage(result.text);
const afterTokens = TokenCounter.estimateMessage(summaryMessage);
// 返回:摘要 + 保护的消息
return {
messages: [summaryMessage, ...toKeep],
freedTokens: beforeTokens - afterTokens,
};
} catch (error) {
console.error('生成摘要失败:', error);
// 失败时返回原消息
return { messages, freedTokens: 0 };
}
}
/**
* 简单压缩:不使用 AI,直接截断旧消息
* 用于没有模型可用或快速压缩的场景
*/
export function simpleCompact(
messages: ModelMessage[],
config: CompressionConfig = DEFAULT_COMPRESSION_CONFIG
): { messages: ModelMessage[]; freedTokens: number } {
const { pruneProtect } = config;
// 计算需要保留的消息
let keptTokens = 0;
let keepFromIndex = messages.length;
for (let i = messages.length - 1; i >= 0; i--) {
const tokens = TokenCounter.estimateMessage(messages[i]);
if (keptTokens + tokens > pruneProtect) {
break;
}
keptTokens += tokens;
keepFromIndex = i;
}
// 确保至少保留最后 N 条消息(强制模式下保留 1 条,否则保留 2 条)
const minKeep = pruneProtect > 0 ? 2 : 1;
keepFromIndex = Math.min(keepFromIndex, messages.length - minKeep);
const removed = messages.slice(0, keepFromIndex);
const kept = messages.slice(keepFromIndex);
if (removed.length === 0) {
return { messages, freedTokens: 0 };
}
// 创建简单摘要
const simpleSummary: ModelMessage = {
role: 'assistant',
content: `${SUMMARY_MARKER}\n[对话历史已压缩,共移除 ${removed.length} 条消息]\n${SUMMARY_MARKER}`,
};
const freedTokens = TokenCounter.estimateMessages(removed);
return {
messages: [simpleSummary, ...kept],
freedTokens,
};
}
+26
View File
@@ -0,0 +1,26 @@
// 类型导出
export type {
TokenUsage,
CompressionConfig,
CompressionContext,
CompressionResult,
} from './types.js';
export {
DEFAULT_COMPRESSION_CONFIG,
COMPACTED_PLACEHOLDER,
SUMMARY_MARKER,
COMPACTED_MARKER,
} from './types.js';
// Token 计数器
export { TokenCounter } from './token-counter.js';
// Prune 策略
export { prune, filterCompacted } from './prune.js';
// Compaction 策略
export { compact, simpleCompact, isSummaryMessage } from './compaction.js';
// 压缩管理器
export { CompressionManager, compressionManager } from './manager.js';
+238
View File
@@ -0,0 +1,238 @@
import type { ModelMessage, LanguageModel } from 'ai';
import { TokenCounter } from './token-counter.js';
import { prune, filterCompacted } from './prune.js';
import { compact, simpleCompact, isSummaryMessage } from './compaction.js';
import {
type TokenUsage,
type CompressionConfig,
type CompressionResult,
DEFAULT_COMPRESSION_CONFIG,
} from './types.js';
/**
* 压缩管理器
* 统一管理对话上下文的压缩策略
*/
export class CompressionManager {
private config: CompressionConfig;
private model: LanguageModel | null = null;
constructor(config: Partial<CompressionConfig> = {}) {
this.config = { ...DEFAULT_COMPRESSION_CONFIG, ...config };
}
/**
* 设置用于生成摘要的模型
*/
setModel(model: LanguageModel): void {
this.model = model;
}
/**
* 获取当前配置
*/
getConfig(): CompressionConfig {
return { ...this.config };
}
/**
* 更新配置
*/
updateConfig(config: Partial<CompressionConfig>): void {
this.config = { ...this.config, ...config };
}
/**
* 计算消息数组的 token 使用情况
*/
calculateUsage(messages: ModelMessage[]): TokenUsage {
const input = TokenCounter.estimateMessages(messages);
const { contextLimit, outputReserve } = this.config;
const available = contextLimit - outputReserve;
const usagePercent = (input / available) * 100;
return {
input,
contextLimit,
available,
usagePercent: Math.min(usagePercent, 100),
};
}
/**
* 检查是否需要压缩(超过溢出阈值)
*/
shouldCompress(messages: ModelMessage[]): boolean {
const usage = this.calculateUsage(messages);
return usage.usagePercent >= this.config.overflowThreshold * 100;
}
/**
* 检查是否溢出(超过可用空间)
*/
isOverflow(messages: ModelMessage[]): boolean {
const usage = this.calculateUsage(messages);
return usage.input >= usage.available;
}
/**
* 执行 prune 策略
*/
prune(messages: ModelMessage[]): { messages: ModelMessage[]; freedTokens: number } {
return prune(messages, this.config);
}
/**
* 执行 compaction 策略
*/
async compact(messages: ModelMessage[]): Promise<{ messages: ModelMessage[]; freedTokens: number }> {
if (this.model) {
return compact(messages, this.model, this.config);
}
// 没有模型时使用简单压缩
return simpleCompact(messages, this.config);
}
/**
* 自动压缩:先 prune,不够再 compact
*/
async compress(messages: ModelMessage[]): Promise<CompressionResult> {
let result = [...messages];
let totalFreed = 0;
let type: CompressionResult['type'] = 'prune';
// 第一步:尝试 prune
const pruneResult = this.prune(result);
if (pruneResult.freedTokens > 0) {
result = pruneResult.messages;
totalFreed += pruneResult.freedTokens;
}
// 检查是否还需要进一步压缩
if (this.shouldCompress(result)) {
// 第二步:执行 compaction
const compactResult = await this.compact(result);
if (compactResult.freedTokens > 0) {
result = compactResult.messages;
totalFreed += compactResult.freedTokens;
type = pruneResult.freedTokens > 0 ? 'both' : 'compaction';
}
}
return {
messages: result,
freedTokens: totalFreed,
type,
};
}
/**
* 强制压缩(用于 /compact 命令)
* 无论是否达到阈值都执行压缩
*/
async forceCompress(messages: ModelMessage[]): Promise<CompressionResult> {
// 消息数量太少时不压缩(至少需要 4 条消息)
if (messages.length <= 4) {
return {
messages,
freedTokens: 0,
type: 'prune',
};
}
let result = [...messages];
let totalFreed = 0;
let type: CompressionResult['type'] = 'prune';
// 先尝试 prune(使用强制配置)
const pruneConfig: CompressionConfig = {
...this.config,
pruneMinimum: 0,
pruneProtect: Math.min(10_000, TokenCounter.estimateMessages(messages) / 4),
};
const pruneResult = prune(result, pruneConfig);
if (pruneResult.freedTokens > 0) {
result = pruneResult.messages;
totalFreed += pruneResult.freedTokens;
}
// 强制 compaction:只保留最后 2 条消息
// 计算保留消息的 tokens
const keepCount = Math.min(2, result.length - 1);
const toKeep = result.slice(-keepCount);
const toCompact = result.slice(0, result.length - keepCount);
if (toCompact.length > 0) {
if (this.model) {
try {
const compactResult = await compact(result, this.model, {
...this.config,
pruneProtect: 0, // 强制模式:不保护任何 tokens
});
if (compactResult.freedTokens > 0) {
result = compactResult.messages;
totalFreed += compactResult.freedTokens;
type = pruneResult.freedTokens > 0 ? 'both' : 'compaction';
}
} catch {
// AI 压缩失败,使用简单压缩
const compactResult = simpleCompact(result, {
...this.config,
pruneProtect: 0,
});
if (compactResult.freedTokens > 0) {
result = compactResult.messages;
totalFreed += compactResult.freedTokens;
type = pruneResult.freedTokens > 0 ? 'both' : 'compaction';
}
}
} else {
const compactResult = simpleCompact(result, {
...this.config,
pruneProtect: 0,
});
if (compactResult.freedTokens > 0) {
result = compactResult.messages;
totalFreed += compactResult.freedTokens;
type = pruneResult.freedTokens > 0 ? 'both' : 'compaction';
}
}
}
return {
messages: result,
freedTokens: totalFreed,
type,
};
}
/**
* 过滤已压缩的内容(发送给模型前调用)
*/
filterCompacted(messages: ModelMessage[]): ModelMessage[] {
return filterCompacted(messages);
}
/**
* 检查消息是否为摘要消息
*/
isSummaryMessage(message: ModelMessage): boolean {
return isSummaryMessage(message);
}
/**
* 格式化 token 使用情况(用于 CLI 显示)
*/
formatUsage(messages: ModelMessage[]): string {
const usage = this.calculateUsage(messages);
const used = TokenCounter.format(usage.input);
const limit = TokenCounter.format(usage.available);
const percent = usage.usagePercent.toFixed(0);
return `${used}/${limit} (${percent}%)`;
}
}
// 导出单例(可选使用)
export const compressionManager = new CompressionManager();
+187
View File
@@ -0,0 +1,187 @@
import type { ModelMessage } from 'ai';
import { TokenCounter } from './token-counter.js';
import {
COMPACTED_PLACEHOLDER,
SUMMARY_MARKER,
COMPACTED_MARKER,
type CompressionConfig,
DEFAULT_COMPRESSION_CONFIG,
} from './types.js';
// 扩展的工具结果类型,支持压缩标记
interface CompactedToolResult {
type: 'tool-result';
toolCallId: string;
result: unknown;
[COMPACTED_MARKER]?: {
compactedAt: number;
originalSize: number;
};
}
/**
* 检查消息是否为摘要消息
*/
function isSummaryMessage(message: ModelMessage): boolean {
if (typeof message.content === 'string') {
return message.content.includes(SUMMARY_MARKER);
}
if (Array.isArray(message.content)) {
return message.content.some(
(part) =>
typeof part === 'object' &&
part !== null &&
'text' in part &&
typeof (part as { text?: unknown }).text === 'string' &&
((part as { text: string }).text).includes(SUMMARY_MARKER)
);
}
return false;
}
/**
* 检查是否为工具结果
*/
function isToolResult(part: unknown): part is CompactedToolResult {
return (
typeof part === 'object' &&
part !== null &&
(part as { type?: unknown }).type === 'tool-result'
);
}
/**
* 检查工具结果是否已压缩
*/
function isCompactedResult(part: unknown): boolean {
if (!isToolResult(part)) return false;
return COMPACTED_MARKER in part;
}
/**
* 获取工具结果的 token 数量
*/
function getToolResultTokens(part: unknown): number {
if (!isToolResult(part)) return 0;
return TokenCounter.estimateText(JSON.stringify(part.result));
}
/**
* 压缩工具结果
*/
function compactToolResult(part: CompactedToolResult): CompactedToolResult {
const originalSize = TokenCounter.estimateText(JSON.stringify(part.result));
return {
...part,
result: COMPACTED_PLACEHOLDER,
[COMPACTED_MARKER]: {
compactedAt: Date.now(),
originalSize,
},
};
}
/**
* Prune 策略:压缩旧的工具调用结果
*
* 逻辑:
* 1. 从后往前遍历消息
* 2. 跳过最近 pruneProtect tokens 的工具结果
* 3. 遇到 summary 消息停止
* 4. 将超出保护范围的 tool-result 替换为占位符
*
* @param messages 消息数组
* @param config 压缩配置
* @returns 处理后的消息数组和释放的 tokens
*/
export function prune(
messages: ModelMessage[],
config: CompressionConfig = DEFAULT_COMPRESSION_CONFIG
): { messages: ModelMessage[]; freedTokens: number } {
const { pruneProtect, pruneMinimum } = config;
// 深拷贝消息数组
const result = JSON.parse(JSON.stringify(messages)) as ModelMessage[];
let protectedTokens = 0;
let freedTokens = 0;
const toPrune: Array<{ msgIndex: number; partIndex: number; tokens: number }> = [];
// 从后往前遍历
for (let msgIndex = result.length - 1; msgIndex >= 0; msgIndex--) {
const message = result[msgIndex];
// 遇到摘要消息停止
if (isSummaryMessage(message)) {
break;
}
// 只处理包含工具结果的消息
if (!Array.isArray(message.content)) continue;
for (let partIndex = message.content.length - 1; partIndex >= 0; partIndex--) {
const part = message.content[partIndex];
// 跳过非工具结果
if (!isToolResult(part)) continue;
// 跳过已压缩的
if (isCompactedResult(part)) {
break; // 遇到已压缩的,说明之前已经 prune 过,停止
}
const tokens = getToolResultTokens(part);
protectedTokens += tokens;
// 超出保护范围的标记为待压缩
if (protectedTokens > pruneProtect) {
toPrune.push({ msgIndex, partIndex, tokens });
freedTokens += tokens;
}
}
}
// 如果释放的 tokens 不够最小量,不执行压缩
if (freedTokens < pruneMinimum) {
return { messages, freedTokens: 0 };
}
// 执行压缩
for (const { msgIndex, partIndex } of toPrune) {
const message = result[msgIndex];
if (Array.isArray(message.content)) {
const part = message.content[partIndex];
if (isToolResult(part)) {
// 使用 any 来绕过严格类型检查,因为我们在运行时知道这是安全的
(message.content as unknown[])[partIndex] = compactToolResult(part);
}
}
}
return { messages: result, freedTokens };
}
/**
* 过滤已压缩的内容(用于发送给模型前)
* 将已压缩的工具结果替换为占位符文本
*/
export function filterCompacted(messages: ModelMessage[]): ModelMessage[] {
return messages.map((message) => {
if (!Array.isArray(message.content)) return message;
const filteredContent = message.content.map((part) => {
if (isCompactedResult(part) && isToolResult(part)) {
return {
...part,
result: COMPACTED_PLACEHOLDER,
};
}
return part;
});
return {
...message,
content: filteredContent,
} as ModelMessage;
});
}
+100
View File
@@ -0,0 +1,100 @@
import type { ModelMessage } from 'ai';
/**
* Token 计数器
* 使用简单的字符估算,不依赖外部库
* 估算规则:
* - 中文字符:约 1.5 字符/token
* - 英文/数字:约 4 字符/token
* - 混合内容取平均
*/
export class TokenCounter {
/**
* 估算文本的 token 数量
*/
static estimateText(text: string): number {
if (!text) return 0;
// 统计中文字符数量
const chineseChars = (text.match(/[\u4e00-\u9fff]/g) || []).length;
// 其他字符数量
const otherChars = text.length - chineseChars;
// 中文约 1.5 字符/token,其他约 4 字符/token
const chineseTokens = chineseChars / 1.5;
const otherTokens = otherChars / 4;
return Math.ceil(chineseTokens + otherTokens);
}
/**
* 估算消息内容的 token 数量
*/
static estimateContent(content: ModelMessage['content']): number {
if (typeof content === 'string') {
return this.estimateText(content);
}
if (Array.isArray(content)) {
let total = 0;
for (const part of content) {
if (typeof part === 'string') {
total += this.estimateText(part);
} else if ('text' in part && typeof part.text === 'string') {
total += this.estimateText(part.text);
} else if ('result' in part) {
// tool-result
total += this.estimateText(JSON.stringify(part.result));
} else if ('args' in part) {
// tool-call
total += this.estimateText(JSON.stringify(part.args));
total += 20; // 工具名称等开销
} else {
// 其他类型,序列化估算
total += this.estimateText(JSON.stringify(part));
}
}
return total;
}
return 0;
}
/**
* 估算单条消息的 token 数量
*/
static estimateMessage(message: ModelMessage): number {
let tokens = 0;
// 角色标记开销
tokens += 4;
// 内容
tokens += this.estimateContent(message.content);
return tokens;
}
/**
* 估算消息数组的总 token 数量
*/
static estimateMessages(messages: ModelMessage[]): number {
let total = 0;
for (const message of messages) {
total += this.estimateMessage(message);
}
// 消息间分隔开销
total += messages.length * 3;
return total;
}
/**
* 格式化 token 数量显示
*/
static format(tokens: number): string {
if (tokens >= 1000) {
return `${(tokens / 1000).toFixed(1)}k`;
}
return `${tokens}`;
}
}
+77
View File
@@ -0,0 +1,77 @@
import type { LanguageModel } from 'ai';
/**
* Token 使用统计
*/
export interface TokenUsage {
/** 输入 tokens(估算) */
input: number;
/** 上下文限制 */
contextLimit: number;
/** 可用空间(contextLimit - outputReserve) */
available: number;
/** 使用百分比 (0-100) */
usagePercent: number;
}
/**
* 压缩配置
*/
export interface CompressionConfig {
/** 模型上下文限制 (默认 200k) */
contextLimit: number;
/** 预留输出 tokens (默认 32k) */
outputReserve: number;
/** 保护最近 tokens 不被 prune (默认 40k) */
pruneProtect: number;
/** 最小清理量才执行 prune (默认 20k) */
pruneMinimum: number;
/** 溢出阈值,超过此比例触发自动压缩 (默认 0.85) */
overflowThreshold: number;
}
/**
* 默认压缩配置
*/
export const DEFAULT_COMPRESSION_CONFIG: CompressionConfig = {
contextLimit: 200_000,
outputReserve: 32_000,
pruneProtect: 40_000,
pruneMinimum: 20_000,
overflowThreshold: 0.85,
};
/**
* 压缩占位符
*/
export const COMPACTED_PLACEHOLDER = '[此工具输出已压缩]';
/**
* 摘要消息标记 key
*/
export const SUMMARY_MARKER = '__summary__';
/**
* 工具结果压缩标记 key
*/
export const COMPACTED_MARKER = '__compacted__';
/**
* 压缩上下文(传递给压缩器的参数)
*/
export interface CompressionContext {
config: CompressionConfig;
model?: LanguageModel;
}
/**
* 压缩结果
*/
export interface CompressionResult {
/** 压缩后的消息 */
messages: import('ai').ModelMessage[];
/** 释放的 tokens */
freedTokens: number;
/** 压缩类型 */
type: 'prune' | 'compaction' | 'both';
}
+637
View File
@@ -0,0 +1,637 @@
import {
generateText,
streamText,
stepCountIs,
type ModelMessage,
type Tool as AITool,
type LanguageModel,
} from 'ai';
import type { Tool, ToolResult, Message, AgentConfig, UserInput, ContentBlock } from '../types/index.js';
import { buildZodSchema } from '../types/index.js';
import { ToolRegistry } from '../tools/registry.js';
import { SessionManager } from '../session/index.js';
import {
CompressionManager,
type TokenUsage,
type CompressionConfig,
} from '../context/index.js';
import type { AgentInfo, ImageData } from '../agent/types.js';
import { agentRegistry, AgentExecutor } from '../agent/index.js';
import { loadVisionConfig } from '../utils/config.js';
import { getModelFactory } from './providers.js';
import { getHookManager } from '../hooks/index.js';
import { getGitManager } from '../git/index.js';
export class Agent {
private getModel: (model: string) => LanguageModel;
private config: AgentConfig;
private conversationHistory: ModelMessage[] = [];
// 工具注册表
private registry: ToolRegistry | null = null;
// 已发现的工具(通过 tool_search 发现的)
private discoveredTools: Set<string> = new Set();
// 会话管理器(可选)
private sessionManager: SessionManager | null = null;
// 压缩管理器
private compressionManager: CompressionManager;
// 当前 Agent 模式(null 表示默认模式)
private currentAgentMode: AgentInfo | null = null;
// 原始 system prompt(用于切换回 default 时恢复)
private originalSystemPrompt: string;
constructor(config: AgentConfig, compressionConfig?: Partial<CompressionConfig>) {
this.config = config;
this.originalSystemPrompt = config.systemPrompt;
this.getModel = getModelFactory(config.provider, {
apiKey: config.apiKey,
baseUrl: config.baseUrl,
});
// 初始化压缩管理器
this.compressionManager = new CompressionManager(compressionConfig);
// 设置模型用于生成摘要
this.compressionManager.setModel(this.getModel(config.model));
}
/**
* 设置工具注册表(新模式:支持动态工具发现)
*/
setRegistry(registry: ToolRegistry): void {
this.registry = registry;
}
/**
* 设置会话管理器(启用会话持久化)
*/
setSessionManager(manager: SessionManager): void {
this.sessionManager = manager;
// 从会话恢复状态
const session = manager.getSession();
if (session) {
this.conversationHistory = [...session.messages];
this.discoveredTools = new Set(session.discoveredTools);
}
}
/**
* 获取会话管理器
*/
getSessionManager(): SessionManager | null {
return this.sessionManager;
}
/**
* 获取当前可用的工具
* 返回核心工具 + 已发现的工具,如果当前有 Agent 模式,应用工具过滤
*/
private getAvailableTools(): Tool[] {
if (!this.registry) {
throw new Error('工具注册表未初始化,请先调用 setRegistry()');
}
// 核心工具 + 已发现的工具
const coreTools = this.registry.getCoreTools();
const discoveredTools = this.registry.getTools([...this.discoveredTools]);
let tools = [...coreTools, ...discoveredTools];
// 应用 Agent 模式的工具过滤
if (this.currentAgentMode?.tools) {
tools = this.filterToolsByAgentConfig(tools);
}
return tools;
}
/**
* 根据 Agent 配置过滤工具
*/
private filterToolsByAgentConfig(tools: Tool[]): Tool[] {
const toolConfig = this.currentAgentMode?.tools;
if (!toolConfig) return tools;
let filteredTools = tools;
// 如果设置了 enabled 列表,只保留这些工具
if (toolConfig.enabled && toolConfig.enabled.length > 0) {
const enabledSet = new Set(toolConfig.enabled);
filteredTools = filteredTools.filter((t) => enabledSet.has(t.name));
}
// 如果设置了 disabled 列表,排除这些工具
if (toolConfig.disabled && toolConfig.disabled.length > 0) {
const disabledSet = new Set(toolConfig.disabled);
filteredTools = filteredTools.filter((t) => !disabledSet.has(t.name));
}
// 如果禁止嵌套 Task,移除 task 工具
if (toolConfig.noTask) {
filteredTools = filteredTools.filter((t) => t.name !== 'task');
}
return filteredTools;
}
/**
* 将工具转换为 Vercel AI SDK 的工具格式
*/
private getVercelTools(): Record<string, AITool> {
const vercelTools: Record<string, AITool> = {};
const availableTools = this.getAvailableTools();
const hookManager = getHookManager();
for (const tool of availableTools) {
const schema = buildZodSchema(tool.parameters);
vercelTools[tool.name] = {
description: tool.description,
inputSchema: schema,
execute: async (params) => {
const args = params as Record<string, unknown>;
const callId = `${tool.name}-${Date.now()}`;
const sessionId = this.sessionManager?.getSession()?.id || 'default';
// 触发工具执行前 hook
let finalArgs = args;
if (hookManager) {
const beforeOutput = await hookManager.triggerToolExecuteBefore({
tool: tool.name,
sessionId,
callId,
args,
});
// 如果 hook 指定跳过,直接返回
if (beforeOutput.skip && beforeOutput.skipResult) {
return beforeOutput.skipResult;
}
finalArgs = beforeOutput.args;
}
// 执行工具
const startTime = Date.now();
let result = await tool.execute(finalArgs);
const duration = Date.now() - startTime;
// 触发工具执行后 hook
if (hookManager) {
const afterOutput = await hookManager.triggerToolExecuteAfter(
{
tool: tool.name,
sessionId,
callId,
args: finalArgs,
duration,
},
result
);
result = afterOutput.result;
// 对于文件操作工具,触发相应的文件 hook 和 Git 自动提交
if (result.success) {
const filePath = finalArgs.path as string | undefined;
if (filePath) {
const gitManager = getGitManager();
if (tool.name === 'write_file') {
await hookManager.triggerFileCreated({
path: filePath,
tool: tool.name,
sessionId,
});
// Git 自动提交
if (gitManager) {
await gitManager.onFileChanged(filePath, 'create');
}
} else if (tool.name === 'edit_file') {
await hookManager.triggerFileEdited({
path: filePath,
tool: tool.name,
sessionId,
});
// Git 自动提交
if (gitManager) {
await gitManager.onFileChanged(filePath, 'modify');
}
} else if (tool.name === 'delete_file') {
await hookManager.triggerFileDeleted({
path: filePath,
tool: tool.name,
sessionId,
});
// Git 自动提交
if (gitManager) {
await gitManager.onFileChanged(filePath, 'delete');
}
}
}
}
}
// 如果是 tool_search 调用,解析结果并注入发现的工具
if (tool.name === 'tool_search' && result.success) {
this.handleToolSearchResult(result.output);
}
return result;
},
} as AITool;
}
return vercelTools;
}
/**
* 处理 tool_search 的结果,将发现的工具添加到可用列表
*/
private handleToolSearchResult(output: string): void {
// 解析输出,提取工具名称
// 格式: "- tool_name: description [category]"
const matches = output.matchAll(/^- (\w+):/gm);
for (const match of matches) {
const toolName = match[1];
if (this.registry?.has(toolName)) {
this.discoveredTools.add(toolName);
}
}
}
/**
* 发送消息并处理响应(流式)
* @param userMessage 用户消息文本或包含图片的 UserInput
* @param onStream 流式输出回调
*/
async chat(userMessage: string | UserInput, onStream?: (text: string) => void): Promise<string> {
// 处理带图片的消息
let processedMessage = userMessage;
if (typeof userMessage !== 'string' && userMessage.images && userMessage.images.length > 0) {
// 检查当前模型是否支持 vision
if (!this.supportsVision()) {
// 不支持 vision,尝试使用 Vision Agent 处理图片
const visionResult = await this.processImagesWithVisionAgent(
userMessage.images,
userMessage.text,
onStream
);
if (visionResult) {
// 成功,将图片分析结果转换为文本消息
processedMessage = visionResult;
} else {
// 失败,返回错误信息
return '无法处理图片:当前模型不支持图片理解,且 Vision 服务未配置或调用失败。';
}
}
}
// 构建消息内容
let messageContent: string | ContentBlock[];
if (typeof processedMessage === 'string') {
// 纯文本消息
messageContent = processedMessage;
} else {
// 带图片的消息
const blocks: ContentBlock[] = [];
// 添加图片
if (processedMessage.images && processedMessage.images.length > 0) {
for (const img of processedMessage.images) {
blocks.push({
type: 'image',
image: img.data,
mimeType: img.mimeType,
});
}
}
// 添加文本
if (processedMessage.text) {
blocks.push({
type: 'text',
text: processedMessage.text,
});
}
messageContent = blocks.length === 1 && blocks[0].type === 'text'
? blocks[0].text
: blocks;
}
// 添加用户消息到历史
this.conversationHistory.push({
role: 'user',
content: messageContent,
} as ModelMessage);
const vercelTools = this.getVercelTools();
let fullResponse = '';
let responseMessages: ModelMessage[] = [];
if (onStream) {
// 流式模式
const result = streamText({
model: this.getModel(this.config.model),
system: this.config.systemPrompt,
messages: this.conversationHistory,
tools: vercelTools,
maxOutputTokens: this.config.maxTokens,
stopWhen: stepCountIs(10), // 允许最多 10 轮工具调用
onChunk: ({ chunk }) => {
if (chunk.type === 'tool-call') {
onStream(`\n[调用工具: ${chunk.toolName}]\n`);
} else if (chunk.type === 'tool-result') {
const output = (chunk as { output?: ToolResult }).output;
if (output && typeof output === 'object') {
if (output.success) {
// 截断过长的输出
const displayOutput =
output.output.length > 500
? output.output.substring(0, 500) + '...(截断)'
: output.output;
onStream(`[结果: ${displayOutput}]\n`);
} else {
onStream(`[错误: ${output.error}]\n`);
}
}
}
},
});
// 流式输出文本
for await (const chunk of result.textStream) {
fullResponse += chunk;
onStream(chunk);
}
// 等待完成并获取完整的响应消息(包括工具调用和结果)
const response = await result.response;
responseMessages = response.messages as ModelMessage[];
} else {
// 非流式模式
const result = await generateText({
model: this.getModel(this.config.model),
system: this.config.systemPrompt,
messages: this.conversationHistory,
tools: vercelTools,
maxOutputTokens: this.config.maxTokens,
stopWhen: stepCountIs(10), // 允许最多 10 轮工具调用
});
fullResponse = result.text;
responseMessages = result.response.messages as ModelMessage[];
}
// 将完整的响应消息添加到历史(包括工具调用和结果)
this.conversationHistory.push(...responseMessages);
// 检查是否需要自动压缩
if (this.compressionManager.shouldCompress(this.conversationHistory)) {
const result = await this.compressionManager.compress(this.conversationHistory);
if (result.freedTokens > 0) {
this.conversationHistory = result.messages;
if (onStream) {
onStream(`\n[自动压缩: 释放了 ${(result.freedTokens / 1000).toFixed(1)}k tokens]\n`);
}
}
}
// 持久化会话
await this.persistSession();
return fullResponse;
}
/**
* 持久化当前会话状态
*/
private async persistSession(): Promise<void> {
if (!this.sessionManager) return;
await this.sessionManager.setMessages(this.conversationHistory);
await this.sessionManager.setDiscoveredTools([...this.discoveredTools]);
}
/**
* 清空对话历史和发现的工具
*/
async clearHistory(): Promise<void> {
this.conversationHistory = [];
this.discoveredTools.clear();
// 如果有会话管理器,创建新会话
if (this.sessionManager) {
await this.sessionManager.newSession();
}
}
/**
* 获取对话历史
*/
getHistory(): Message[] {
return this.conversationHistory
.filter(
(msg): msg is ModelMessage & { role: 'user' | 'assistant' } =>
msg.role === 'user' || msg.role === 'assistant'
)
.map((msg) => ({
role: msg.role,
content: typeof msg.content === 'string' ? msg.content : JSON.stringify(msg.content),
}));
}
/**
* 获取当前可用工具的数量
*/
getToolCount(): { core: number; discovered: number; total: number } {
if (!this.registry) {
return { core: 0, discovered: 0, total: 0 };
}
const coreCount = this.registry.getCoreTools().length;
const discoveredCount = this.discoveredTools.size;
return {
core: coreCount,
discovered: discoveredCount,
total: coreCount + discoveredCount,
};
}
/**
* 获取当前上下文使用情况
*/
getContextUsage(): TokenUsage {
return this.compressionManager.calculateUsage(this.conversationHistory);
}
/**
* 获取格式化的上下文使用情况(用于 CLI 显示)
*/
getContextUsageFormatted(): string {
return this.compressionManager.formatUsage(this.conversationHistory);
}
/**
* 获取压缩管理器
*/
getCompressionManager(): CompressionManager {
return this.compressionManager;
}
/**
* 手动压缩对话历史(用于 /compact 命令)
*/
async compactHistory(): Promise<{ freedTokens: number; type: string }> {
const result = await this.compressionManager.forceCompress(this.conversationHistory);
if (result.freedTokens > 0) {
this.conversationHistory = result.messages;
await this.persistSession();
}
return {
freedTokens: result.freedTokens,
type: result.type,
};
}
/**
* 切换 Agent 模式
*/
setAgentMode(agent: AgentInfo | null): void {
this.currentAgentMode = agent;
if (agent?.prompt) {
// 切换到指定 Agent,使用其 prompt
this.config = {
...this.config,
systemPrompt: agent.prompt,
};
} else {
// 切换回 default,恢复原始 prompt
this.config = {
...this.config,
systemPrompt: this.originalSystemPrompt,
};
}
}
/**
* 获取当前 Agent 模式
*/
getAgentMode(): AgentInfo | null {
return this.currentAgentMode;
}
/**
* 获取当前 Agent 名称
*/
getAgentModeName(): string {
return this.currentAgentMode?.name ?? 'default';
}
/**
* 使用 Vision Agent 处理图片
* 当主模型不支持 vision 时,委托给 Vision Agent 分析图片
* @returns 包含图片分析结果的文本消息,或 null 表示失败
*/
private async processImagesWithVisionAgent(
images: ImageData[],
userText?: string,
onStream?: (text: string) => void
): Promise<string | null> {
// 检查 Vision 配置是否可用
const visionConfig = loadVisionConfig();
if (!visionConfig) {
onStream?.('\n⚠ Vision 服务未配置,无法处理图片\n');
return null;
}
// 获取 Vision Agent
const visionAgent = agentRegistry.get('vision');
if (!visionAgent) {
onStream?.('\n⚠ Vision Agent 未注册\n');
return null;
}
// 确保有工具注册表
if (!this.registry) {
onStream?.('\n⚠ 工具注册表未初始化\n');
return null;
}
onStream?.(`\n[委托 Vision Agent (${visionConfig.model}) 分析图片...]\n`);
// 构建 Vision 配置
const visionAgentConfig: AgentConfig = {
...this.config,
provider: visionConfig.provider,
apiKey: visionConfig.apiKey,
model: visionConfig.model,
baseUrl: visionConfig.baseUrl,
};
// 创建 Vision Agent 执行器
const executor = new AgentExecutor(visionAgent, visionAgentConfig, this.registry);
// 构建提示词
const prompt = userText || '请详细描述这张图片的内容';
// 执行 Vision 分析
const result = await executor.execute(prompt, {
workdir: process.cwd(),
images,
onStream: undefined, // Vision Agent 不使用流式输出
});
if (!result.success) {
onStream?.(`\n⚠ Vision 分析失败: ${result.error}\n`);
return null;
}
onStream?.('\n[Vision 分析完成]\n');
// 构建带分析结果的文本消息
const combinedText = `[图片分析结果 - 由 ${visionConfig.model} 提供]\n${result.text}\n\n用户问题: ${userText || '(无附加问题)'}`;
return combinedText;
}
/**
* 检查当前模型是否支持 vision(图片理解)
*/
supportsVision(): boolean {
const model = this.config.model.toLowerCase();
// Anthropic Claude 模型支持 vision
if (this.config.provider === 'anthropic') {
// Claude 3 及以上版本支持 vision
return model.includes('claude-3') || model.includes('claude-4');
}
// OpenAI GPT-4 系列支持 vision
if (this.config.provider === 'openai') {
// GPT-4o, GPT-4 Turbo, GPT-4 Vision 等支持
return model.includes('gpt-4');
}
// DeepSeek 目前不支持 vision
if (this.config.provider === 'deepseek') {
return false;
}
return false;
}
/**
* 获取当前配置
*/
getConfig(): AgentConfig {
return { ...this.config };
}
}
+65
View File
@@ -0,0 +1,65 @@
import { createAnthropic } from '@ai-sdk/anthropic';
import { createDeepSeek } from '@ai-sdk/deepseek';
import { createOpenAI } from '@ai-sdk/openai';
import { createQwen } from 'qwen-ai-provider-v5';
import type { LanguageModel } from 'ai';
import type { ProviderType } from '../types/index.js';
/**
* Provider 配置选项
*/
export interface ProviderOptions {
apiKey: string;
baseUrl?: string;
}
/**
* Provider 工厂函数类型
*/
export type ProviderFactory = (options: ProviderOptions) => (model: string) => LanguageModel;
/**
* 检查 baseUrl 是否为阿里云百炼/DashScope
*/
function isDashScopeUrl(baseUrl?: string): boolean {
if (!baseUrl) return false;
return baseUrl.includes('dashscope');
}
/**
* Provider 注册表
* 支持 Anthropic、DeepSeek、OpenAI 及 OpenAI 兼容 API(如阿里云百炼)
*/
export const providers: Record<ProviderType, ProviderFactory> = {
anthropic: ({ apiKey, baseUrl }) => {
const client = createAnthropic({ apiKey, baseURL: baseUrl });
return (model) => client(model);
},
deepseek: ({ apiKey, baseUrl }) => {
const client = createDeepSeek({ apiKey, baseURL: baseUrl });
return (model) => client(model);
},
openai: ({ apiKey, baseUrl }) => {
// 如果是百炼的 URL,使用 qwen provider
if (isDashScopeUrl(baseUrl)) {
const client = createQwen({ apiKey, baseURL: baseUrl });
return (model) => client(model);
}
const client = createOpenAI({ apiKey, baseURL: baseUrl });
return (model) => client(model);
},
};
/**
* 获取模型工厂函数
*/
export function getModelFactory(
provider: ProviderType,
options: ProviderOptions
): (model: string) => LanguageModel {
const factory = providers[provider];
if (!factory) {
throw new Error(`不支持的 provider: ${provider}`);
}
return factory(options);
}
+459
View File
@@ -0,0 +1,459 @@
/**
* 编辑应用器
*
* 将编辑操作应用到文件
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import type {
Edit,
WholeFileEdit,
SearchReplaceEdit,
DiffEdit,
EditApplyResult,
EditStats,
BatchEdit,
BatchEditResult,
} from './types.js';
import { validateEdit } from './validator.js';
import { applyDiffPatch, normalizeSearchString, findSearchPositions } from './parsers.js';
import { touchFile, getFormattedFileDiagnostics, isLanguageSupported } from '../lsp/index.js';
/**
* 应用编辑选项
*/
export interface ApplyEditOptions {
/** 是否在应用前验证 */
validate?: boolean;
/** 是否创建备份 */
backup?: boolean;
/** 备份目录 */
backupDir?: string;
/** 是否运行 LSP 诊断 */
runDiagnostics?: boolean;
/** 是否为试运行(不实际写入) */
dryRun?: boolean;
}
const DEFAULT_OPTIONS: ApplyEditOptions = {
validate: true,
backup: false,
runDiagnostics: true,
dryRun: false,
};
/**
* 应用单个编辑
*/
export async function applyEdit(
edit: Edit,
options: ApplyEditOptions = {}
): Promise<EditApplyResult> {
const opts = { ...DEFAULT_OPTIONS, ...options };
// 验证编辑
if (opts.validate) {
const validation = await validateEdit(edit);
if (!validation.valid) {
const errorMessages = validation.errors.map(e => e.message).join('; ');
return {
success: false,
filePath: edit.filePath,
error: `验证失败: ${errorMessages}`,
};
}
}
// 读取原始内容
let originalContent: string | null = null;
try {
originalContent = await fs.readFile(edit.filePath, 'utf-8');
} catch {
// 文件不存在,可能是新建文件
}
// 计算新内容
let newContent: string;
try {
newContent = await computeNewContent(edit, originalContent);
} catch (error) {
return {
success: false,
filePath: edit.filePath,
error: error instanceof Error ? error.message : String(error),
};
}
// 计算统计
const stats = computeEditStats(originalContent, newContent, edit);
// 试运行模式
if (opts.dryRun) {
return {
success: true,
filePath: edit.filePath,
originalContent: originalContent ?? undefined,
newContent,
stats,
};
}
// 创建备份
if (opts.backup && originalContent !== null) {
await createBackup(edit.filePath, originalContent, opts.backupDir);
}
// 确保目录存在
try {
await fs.mkdir(path.dirname(edit.filePath), { recursive: true });
} catch (error) {
return {
success: false,
filePath: edit.filePath,
error: `创建目录失败: ${error instanceof Error ? error.message : String(error)}`,
};
}
// 写入新内容
try {
await fs.writeFile(edit.filePath, newContent, 'utf-8');
} catch (error) {
return {
success: false,
filePath: edit.filePath,
error: `写入文件失败: ${error instanceof Error ? error.message : String(error)}`,
};
}
// 运行 LSP 诊断
let diagnostics: string | undefined;
if (opts.runDiagnostics && isLanguageSupported(edit.filePath)) {
try {
const isFirstStart = await touchFile(edit.filePath, originalContent === null);
const waitTime = isFirstStart ? 2000 : 300;
await new Promise(resolve => setTimeout(resolve, waitTime));
diagnostics = await getFormattedFileDiagnostics(edit.filePath) ?? undefined;
} catch {
// LSP 错误不影响主流程
}
}
return {
success: true,
filePath: edit.filePath,
originalContent: originalContent ?? undefined,
newContent,
stats,
diagnostics,
};
}
/**
* 计算新内容
*/
async function computeNewContent(edit: Edit, originalContent: string | null): Promise<string> {
switch (edit.mode) {
case 'whole':
return computeWholeFileContent(edit);
case 'search-replace':
return computeSearchReplaceContent(edit, originalContent);
case 'diff':
return computeDiffContent(edit, originalContent);
default:
throw new Error(`不支持的编辑模式: ${(edit as Edit).mode}`);
}
}
/**
* 计算整文件替换的新内容
*/
function computeWholeFileContent(edit: WholeFileEdit): string {
return edit.content;
}
/**
* 计算搜索替换的新内容
*/
function computeSearchReplaceContent(
edit: SearchReplaceEdit,
originalContent: string | null
): string {
if (originalContent === null) {
throw new Error('搜索替换模式需要文件已存在');
}
let content = originalContent;
for (const block of edit.blocks) {
// 首先尝试直接匹配
let positions = findSearchPositions(content, block.search);
// 如果没有找到,尝试规范化后匹配
if (positions.length === 0) {
const normalizedSearch = normalizeSearchString(block.search, {
trimTrailingWhitespace: true,
normalizeLineEndings: true,
});
// 尝试在规范化的内容中查找
const normalizedContent = normalizeSearchString(content, {
trimTrailingWhitespace: true,
normalizeLineEndings: true,
});
const normalizedPositions = findSearchPositions(normalizedContent, normalizedSearch);
if (normalizedPositions.length === 1) {
// 找到规范化匹配,需要在原始内容中找到对应位置
// 使用逐行匹配的方式
content = replaceWithNormalization(content, block.search, block.replace);
continue;
}
if (normalizedPositions.length === 0) {
throw new Error(`未找到要替换的内容`);
}
if (normalizedPositions.length > 1) {
throw new Error(`找到 ${normalizedPositions.length} 处匹配,请提供更多上下文`);
}
}
if (positions.length === 0) {
throw new Error(`未找到要替换的内容`);
}
if (positions.length > 1) {
throw new Error(`找到 ${positions.length} 处匹配,请提供更多上下文`);
}
// 执行替换
content = content.replace(block.search, block.replace);
}
return content;
}
/**
* 使用规范化方式替换内容
*/
function replaceWithNormalization(
content: string,
search: string,
replace: string
): string {
const searchLines = search.split('\n');
const contentLines = content.split('\n');
// 找到第一行的匹配位置
const firstSearchLine = searchLines[0].trimEnd();
for (let i = 0; i <= contentLines.length - searchLines.length; i++) {
if (contentLines[i].trimEnd() === firstSearchLine) {
// 检查后续行是否都匹配
let allMatch = true;
for (let j = 1; j < searchLines.length; j++) {
if (contentLines[i + j].trimEnd() !== searchLines[j].trimEnd()) {
allMatch = false;
break;
}
}
if (allMatch) {
// 找到匹配,执行替换
const replaceLines = replace.split('\n');
const before = contentLines.slice(0, i);
const after = contentLines.slice(i + searchLines.length);
return [...before, ...replaceLines, ...after].join('\n');
}
}
}
throw new Error('规范化替换失败');
}
/**
* 计算 diff 应用后的新内容
*/
function computeDiffContent(edit: DiffEdit, originalContent: string | null): string {
if (originalContent === null) {
// 新文件的 diff
if (edit.patch.includes('new file mode')) {
// 从 diff 中提取新增的内容
const lines = edit.patch.split('\n');
const newLines: string[] = [];
for (const line of lines) {
if (line.startsWith('+') && !line.startsWith('+++')) {
newLines.push(line.slice(1));
}
}
return newLines.join('\n');
}
throw new Error('Diff 模式需要文件已存在');
}
return applyDiffPatch(originalContent, edit.patch);
}
/**
* 计算编辑统计
*/
function computeEditStats(
originalContent: string | null,
newContent: string,
edit: Edit
): EditStats {
const originalLines = originalContent?.split('\n') || [];
const newLines = newContent.split('\n');
// 简单统计:比较行数差异
const additions = Math.max(0, newLines.length - originalLines.length);
const deletions = Math.max(0, originalLines.length - newLines.length);
// 更精确的统计
let actualAdditions = 0;
let actualDeletions = 0;
if (originalContent !== null) {
const originalSet = new Set(originalLines);
const newSet = new Set(newLines);
for (const line of newLines) {
if (!originalSet.has(line)) {
actualAdditions++;
}
}
for (const line of originalLines) {
if (!newSet.has(line)) {
actualDeletions++;
}
}
} else {
actualAdditions = newLines.length;
}
const blocksApplied = edit.mode === 'search-replace' ? edit.blocks.length : 1;
return {
additions: actualAdditions || additions,
deletions: actualDeletions || deletions,
blocksApplied,
};
}
/**
* 创建备份
*/
async function createBackup(
filePath: string,
content: string,
backupDir?: string
): Promise<string> {
const dir = backupDir || path.join(path.dirname(filePath), '.edit-backups');
await fs.mkdir(dir, { recursive: true });
const timestamp = new Date().toISOString().replace(/[:.]/g, '-');
const filename = path.basename(filePath);
const backupPath = path.join(dir, `${filename}.${timestamp}.bak`);
await fs.writeFile(backupPath, content, 'utf-8');
return backupPath;
}
/**
* 应用批量编辑
*/
export async function applyBatchEdits(
batch: BatchEdit,
options: ApplyEditOptions = {}
): Promise<BatchEditResult> {
const results: EditApplyResult[] = [];
const opts = { ...DEFAULT_OPTIONS, ...options };
// 如果是原子操作,先验证所有编辑
if (batch.atomic) {
for (const edit of batch.edits) {
const validation = await validateEdit(edit);
if (!validation.valid) {
return {
success: false,
results: [{
success: false,
filePath: edit.filePath,
error: `验证失败: ${validation.errors.map(e => e.message).join('; ')}`,
}],
totalStats: { additions: 0, deletions: 0, blocksApplied: 0 },
};
}
}
}
// 应用所有编辑
const backups: Array<{ filePath: string; content: string }> = [];
for (const edit of batch.edits) {
// 对于原子操作,先保存原始内容用于回滚
if (batch.atomic) {
try {
const content = await fs.readFile(edit.filePath, 'utf-8');
backups.push({ filePath: edit.filePath, content });
} catch {
// 文件不存在,不需要备份
}
}
const result = await applyEdit(edit, { ...opts, validate: !batch.atomic });
results.push(result);
// 原子操作下,如果有失败则回滚
if (batch.atomic && !result.success) {
// 回滚已完成的编辑
for (const backup of backups) {
try {
await fs.writeFile(backup.filePath, backup.content, 'utf-8');
} catch {
// 回滚失败,记录但不抛出
console.error(`回滚失败: ${backup.filePath}`);
}
}
return {
success: false,
results,
totalStats: { additions: 0, deletions: 0, blocksApplied: 0 },
};
}
}
// 计算总统计
const totalStats: EditStats = {
additions: results.reduce((sum, r) => sum + (r.stats?.additions || 0), 0),
deletions: results.reduce((sum, r) => sum + (r.stats?.deletions || 0), 0),
blocksApplied: results.reduce((sum, r) => sum + (r.stats?.blocksApplied || 0), 0),
};
return {
success: results.every(r => r.success),
results,
totalStats,
};
}
/**
* 预览编辑效果
*/
export async function previewEdit(edit: Edit): Promise<EditApplyResult> {
return applyEdit(edit, { dryRun: true, runDiagnostics: false });
}
/**
* 预览批量编辑效果
*/
export async function previewBatchEdits(batch: BatchEdit): Promise<BatchEditResult> {
return applyBatchEdits(batch, { dryRun: true, runDiagnostics: false });
}
+115
View File
@@ -0,0 +1,115 @@
/**
* 编辑模式模块
*
* 提供统一的代码编辑接口,支持多种编辑模式:
* - whole: 整文件替换
* - search-replace: 搜索替换(支持多块)
* - diff: 统一 diff 格式
*/
// 类型导出
export type {
EditMode,
Edit,
WholeFileEdit,
SearchReplaceEdit,
DiffEdit,
SearchReplaceBlock,
EditOperation,
EditValidationResult,
EditValidationError,
EditValidationWarning,
EditApplyResult,
EditStats,
EditPreview,
DiffHunk,
DiffChange,
BatchEdit,
BatchEditResult,
EditorConfig,
} from './types.js';
export { DEFAULT_EDITOR_CONFIG } from './types.js';
// 解析器导出
export {
parseSearchReplaceBlocks,
createWholeFileEdit,
createSearchReplaceEdit,
createSingleSearchReplaceEdit,
createDiffEdit,
parseDiffPatch,
applyDiffPatch,
detectEditMode,
normalizeSearchString,
findSearchPositions,
getSearchLineNumbers,
} from './parsers.js';
// 验证器导出
export {
validateEdit,
validateEdits,
areAllEditsValid,
} from './validator.js';
// 应用器导出
export type { ApplyEditOptions } from './applier.js';
export {
applyEdit,
applyBatchEdits,
previewEdit,
previewBatchEdits,
} from './applier.js';
// 便捷函数
import type { Edit, SearchReplaceBlock, EditApplyResult } from './types.js';
import { createWholeFileEdit, createSearchReplaceEdit, createSingleSearchReplaceEdit } from './parsers.js';
import { applyEdit, type ApplyEditOptions } from './applier.js';
/**
* 快速写入文件(整文件模式)
*/
export async function writeFile(
filePath: string,
content: string,
options?: ApplyEditOptions
): Promise<EditApplyResult> {
const edit = createWholeFileEdit(filePath, content);
return applyEdit(edit, options);
}
/**
* 快速编辑文件(单块搜索替换)
*/
export async function editFile(
filePath: string,
search: string,
replace: string,
options?: ApplyEditOptions
): Promise<EditApplyResult> {
const edit = createSingleSearchReplaceEdit(filePath, search, replace);
return applyEdit(edit, options);
}
/**
* 批量搜索替换
*/
export async function editFileMultiple(
filePath: string,
blocks: SearchReplaceBlock[],
options?: ApplyEditOptions
): Promise<EditApplyResult> {
const edit = createSearchReplaceEdit(filePath, blocks);
return applyEdit(edit, options);
}
/**
* 应用任意编辑
*/
export async function apply(
edit: Edit,
options?: ApplyEditOptions
): Promise<EditApplyResult> {
return applyEdit(edit, options);
}
+297
View File
@@ -0,0 +1,297 @@
/**
* 编辑内容解析器
*
* 将各种格式的编辑指令解析为统一的 Edit 对象
*/
import type {
Edit,
WholeFileEdit,
SearchReplaceEdit,
SearchReplaceBlock,
DiffEdit,
} from './types.js';
/**
* 解析搜索替换块
*
* 支持多种格式:
* 1. 简单格式: { search: "...", replace: "..." }
* 2. 标记格式:
* <<<<<<< SEARCH
* 要搜索的内容
* =======
* 替换后的内容
* >>>>>>> REPLACE
*/
export function parseSearchReplaceBlocks(content: string): SearchReplaceBlock[] {
const blocks: SearchReplaceBlock[] = [];
// 尝试解析标记格式
const markerPattern = /<<<<<<<?[ ]*SEARCH\n([\s\S]*?)\n?=======\n?([\s\S]*?)\n?>>>>>>>?[ ]*REPLACE/g;
let match;
while ((match = markerPattern.exec(content)) !== null) {
blocks.push({
search: match[1],
replace: match[2],
});
}
// 如果找到了标记格式的块,直接返回
if (blocks.length > 0) {
return blocks;
}
// 尝试解析 JSON 数组格式
try {
const parsed = JSON.parse(content);
if (Array.isArray(parsed)) {
for (const item of parsed) {
if (typeof item.search === 'string' && typeof item.replace === 'string') {
blocks.push({
search: item.search,
replace: item.replace,
});
}
}
return blocks;
}
} catch {
// 不是 JSON 格式,继续
}
return blocks;
}
/**
* 创建整文件替换编辑
*/
export function createWholeFileEdit(filePath: string, content: string): WholeFileEdit {
return {
mode: 'whole',
filePath,
content,
};
}
/**
* 创建搜索替换编辑
*/
export function createSearchReplaceEdit(
filePath: string,
blocks: SearchReplaceBlock[]
): SearchReplaceEdit {
return {
mode: 'search-replace',
filePath,
blocks,
};
}
/**
* 创建单个搜索替换编辑
*/
export function createSingleSearchReplaceEdit(
filePath: string,
search: string,
replace: string
): SearchReplaceEdit {
return {
mode: 'search-replace',
filePath,
blocks: [{ search, replace }],
};
}
/**
* 创建 Diff 编辑
*/
export function createDiffEdit(filePath: string, patch: string): DiffEdit {
return {
mode: 'diff',
filePath,
patch,
};
}
/**
* 从统一 diff 格式解析编辑
*
* 支持标准的 unified diff 格式:
* --- a/file.txt
* +++ b/file.txt
* @@ -1,3 +1,4 @@
* context line
* -removed line
* +added line
* context line
*/
export function parseDiffPatch(patch: string): Map<string, string> {
const filePatches = new Map<string, string>();
// 按文件分割
const filePattern = /^diff --git a\/(.*) b\/(.*)$/gm;
const sections = patch.split(/(?=^diff --git)/m).filter(Boolean);
for (const section of sections) {
const headerMatch = section.match(/^diff --git a\/(.*) b\/(.*)$/m);
if (headerMatch) {
const filePath = headerMatch[2];
filePatches.set(filePath, section);
}
}
return filePatches;
}
/**
* 应用统一 diff 补丁到内容
*
* 简化实现:提取搜索替换块
*/
export function applyDiffPatch(originalContent: string, patch: string): string {
const lines = originalContent.split('\n');
const patchLines = patch.split('\n');
const result: string[] = [];
let lineIndex = 0;
let patchIndex = 0;
// 跳过 diff 头部
while (patchIndex < patchLines.length && !patchLines[patchIndex].startsWith('@@')) {
patchIndex++;
}
while (patchIndex < patchLines.length) {
const line = patchLines[patchIndex];
if (line.startsWith('@@')) {
// 解析 hunk 头部: @@ -start,count +start,count @@
const hunkMatch = line.match(/@@ -(\d+),?(\d*) \+(\d+),?(\d*) @@/);
if (hunkMatch) {
const oldStart = parseInt(hunkMatch[1], 10);
// 复制 hunk 之前的未修改行
while (lineIndex < oldStart - 1 && lineIndex < lines.length) {
result.push(lines[lineIndex]);
lineIndex++;
}
}
patchIndex++;
continue;
}
if (line.startsWith('+') && !line.startsWith('+++')) {
// 新增行
result.push(line.slice(1));
patchIndex++;
} else if (line.startsWith('-') && !line.startsWith('---')) {
// 删除行 - 跳过原始内容
lineIndex++;
patchIndex++;
} else if (line.startsWith(' ') || line === '') {
// 上下文行
if (lineIndex < lines.length) {
result.push(lines[lineIndex]);
lineIndex++;
}
patchIndex++;
} else {
patchIndex++;
}
}
// 复制剩余的原始行
while (lineIndex < lines.length) {
result.push(lines[lineIndex]);
lineIndex++;
}
return result.join('\n');
}
/**
* 检测编辑模式
*
* 根据内容特征自动检测合适的编辑模式
*/
export function detectEditMode(content: string, fileSize: number): 'whole' | 'search-replace' | 'diff' {
// 检测 diff 格式
if (content.includes('diff --git') || content.match(/^@@\s+-\d+,?\d*\s+\+\d+,?\d*\s+@@/m)) {
return 'diff';
}
// 检测搜索替换标记格式
if (content.includes('<<<<<<< SEARCH') || content.includes('<<<<<<SEARCH')) {
return 'search-replace';
}
// 大文件优先使用 search-replace(如果有明显的结构)
if (fileSize > 50 * 1024) { // 50KB
return 'search-replace';
}
// 默认使用 whole
return 'whole';
}
/**
* 规范化搜索字符串
*
* 处理常见的格式问题:
* - 行尾空格
* - 不同的换行符
* - 缩进差异
*/
export function normalizeSearchString(search: string, options: {
trimTrailingWhitespace?: boolean;
normalizeLineEndings?: boolean;
normalizeIndentation?: boolean;
} = {}): string {
let result = search;
// 规范化换行符
if (options.normalizeLineEndings !== false) {
result = result.replace(/\r\n/g, '\n').replace(/\r/g, '\n');
}
// 去除行尾空格
if (options.trimTrailingWhitespace) {
result = result.split('\n').map(line => line.trimEnd()).join('\n');
}
return result;
}
/**
* 在内容中查找搜索字符串的位置
*/
export function findSearchPositions(content: string, search: string): number[] {
const positions: number[] = [];
let pos = 0;
while (true) {
const index = content.indexOf(search, pos);
if (index === -1) break;
positions.push(index);
pos = index + 1;
}
return positions;
}
/**
* 获取搜索字符串在内容中的行号
*/
export function getSearchLineNumbers(content: string, search: string): number[] {
const positions = findSearchPositions(content, search);
const lineNumbers: number[] = [];
for (const pos of positions) {
const lineNumber = content.slice(0, pos).split('\n').length;
lineNumbers.push(lineNumber);
}
return lineNumbers;
}
+234
View File
@@ -0,0 +1,234 @@
/**
* 编辑模式类型定义
*
* 支持多种编辑模式:
* - whole: 整文件替换
* - search-replace: 搜索替换(支持多块)
* - diff: 统一 diff 格式(未来扩展)
*/
/**
* 编辑模式类型
*/
export type EditMode = 'whole' | 'search-replace' | 'diff';
/**
* 单个搜索替换块
*/
export interface SearchReplaceBlock {
/** 要搜索的原始字符串 */
search: string;
/** 替换后的新字符串 */
replace: string;
}
/**
* 编辑操作(通用接口)
*/
export interface EditOperation {
/** 编辑模式 */
mode: EditMode;
/** 目标文件路径 */
filePath: string;
}
/**
* 整文件替换操作
*/
export interface WholeFileEdit extends EditOperation {
mode: 'whole';
/** 新的文件内容 */
content: string;
}
/**
* 搜索替换操作
*/
export interface SearchReplaceEdit extends EditOperation {
mode: 'search-replace';
/** 替换块列表 */
blocks: SearchReplaceBlock[];
}
/**
* Diff 格式操作(统一 diff)
*/
export interface DiffEdit extends EditOperation {
mode: 'diff';
/** 统一 diff 格式的补丁内容 */
patch: string;
}
/**
* 所有编辑操作的联合类型
*/
export type Edit = WholeFileEdit | SearchReplaceEdit | DiffEdit;
/**
* 编辑验证结果
*/
export interface EditValidationResult {
/** 是否有效 */
valid: boolean;
/** 错误信息列表 */
errors: EditValidationError[];
/** 警告信息列表 */
warnings: EditValidationWarning[];
}
/**
* 验证错误
*/
export interface EditValidationError {
/** 错误类型 */
type: 'not_found' | 'ambiguous' | 'conflict' | 'syntax' | 'permission';
/** 错误消息 */
message: string;
/** 相关的搜索字符串(如果适用) */
search?: string;
/** 在文件中找到的匹配数量(如果适用) */
occurrences?: number;
/** 行号(如果适用) */
line?: number;
}
/**
* 验证警告
*/
export interface EditValidationWarning {
/** 警告类型 */
type: 'whitespace' | 'large_change' | 'binary';
/** 警告消息 */
message: string;
}
/**
* 编辑应用结果
*/
export interface EditApplyResult {
/** 是否成功 */
success: boolean;
/** 目标文件路径 */
filePath: string;
/** 原始内容 */
originalContent?: string;
/** 新内容 */
newContent?: string;
/** 错误信息 */
error?: string;
/** 变更统计 */
stats?: EditStats;
/** 代码诊断结果(如果启用 LSP) */
diagnostics?: string;
}
/**
* 编辑统计
*/
export interface EditStats {
/** 新增行数 */
additions: number;
/** 删除行数 */
deletions: number;
/** 修改的块数 */
blocksApplied: number;
}
/**
* 编辑预览
*/
export interface EditPreview {
/** 文件路径 */
filePath: string;
/** 是否是新文件 */
isNewFile: boolean;
/** 原始内容(null 表示新文件) */
originalContent: string | null;
/** 预览的新内容 */
previewContent: string;
/** 变更统计 */
stats: EditStats;
/** Diff hunks(用于显示) */
hunks: DiffHunk[];
}
/**
* Diff Hunk(差异块)
*/
export interface DiffHunk {
/** 原文件起始行 */
oldStart: number;
/** 原文件行数 */
oldLines: number;
/** 新文件起始行 */
newStart: number;
/** 新文件行数 */
newLines: number;
/** 变更行 */
changes: DiffChange[];
}
/**
* Diff 变更行
*/
export interface DiffChange {
/** 变更类型 */
type: 'add' | 'remove' | 'context';
/** 行内容 */
content: string;
/** 原文件行号(删除和上下文行有效) */
oldLineNumber?: number;
/** 新文件行号(新增和上下文行有效) */
newLineNumber?: number;
}
/**
* 批量编辑操作
*/
export interface BatchEdit {
/** 编辑列表 */
edits: Edit[];
/** 是否原子操作(全部成功或全部回滚) */
atomic?: boolean;
}
/**
* 批量编辑结果
*/
export interface BatchEditResult {
/** 是否全部成功 */
success: boolean;
/** 各个编辑的结果 */
results: EditApplyResult[];
/** 总体统计 */
totalStats: EditStats;
}
/**
* 编辑器配置
*/
export interface EditorConfig {
/** 默认编辑模式 */
defaultMode: EditMode;
/** 是否启用 LSP 诊断 */
enableDiagnostics: boolean;
/** 大文件阈值(超过此字节数使用 search-replace 模式) */
largeFileThreshold: number;
/** 是否在应用前验证 */
validateBeforeApply: boolean;
/** 是否创建备份 */
createBackup: boolean;
/** 备份目录 */
backupDir?: string;
}
/**
* 默认编辑器配置
*/
export const DEFAULT_EDITOR_CONFIG: EditorConfig = {
defaultMode: 'search-replace',
enableDiagnostics: true,
largeFileThreshold: 100 * 1024, // 100KB
validateBeforeApply: true,
createBackup: false,
};
+414
View File
@@ -0,0 +1,414 @@
/**
* 编辑验证器
*
* 在应用编辑之前验证编辑的有效性和安全性
*/
import * as fs from 'fs/promises';
import type {
Edit,
WholeFileEdit,
SearchReplaceEdit,
DiffEdit,
EditValidationResult,
EditValidationError,
EditValidationWarning,
SearchReplaceBlock,
} from './types.js';
import { findSearchPositions, normalizeSearchString } from './parsers.js';
/**
* 验证编辑操作
*/
export async function validateEdit(edit: Edit): Promise<EditValidationResult> {
const errors: EditValidationError[] = [];
const warnings: EditValidationWarning[] = [];
// 检查文件是否存在(对于非新建文件的操作)
let fileContent: string | null = null;
let fileExists = false;
try {
fileContent = await fs.readFile(edit.filePath, 'utf-8');
fileExists = true;
} catch {
// 文件不存在
}
// 根据编辑模式进行验证
switch (edit.mode) {
case 'whole':
validateWholeFileEdit(edit, fileContent, fileExists, errors, warnings);
break;
case 'search-replace':
validateSearchReplaceEdit(edit, fileContent, fileExists, errors, warnings);
break;
case 'diff':
validateDiffEdit(edit, fileContent, fileExists, errors, warnings);
break;
}
return {
valid: errors.length === 0,
errors,
warnings,
};
}
/**
* 验证整文件替换编辑
*/
function validateWholeFileEdit(
edit: WholeFileEdit,
fileContent: string | null,
_fileExists: boolean,
errors: EditValidationError[],
warnings: EditValidationWarning[]
): void {
// 检查是否是二进制文件
if (fileContent !== null && isBinaryContent(fileContent)) {
warnings.push({
type: 'binary',
message: '目标文件可能是二进制文件',
});
}
// 检查变更幅度
if (fileContent !== null) {
const changeRatio = Math.abs(edit.content.length - fileContent.length) / fileContent.length;
if (changeRatio > 0.8 && fileContent.length > 1000) {
warnings.push({
type: 'large_change',
message: `文件内容变化较大 (${Math.round(changeRatio * 100)}%),请确认是否正确`,
});
}
}
// 检查空白字符问题
if (edit.content.includes('\t') && fileContent?.includes(' ')) {
warnings.push({
type: 'whitespace',
message: '新内容使用 Tab 缩进,但原文件使用空格缩进',
});
} else if (edit.content.includes(' ') && fileContent?.includes('\t')) {
warnings.push({
type: 'whitespace',
message: '新内容使用空格缩进,但原文件使用 Tab 缩进',
});
}
}
/**
* 验证搜索替换编辑
*/
function validateSearchReplaceEdit(
edit: SearchReplaceEdit,
fileContent: string | null,
fileExists: boolean,
errors: EditValidationError[],
warnings: EditValidationWarning[]
): void {
// 文件必须存在
if (!fileExists || fileContent === null) {
errors.push({
type: 'not_found',
message: `文件不存在: ${edit.filePath}`,
});
return;
}
// 验证每个搜索替换块
for (const block of edit.blocks) {
validateSearchReplaceBlock(block, fileContent, errors, warnings);
}
// 检查是否有重叠的搜索区域
checkOverlappingBlocks(edit.blocks, fileContent, errors);
}
/**
* 验证单个搜索替换块
*/
function validateSearchReplaceBlock(
block: SearchReplaceBlock,
fileContent: string,
errors: EditValidationError[],
warnings: EditValidationWarning[]
): void {
const { search, replace } = block;
// 检查搜索字符串是否为空
if (!search || search.length === 0) {
errors.push({
type: 'syntax',
message: '搜索字符串不能为空',
search,
});
return;
}
// 查找匹配
const positions = findSearchPositions(fileContent, search);
if (positions.length === 0) {
// 尝试规范化后再查找
const normalizedSearch = normalizeSearchString(search, {
trimTrailingWhitespace: true,
normalizeLineEndings: true,
});
const normalizedContent = normalizeSearchString(fileContent, {
trimTrailingWhitespace: true,
normalizeLineEndings: true,
});
const normalizedPositions = findSearchPositions(normalizedContent, normalizedSearch);
if (normalizedPositions.length > 0) {
warnings.push({
type: 'whitespace',
message: '搜索字符串与文件内容存在空白字符差异,已自动处理',
});
} else {
// 提供更友好的错误信息
const similarMatch = findSimilarMatch(fileContent, search);
let errorMessage = `未找到要替换的内容`;
if (similarMatch) {
errorMessage += `。找到相似内容在第 ${similarMatch.line} 行`;
}
errors.push({
type: 'not_found',
message: errorMessage,
search: search.length > 100 ? search.slice(0, 100) + '...' : search,
occurrences: 0,
});
}
return;
}
if (positions.length > 1) {
// 找到多个匹配
const lineNumbers = positions.map(pos =>
fileContent.slice(0, pos).split('\n').length
);
errors.push({
type: 'ambiguous',
message: `找到 ${positions.length} 处匹配(行 ${lineNumbers.join(', ')})。请提供更多上下文使搜索字符串唯一`,
search: search.length > 100 ? search.slice(0, 100) + '...' : search,
occurrences: positions.length,
});
return;
}
// 检查替换后是否有意义
if (search === replace) {
warnings.push({
type: 'whitespace',
message: '搜索和替换内容相同,此操作无效',
});
}
}
/**
* 验证 Diff 编辑
*/
function validateDiffEdit(
edit: DiffEdit,
fileContent: string | null,
fileExists: boolean,
errors: EditValidationError[],
_warnings: EditValidationWarning[]
): void {
// 文件必须存在(除非是新建文件的 diff)
if (!fileExists && !edit.patch.includes('new file mode')) {
errors.push({
type: 'not_found',
message: `文件不存在: ${edit.filePath}`,
});
return;
}
// 验证 diff 格式
if (!edit.patch.includes('@@') && !edit.patch.includes('diff --git')) {
errors.push({
type: 'syntax',
message: 'Diff 格式无效,缺少 hunk 头部 (@@) 或 diff 头部',
});
return;
}
// 验证上下文行是否匹配
if (fileContent) {
const contextErrors = validateDiffContext(edit.patch, fileContent);
errors.push(...contextErrors);
}
}
/**
* 验证 diff 中的上下文行是否与文件内容匹配
*/
function validateDiffContext(patch: string, fileContent: string): EditValidationError[] {
const errors: EditValidationError[] = [];
const fileLines = fileContent.split('\n');
const patchLines = patch.split('\n');
let currentHunkOldStart = 0;
let oldLineOffset = 0;
for (const line of patchLines) {
// 解析 hunk 头部
const hunkMatch = line.match(/@@ -(\d+),?(\d*) \+(\d+),?(\d*) @@/);
if (hunkMatch) {
currentHunkOldStart = parseInt(hunkMatch[1], 10);
oldLineOffset = 0;
continue;
}
// 检查上下文行
if (line.startsWith(' ')) {
const contextContent = line.slice(1);
const fileLineIndex = currentHunkOldStart - 1 + oldLineOffset;
if (fileLineIndex < fileLines.length) {
if (fileLines[fileLineIndex] !== contextContent) {
errors.push({
type: 'conflict',
message: `第 ${fileLineIndex + 1} 行的上下文不匹配`,
line: fileLineIndex + 1,
});
}
}
oldLineOffset++;
} else if (line.startsWith('-')) {
oldLineOffset++;
}
}
return errors;
}
/**
* 检查搜索块是否有重叠
*/
function checkOverlappingBlocks(
blocks: SearchReplaceBlock[],
fileContent: string,
errors: EditValidationError[]
): void {
const ranges: Array<{ start: number; end: number; index: number }> = [];
for (let i = 0; i < blocks.length; i++) {
const positions = findSearchPositions(fileContent, blocks[i].search);
for (const pos of positions) {
ranges.push({
start: pos,
end: pos + blocks[i].search.length,
index: i,
});
}
}
// 检查重叠
ranges.sort((a, b) => a.start - b.start);
for (let i = 0; i < ranges.length - 1; i++) {
if (ranges[i].end > ranges[i + 1].start && ranges[i].index !== ranges[i + 1].index) {
errors.push({
type: 'conflict',
message: `搜索块 ${ranges[i].index + 1} 和 ${ranges[i + 1].index + 1} 存在重叠区域`,
});
}
}
}
/**
* 检测是否是二进制内容
*/
function isBinaryContent(content: string): boolean {
// 检查是否包含 null 字符或大量不可打印字符
const nonPrintableCount = (content.match(/[\x00-\x08\x0B\x0C\x0E-\x1F]/g) || []).length;
return nonPrintableCount > content.length * 0.1;
}
/**
* 查找相似匹配(用于错误提示)
*/
function findSimilarMatch(
content: string,
search: string
): { line: number; similarity: number } | null {
// 取搜索字符串的第一行作为关键内容
const firstLine = search.split('\n')[0].trim();
if (firstLine.length < 5) {
return null;
}
const lines = content.split('\n');
let bestMatch: { line: number; similarity: number } | null = null;
for (let i = 0; i < lines.length; i++) {
const line = lines[i];
if (line.includes(firstLine.slice(0, Math.min(20, firstLine.length)))) {
const similarity = calculateSimilarity(line, firstLine);
if (!bestMatch || similarity > bestMatch.similarity) {
bestMatch = { line: i + 1, similarity };
}
}
}
return bestMatch && bestMatch.similarity > 0.5 ? bestMatch : null;
}
/**
* 计算字符串相似度(简单实现)
*/
function calculateSimilarity(a: string, b: string): number {
if (a === b) return 1;
if (a.length === 0 || b.length === 0) return 0;
const longer = a.length > b.length ? a : b;
const shorter = a.length > b.length ? b : a;
const longerLength = longer.length;
if (longerLength === 0) return 1;
// 简单的包含检查
if (longer.includes(shorter)) {
return shorter.length / longerLength;
}
// 计算共同字符数
let matches = 0;
const shorterChars = shorter.split('');
const longerChars = longer.split('');
for (const char of shorterChars) {
const idx = longerChars.indexOf(char);
if (idx !== -1) {
matches++;
longerChars.splice(idx, 1);
}
}
return matches / longerLength;
}
/**
* 批量验证编辑
*/
export async function validateEdits(edits: Edit[]): Promise<EditValidationResult[]> {
return Promise.all(edits.map(edit => validateEdit(edit)));
}
/**
* 检查所有编辑是否都有效
*/
export async function areAllEditsValid(edits: Edit[]): Promise<{
valid: boolean;
results: EditValidationResult[];
}> {
const results = await validateEdits(edits);
const valid = results.every(r => r.valid);
return { valid, results };
}
+196
View File
@@ -0,0 +1,196 @@
/**
* 自动提交管理器
*
* 参考 aider 的 auto_commit 实现
* 支持 immediate、batch、manual 三种模式
*/
import { minimatch } from 'minimatch';
import type { GitRepo } from './repo.js';
import type { AutoCommitConfig, CommitResult } from './types.js';
import { MessageGenerator } from './message-generator.js';
export class AutoCommitManager {
private repo: GitRepo;
private config: AutoCommitConfig;
private messageGenerator: MessageGenerator;
/** 待提交的文件 */
private pendingFiles: Set<string> = new Set();
/** 批量提交定时器 */
private batchTimer: NodeJS.Timeout | null = null;
/** 提交回调 */
private onCommitCallback?: (result: CommitResult) => void;
constructor(repo: GitRepo, config: AutoCommitConfig, messageGenerator: MessageGenerator) {
this.repo = repo;
this.config = config;
this.messageGenerator = messageGenerator;
}
/**
* 设置提交回调
*/
setOnCommit(callback: (result: CommitResult) => void): void {
this.onCommitCallback = callback;
}
/**
* 文件变更后调用
*/
async onFileChanged(filePath: string, changeType: 'create' | 'modify' | 'delete'): Promise<void> {
if (!this.config.enabled) {
return;
}
// 检查是否应该排除
if (this.shouldExclude(filePath)) {
return;
}
// 如果启用脏文件提交,检查文件是否为脏状态
if (this.config.dirtyCommits && changeType !== 'create') {
const isDirty = await this.repo.isDirty(filePath);
if (isDirty) {
// 在新的编辑之前,先提交脏文件
await this.commitDirtyFile(filePath);
}
}
// 添加到待提交列表
this.pendingFiles.add(filePath);
// 根据模式处理
switch (this.config.mode) {
case 'immediate':
await this.executeCommit();
break;
case 'batch':
this.scheduleBatchCommit();
break;
case 'manual':
// 不自动提交,只记录
break;
}
}
/**
* 计划批量提交
*/
private scheduleBatchCommit(): void {
if (this.batchTimer) {
clearTimeout(this.batchTimer);
}
this.batchTimer = setTimeout(async () => {
this.batchTimer = null;
await this.executeCommit();
}, this.config.batchDelay);
}
/**
* 执行提交
*/
private async executeCommit(): Promise<CommitResult | null> {
if (this.pendingFiles.size === 0) {
return null;
}
const files = Array.from(this.pendingFiles);
this.pendingFiles.clear();
try {
// 获取差异用于生成消息
const diff = await this.repo.getDiff({ staged: false });
// 生成提交消息
const message = this.messageGenerator.generate(diff, files);
// 执行提交
const result = await this.repo.commit({
files,
message,
aiEdits: true,
});
// 调用回调
if (result.success && this.onCommitCallback) {
this.onCommitCallback(result);
}
return result;
} catch (error) {
// 恢复 pending 状态
files.forEach((f) => this.pendingFiles.add(f));
return {
success: false,
error: error instanceof Error ? error.message : String(error),
};
}
}
/**
* 提交脏文件
*/
private async commitDirtyFile(filePath: string): Promise<void> {
try {
await this.repo.commit({
files: [filePath],
message: `chore: save changes to ${filePath} before AI edit`,
aiEdits: false, // 用户的脏文件,不标记为 AI 编辑
});
} catch {
// 脏文件提交失败不影响主流程
}
}
/**
* 检查文件是否应该排除
*/
private shouldExclude(filePath: string): boolean {
return this.config.excludePatterns.some((pattern) =>
minimatch(filePath, pattern, { matchBase: true })
);
}
/**
* 强制立即提交
*/
async flush(): Promise<CommitResult | null> {
if (this.batchTimer) {
clearTimeout(this.batchTimer);
this.batchTimer = null;
}
return this.executeCommit();
}
/**
* 获取待提交文件
*/
getPendingFiles(): string[] {
return Array.from(this.pendingFiles);
}
/**
* 清除待提交文件
*/
clearPending(): void {
if (this.batchTimer) {
clearTimeout(this.batchTimer);
this.batchTimer = null;
}
this.pendingFiles.clear();
}
/**
* 是否有待提交文件
*/
hasPendingFiles(): boolean {
return this.pendingFiles.size > 0;
}
}
+53
View File
@@ -0,0 +1,53 @@
/**
* Git 深度集成模块
*
* 提供自动提交、智能 commit message 生成、undo 等功能
* 参考 aider 的实现
*/
// Git 管理器
export {
GitManager,
getGitManager,
initGitManager,
resetGitManager,
} from './manager.js';
// GitRepo
export { GitRepo } from './repo.js';
// 自动提交
export { AutoCommitManager } from './auto-commit.js';
// 消息生成
export { MessageGenerator } from './message-generator.js';
// Undo 管理
export { UndoManager } from './undo-manager.js';
// 类型导出
export type {
GitConfig,
AutoCommitConfig,
UndoConfig,
MessageFormatConfig,
AttributionConfig,
GitStatus,
FileChange,
ChangeStatus,
CommitInfo,
DiffResult,
FileDiff,
DiffHunk,
DiffStats,
UndoEntry,
UndoResult,
CommitOptions,
CommitResult,
BranchInfo,
GitEventType,
GitEvent,
GitEventListener,
} from './types.js';
export { DEFAULT_GIT_CONFIG } from './types.js';
+329
View File
@@ -0,0 +1,329 @@
/**
* Git 管理器
*
* 整合 GitRepo、AutoCommitManager、MessageGenerator、UndoManager
* 提供统一的 Git 集成接口
*/
import type {
GitConfig,
GitStatus,
CommitResult,
UndoResult,
UndoEntry,
DiffResult,
CommitInfo,
GitEvent,
GitEventListener,
} from './types.js';
import { DEFAULT_GIT_CONFIG } from './types.js';
import { GitRepo } from './repo.js';
import { AutoCommitManager } from './auto-commit.js';
import { MessageGenerator } from './message-generator.js';
import { UndoManager } from './undo-manager.js';
export class GitManager {
private repo: GitRepo;
private autoCommit: AutoCommitManager;
private messageGenerator: MessageGenerator;
private undoManager: UndoManager;
private config: GitConfig;
private workdir: string;
/** 事件监听器 */
private eventListeners: GitEventListener[] = [];
/** 是否已初始化 */
private initialized: boolean = false;
constructor(workdir: string, config?: Partial<GitConfig>) {
this.workdir = workdir;
this.config = { ...DEFAULT_GIT_CONFIG, ...config };
// 创建组件
this.repo = new GitRepo(workdir, this.config);
this.messageGenerator = new MessageGenerator(this.config.messageFormat);
this.autoCommit = new AutoCommitManager(
this.repo,
this.config.autoCommit,
this.messageGenerator
);
this.undoManager = new UndoManager(this.repo, this.config.undo);
// 设置自动提交回调
this.autoCommit.setOnCommit((result) => {
if (result.success && result.shortHash && result.message) {
// 记录到 undo 历史
const files = this.autoCommit.getPendingFiles();
this.undoManager.recordCommit(
result.hash!,
result.shortHash,
result.message,
files
);
// 发送事件
this.emitEvent('commit', {
hash: result.shortHash,
message: result.message,
auto: true,
});
}
});
}
/**
* 初始化 Git 管理器
*/
async initialize(): Promise<boolean> {
if (!this.config.enabled) {
return false;
}
const isRepo = await this.repo.initialize();
this.initialized = isRepo;
return isRepo;
}
/**
* 检查是否已初始化
*/
isInitialized(): boolean {
return this.initialized;
}
/**
* 获取仓库状态
*/
async getStatus(): Promise<GitStatus> {
return this.repo.getStatus();
}
/**
* 文件变更后调用(由工具执行器调用)
*/
async onFileChanged(
filePath: string,
changeType: 'create' | 'modify' | 'delete'
): Promise<void> {
if (!this.initialized || !this.config.enabled) {
return;
}
await this.autoCommit.onFileChanged(filePath, changeType);
}
/**
* 手动提交
*/
async commit(options: {
message?: string;
files?: string[];
all?: boolean;
} = {}): Promise<CommitResult> {
if (!this.initialized) {
return { success: false, error: 'Git not initialized' };
}
// 如果有待提交的文件,先处理
if (this.autoCommit.hasPendingFiles()) {
await this.autoCommit.flush();
}
// 获取差异用于生成消息
const diff = await this.repo.getDiff({ staged: options.all });
const message = options.message || this.messageGenerator.generate(diff);
const result = await this.repo.commit({
message,
files: options.files,
all: options.all,
aiEdits: true,
});
if (result.success && result.shortHash && result.message) {
// 记录到 undo 历史
const files = options.files || [];
this.undoManager.recordCommit(
result.hash!,
result.shortHash,
result.message,
files
);
// 发送事件
this.emitEvent('commit', {
hash: result.shortHash,
message: result.message,
auto: false,
});
}
return result;
}
/**
* 撤销上一次 AI 提交
*/
async undo(): Promise<UndoResult> {
if (!this.initialized) {
return { success: false, message: 'Git not initialized' };
}
const result = await this.undoManager.undo();
if (result.success) {
this.emitEvent('undo', {
hash: result.commitHash,
files: result.restoredFiles,
});
}
return result;
}
/**
* 获取 undo 预览
*/
getUndoPreview(): UndoEntry | null {
return this.undoManager.getUndoPreview();
}
/**
* 获取 undo 历史
*/
getUndoHistory(): UndoEntry[] {
return this.undoManager.getHistory();
}
/**
* 检查是否可以 undo
*/
async canUndo(): Promise<{ canUndo: boolean; reason?: string }> {
return this.undoManager.canUndo();
}
/**
* 获取差异
*/
async getDiff(options: {
staged?: boolean;
file?: string;
commit?: string;
} = {}): Promise<DiffResult> {
return this.repo.getDiff(options);
}
/**
* 获取最近的提交
*/
async getRecentCommits(count: number = 10): Promise<CommitInfo[]> {
return this.repo.getRecentCommits(count);
}
/**
* 获取当前分支
*/
async getCurrentBranch(): Promise<string> {
return this.repo.getCurrentBranch();
}
/**
* 强制刷新待提交
*/
async flushPendingCommits(): Promise<CommitResult | null> {
return this.autoCommit.flush();
}
/**
* 获取待提交文件
*/
getPendingFiles(): string[] {
return this.autoCommit.getPendingFiles();
}
/**
* 添加事件监听器
*/
addEventListener(listener: GitEventListener): void {
this.eventListeners.push(listener);
}
/**
* 移除事件监听器
*/
removeEventListener(listener: GitEventListener): void {
const index = this.eventListeners.indexOf(listener);
if (index !== -1) {
this.eventListeners.splice(index, 1);
}
}
/**
* 发送事件
*/
private emitEvent(type: GitEvent['type'], data: unknown): void {
const event: GitEvent = {
type,
timestamp: Date.now(),
data,
};
for (const listener of this.eventListeners) {
try {
listener(event);
} catch (error) {
console.error('Git event listener error:', error);
}
}
}
/**
* 获取配置
*/
getConfig(): GitConfig {
return { ...this.config };
}
/**
* 更新配置
*/
updateConfig(config: Partial<GitConfig>): void {
this.config = { ...this.config, ...config };
}
}
// 全局 Git 管理器实例
let gitManager: GitManager | null = null;
/**
* 获取全局 Git 管理器
*/
export function getGitManager(): GitManager | null {
return gitManager;
}
/**
* 初始化全局 Git 管理器
*/
export async function initGitManager(
workdir: string,
config?: Partial<GitConfig>
): Promise<GitManager | null> {
gitManager = new GitManager(workdir, config);
const initialized = await gitManager.initialize();
if (!initialized) {
gitManager = null;
return null;
}
return gitManager;
}
/**
* 重置全局 Git 管理器
*/
export function resetGitManager(): void {
gitManager = null;
}
+276
View File
@@ -0,0 +1,276 @@
/**
* Commit Message 生成器
*
* 参考 aider 的 message_generator 实现
* 支持 conventional、simple、detailed 三种格式
*/
import * as path from 'path';
import type { DiffResult, MessageFormatConfig, FileDiff } from './types.js';
export class MessageGenerator {
private config: MessageFormatConfig;
constructor(config: MessageFormatConfig) {
this.config = config;
}
/**
* 生成 commit message
*/
generate(diff: DiffResult, files?: string[]): string {
const fileList = files || diff.files.map((f) => f.path);
switch (this.config.style) {
case 'conventional':
return this.generateConventional(diff, fileList);
case 'simple':
return this.generateSimple(diff, fileList);
case 'detailed':
return this.generateDetailed(diff, fileList);
default:
return this.generateSimple(diff, fileList);
}
}
/**
* Conventional Commits 格式
* type(scope): subject
*/
private generateConventional(diff: DiffResult, files: string[]): string {
const type = this.detectChangeType(diff, files);
const scope = this.detectScope(files);
const subject = this.generateSubject(diff, files);
let message = type;
if (scope) {
message += `(${scope})`;
}
message += `: ${subject}`;
// 截断到最大长度
if (message.length > this.config.maxLength) {
message = message.slice(0, this.config.maxLength - 3) + '...';
}
// 添加文件列表
if (this.config.includeFileList && files.length <= 5) {
const fileListStr = files.map((f) => `- ${f}`).join('\n');
message += `\n\nFiles:\n${fileListStr}`;
}
return message;
}
/**
* 简单格式
* action file(s)
*/
private generateSimple(diff: DiffResult, files: string[]): string {
const action = this.detectAction(diff, files);
if (files.length === 1) {
return `${action} ${path.basename(files[0])}`;
} else if (files.length <= 3) {
const fileNames = files.map((f) => path.basename(f));
return `${action} ${fileNames.join(', ')}`;
} else {
return `${action} ${files.length} files`;
}
}
/**
* 详细格式
* type(scope): subject
*
* body with file list
*
* footer with stats
*/
private generateDetailed(diff: DiffResult, files: string[]): string {
const header = this.generateConventional(diff, files).split('\n')[0];
const body = this.generateBody(diff, files);
const footer = this.generateFooter(diff);
return [header, '', body, '', footer].filter(Boolean).join('\n');
}
/**
* 检测变更类型
*/
private detectChangeType(diff: DiffResult, files: string[]): string {
const hasNewFiles = diff.files.some((f) => f.status === 'added');
const hasDeletedFiles = diff.files.some((f) => f.status === 'deleted');
// 检查特殊文件类型
const hasTestFiles = files.some((f) =>
f.includes('test') || f.includes('spec') || f.includes('.test.')
);
const hasDocFiles = files.some((f) =>
f.match(/\.(md|txt|rst|doc)$/i)
);
const hasConfigFiles = files.some((f) =>
f.match(/(config|\.json|\.yaml|\.yml|\.toml|\.ini)$/i) ||
f.includes('package.json') ||
f.includes('tsconfig')
);
const hasStyleFiles = files.some((f) =>
f.match(/\.(css|scss|less|sass|styl)$/i)
);
// 根据文件类型判断
if (hasTestFiles && files.every((f) => f.includes('test') || f.includes('spec'))) {
return 'test';
}
if (hasDocFiles && files.every((f) => f.match(/\.(md|txt|rst|doc)$/i))) {
return 'docs';
}
if (hasConfigFiles && files.every((f) =>
f.match(/(config|\.json|\.yaml|\.yml|\.toml|\.ini)$/i) ||
f.includes('package.json')
)) {
return 'chore';
}
if (hasStyleFiles && files.every((f) => f.match(/\.(css|scss|less|sass|styl)$/i))) {
return 'style';
}
// 根据变更类型判断
if (hasNewFiles && !hasDeletedFiles) {
return 'feat';
}
if (hasDeletedFiles && !hasNewFiles) {
return 'refactor';
}
// 分析内容判断是 fix 还是 feat
const content = diff.files
.map((f) => f.hunks.map((h) => h.content).join(''))
.join('');
if (
content.toLowerCase().includes('fix') ||
content.toLowerCase().includes('bug') ||
content.toLowerCase().includes('error') ||
content.toLowerCase().includes('issue')
) {
return 'fix';
}
// 默认为 feat
return 'feat';
}
/**
* 检测范围 (scope)
*/
private detectScope(files: string[]): string | null {
if (files.length === 0) return null;
// 找共同目录
const dirs = files.map((f) => {
const parts = f.split('/');
// 排除 src 目录,取第一个有意义的目录
return parts.filter((p) => p !== 'src' && p !== '.' && !p.includes('.'))[0];
});
// 如果所有文件在同一目录下
const uniqueDirs = [...new Set(dirs.filter(Boolean))];
if (uniqueDirs.length === 1) {
return uniqueDirs[0];
}
return null;
}
/**
* 生成主题行
*/
private generateSubject(diff: DiffResult, files: string[]): string {
const action = this.detectAction(diff, files);
if (files.length === 1) {
const fileName = path.basename(files[0]);
const ext = path.extname(fileName);
const name = path.basename(fileName, ext);
return `${action} ${name}`;
}
// 找共同特征
const extensions = [...new Set(files.map((f) => path.extname(f)))];
if (extensions.length === 1 && extensions[0]) {
return `${action} ${files.length} ${extensions[0].slice(1)} files`;
}
const dirs = [...new Set(files.map((f) => f.split('/')[0]))];
if (dirs.length === 1 && dirs[0] !== '.') {
return `${action} ${dirs[0]} module`;
}
return `${action} ${files.length} files`;
}
/**
* 检测操作动词
*/
private detectAction(diff: DiffResult, files: string[]): string {
const hasAdded = diff.files.some((f) => f.status === 'added');
const hasDeleted = diff.files.some((f) => f.status === 'deleted');
const hasModified = diff.files.some((f) => f.status === 'modified');
const hasRenamed = diff.files.some((f) => f.status === 'renamed');
if (hasAdded && !hasModified && !hasDeleted) return 'add';
if (hasDeleted && !hasModified && !hasAdded) return 'remove';
if (hasRenamed) return 'rename';
if (hasModified && !hasAdded && !hasDeleted) return 'update';
return 'update';
}
/**
* 生成正文
*/
private generateBody(diff: DiffResult, files: string[]): string {
const changes: string[] = [];
for (const file of files) {
const fileDiff = diff.files.find((f) => f.path === file);
const action = fileDiff
? this.getActionWord(fileDiff.status)
: 'Modified';
changes.push(`- ${action}: ${file}`);
}
return changes.join('\n');
}
/**
* 获取操作词
*/
private getActionWord(status: string): string {
switch (status) {
case 'added':
return 'Added';
case 'deleted':
return 'Removed';
case 'renamed':
return 'Renamed';
case 'copied':
return 'Copied';
case 'modified':
default:
return 'Modified';
}
}
/**
* 生成页脚
*/
private generateFooter(diff: DiffResult): string {
const { stats } = diff;
if (stats.filesChanged === 0) {
return '';
}
return `Stats: ${stats.filesChanged} file(s), +${stats.insertions}/-${stats.deletions}`;
}
}
+534
View File
@@ -0,0 +1,534 @@
/**
* Git 仓库管理类
*
* 封装 simple-git,提供 Git 操作的基础功能
* 参考 aider 的 repo.py 实现
*/
import { simpleGit, type SimpleGit, type StatusResult } from 'simple-git';
import * as path from 'path';
import type {
GitStatus,
FileChange,
CommitInfo,
DiffResult,
FileDiff,
DiffHunk,
ChangeStatus,
BranchInfo,
GitConfig,
CommitOptions,
CommitResult,
AttributionConfig,
} from './types.js';
export class GitRepo {
private git: SimpleGit;
private workdir: string;
private config: GitConfig;
/** 当前会话中 AI 生成的提交哈希 */
private aiCommitHashes: Set<string> = new Set();
constructor(workdir: string, config: GitConfig) {
this.workdir = workdir;
this.config = config;
this.git = simpleGit(workdir);
}
/**
* 初始化并验证是否为 Git 仓库
*/
async initialize(): Promise<boolean> {
try {
const isRepo = await this.git.checkIsRepo();
return isRepo;
} catch {
return false;
}
}
/**
* 获取仓库根目录
*/
async getRoot(): Promise<string> {
const root = await this.git.revparse(['--show-toplevel']);
return root.trim();
}
/**
* 获取仓库状态
*/
async getStatus(): Promise<GitStatus> {
const status = await this.git.status();
return {
branch: status.current || 'HEAD',
ahead: status.ahead,
behind: status.behind,
staged: this.parseStatusFiles(status.staged, status),
unstaged: this.parseStatusFiles(status.modified, status),
untracked: status.not_added,
hasConflicts: status.conflicted.length > 0,
conflicts: status.conflicted,
isDirty: status.files.length > 0,
};
}
/**
* 检查文件是否为脏状态
*/
async isDirty(filePath?: string): Promise<boolean> {
const status = await this.git.status();
if (!filePath) {
return status.files.length > 0;
}
const normalizedPath = this.normalizePath(filePath);
return status.files.some(
(f) => this.normalizePath(f.path) === normalizedPath
);
}
/**
* 获取差异
*/
async getDiff(options: {
staged?: boolean;
file?: string;
commit?: string;
} = {}): Promise<DiffResult> {
let diffOutput: string;
try {
if (options.commit) {
diffOutput = await this.git.diff([`${options.commit}^`, options.commit]);
} else if (options.staged) {
const args = ['--cached'];
if (options.file) args.push('--', options.file);
diffOutput = await this.git.diff(args);
} else {
const args: string[] = [];
if (options.file) args.push('--', options.file);
diffOutput = await this.git.diff(args);
}
} catch {
diffOutput = '';
}
return this.parseDiff(diffOutput);
}
/**
* 获取两个提交之间的差异
*/
async getDiffBetween(fromCommit: string, toCommit: string): Promise<DiffResult> {
const diffOutput = await this.git.diff([fromCommit, toCommit]);
return this.parseDiff(diffOutput);
}
/**
* 暂存文件
*/
async stage(files: string | string[]): Promise<void> {
const fileList = Array.isArray(files) ? files : [files];
await this.git.add(fileList);
}
/**
* 取消暂存文件
*/
async unstage(files: string | string[]): Promise<void> {
const fileList = Array.isArray(files) ? files : [files];
await this.git.reset(['HEAD', '--', ...fileList]);
}
/**
* 提交变更
*/
async commit(options: CommitOptions = {}): Promise<CommitResult> {
try {
// 暂存文件
if (options.files && options.files.length > 0) {
await this.stage(options.files);
} else if (options.all) {
await this.git.add('.');
}
// 检查是否有可提交的内容
const status = await this.git.status();
if (status.staged.length === 0) {
return {
success: false,
error: 'Nothing to commit',
};
}
// 构建提交消息
let message = options.message || 'Update files';
// 添加属性标记
if (options.aiEdits && this.config.attribution.coAuthoredBy) {
message += `\n\nCo-authored-by: ${this.config.attribution.markerName} <noreply@ai-assistant.local>`;
}
// 设置作者信息 (如果需要标记)
const commitOptions: string[] = [];
if (options.aiEdits && this.config.attribution.attributeAuthor) {
const originalName = await this.getConfigValue('user.name');
if (originalName) {
commitOptions.push(
'--author',
`${originalName} (${this.config.attribution.markerName}) <${await this.getConfigValue('user.email')}>`
);
}
}
// 执行提交
const result = await this.git.commit(message, undefined, {
'--no-verify': null, // 跳过 hooks 以避免干扰
...Object.fromEntries(commitOptions.map((opt, i) =>
i % 2 === 0 ? [opt, commitOptions[i + 1]] : []
).filter(arr => arr.length > 0)),
});
const hash = result.commit;
const shortHash = hash.slice(0, 7);
// 记录 AI 提交
if (options.aiEdits) {
this.aiCommitHashes.add(shortHash);
}
return {
success: true,
hash,
shortHash,
message: message.split('\n')[0], // 只返回第一行
};
} catch (error) {
return {
success: false,
error: error instanceof Error ? error.message : String(error),
};
}
}
/**
* 获取提交信息
*/
async getCommitInfo(hash: string): Promise<CommitInfo | null> {
try {
const log = await this.git.log({
maxCount: 1,
from: hash,
to: hash,
});
if (log.all.length === 0) {
return null;
}
const entry = log.all[0];
const files = await this.getCommitFiles(hash);
return {
hash: entry.hash,
shortHash: entry.hash.slice(0, 7),
message: entry.message,
author: entry.author_name,
email: entry.author_email,
date: new Date(entry.date),
files,
isAIGenerated: this.aiCommitHashes.has(entry.hash.slice(0, 7)),
};
} catch {
return null;
}
}
/**
* 获取最近的提交
*/
async getRecentCommits(count: number = 10): Promise<CommitInfo[]> {
const log = await this.git.log({ maxCount: count });
const commits: CommitInfo[] = [];
for (const entry of log.all) {
const files = await this.getCommitFiles(entry.hash);
commits.push({
hash: entry.hash,
shortHash: entry.hash.slice(0, 7),
message: entry.message,
author: entry.author_name,
email: entry.author_email,
date: new Date(entry.date),
files,
isAIGenerated: this.aiCommitHashes.has(entry.hash.slice(0, 7)),
});
}
return commits;
}
/**
* 获取 HEAD 提交哈希
*/
async getHeadCommit(): Promise<string | null> {
try {
const hash = await this.git.revparse(['HEAD']);
return hash.trim();
} catch {
return null;
}
}
/**
* 获取 HEAD 短哈希
*/
async getHeadShortHash(): Promise<string | null> {
try {
const hash = await this.git.revparse(['--short', 'HEAD']);
return hash.trim();
} catch {
return null;
}
}
/**
* 检查提交是否已推送到远程
*/
async isCommitPushed(hash: string): Promise<boolean> {
try {
const branch = await this.git.revparse(['--abbrev-ref', 'HEAD']);
const remoteBranch = `origin/${branch.trim()}`;
// 检查远程分支是否存在
try {
await this.git.revparse([remoteBranch]);
} catch {
return false; // 远程分支不存在
}
// 检查提交是否在远程分支上
const result = await this.git.raw([
'branch', '-r', '--contains', hash
]);
return result.includes(remoteBranch);
} catch {
return false;
}
}
/**
* 软重置到指定提交
*/
async resetSoft(commit: string): Promise<void> {
await this.git.reset(['--soft', commit]);
}
/**
* 检出文件
*/
async checkoutFile(commit: string, filePath: string): Promise<void> {
await this.git.checkout([commit, '--', filePath]);
}
/**
* 获取分支列表
*/
async getBranches(): Promise<BranchInfo[]> {
const summary = await this.git.branch(['-vv']);
const branches: BranchInfo[] = [];
for (const [name, data] of Object.entries(summary.branches)) {
branches.push({
name,
current: data.current,
remote: undefined, // simple-git 不直接提供
upstream: undefined,
ahead: 0,
behind: 0,
});
}
return branches;
}
/**
* 获取当前分支名
*/
async getCurrentBranch(): Promise<string> {
const branch = await this.git.revparse(['--abbrev-ref', 'HEAD']);
return branch.trim();
}
/**
* 获取提交涉及的文件
*/
private async getCommitFiles(hash: string): Promise<string[]> {
try {
const result = await this.git.raw([
'diff-tree', '--no-commit-id', '--name-only', '-r', hash
]);
return result.trim().split('\n').filter(Boolean);
} catch {
return [];
}
}
/**
* 获取 Git 配置值
*/
private async getConfigValue(key: string): Promise<string | null> {
try {
const value = await this.git.getConfig(key);
return value.value || null;
} catch {
return null;
}
}
/**
* 解析状态文件
*/
private parseStatusFiles(files: string[], status: StatusResult): FileChange[] {
return files.map((file) => {
const fileStatus = status.files.find((f) => f.path === file);
return {
path: file,
status: this.mapStatus(fileStatus?.index || 'M'),
additions: 0,
deletions: 0,
};
});
}
/**
* 映射状态字符到 ChangeStatus
*/
private mapStatus(code: string): ChangeStatus {
switch (code) {
case 'A':
case '?':
return 'added';
case 'D':
return 'deleted';
case 'R':
return 'renamed';
case 'C':
return 'copied';
case 'M':
default:
return 'modified';
}
}
/**
* 解析 diff 输出
*/
private parseDiff(diffOutput: string): DiffResult {
const files: FileDiff[] = [];
let totalInsertions = 0;
let totalDeletions = 0;
if (!diffOutput.trim()) {
return { files, stats: { filesChanged: 0, insertions: 0, deletions: 0 } };
}
const fileSections = diffOutput.split(/(?=diff --git)/);
for (const section of fileSections) {
if (!section.trim()) continue;
const pathMatch = section.match(/diff --git a\/(.*) b\/(.*)/);
if (!pathMatch) continue;
const hunks: DiffHunk[] = [];
const hunkMatches = section.matchAll(
/@@ -(\d+),?(\d*) \+(\d+),?(\d*) @@([\s\S]*?)(?=@@|$)/g
);
for (const match of hunkMatches) {
const content = match[5] || '';
const adds = (content.match(/^\+[^+]/gm) || []).length;
const dels = (content.match(/^-[^-]/gm) || []).length;
totalInsertions += adds;
totalDeletions += dels;
hunks.push({
oldStart: parseInt(match[1]),
oldLines: parseInt(match[2]) || 1,
newStart: parseInt(match[3]),
newLines: parseInt(match[4]) || 1,
content,
});
}
files.push({
path: pathMatch[2],
oldPath: pathMatch[1] !== pathMatch[2] ? pathMatch[1] : undefined,
status: this.detectDiffStatus(section),
hunks,
binary: section.includes('Binary files'),
});
}
return {
files,
stats: {
filesChanged: files.length,
insertions: totalInsertions,
deletions: totalDeletions,
},
};
}
/**
* 检测 diff 中的变更状态
*/
private detectDiffStatus(diffSection: string): ChangeStatus {
if (diffSection.includes('new file mode')) return 'added';
if (diffSection.includes('deleted file mode')) return 'deleted';
if (diffSection.includes('rename from')) return 'renamed';
if (diffSection.includes('copy from')) return 'copied';
return 'modified';
}
/**
* 规范化文件路径
*/
private normalizePath(filePath: string): string {
return path.normalize(filePath).replace(/\\/g, '/');
}
/**
* 获取 AI 提交哈希集合
*/
getAICommitHashes(): Set<string> {
return this.aiCommitHashes;
}
/**
* 检查提交是否为 AI 生成
*/
isAICommit(shortHash: string): boolean {
return this.aiCommitHashes.has(shortHash);
}
/**
* 添加 AI 提交哈希
*/
addAICommitHash(shortHash: string): void {
this.aiCommitHashes.add(shortHash);
}
/**
* 移除 AI 提交哈希
*/
removeAICommitHash(shortHash: string): void {
this.aiCommitHashes.delete(shortHash);
}
}
+346
View File
@@ -0,0 +1,346 @@
/**
* Git 深度集成类型定义
*
* 参考 aider 的实现
*/
/**
* Git 配置
*/
export interface GitConfig {
/** 是否启用 Git 集成 */
enabled: boolean;
/** 自动提交配置 */
autoCommit: AutoCommitConfig;
/** Undo 配置 */
undo: UndoConfig;
/** Commit message 格式配置 */
messageFormat: MessageFormatConfig;
/** 属性配置 */
attribution: AttributionConfig;
}
/**
* 自动提交配置
*/
export interface AutoCommitConfig {
/** 是否启用自动提交 */
enabled: boolean;
/** 提交模式 */
mode: 'immediate' | 'batch' | 'manual';
/** batch 模式下的延迟 (ms) */
batchDelay: number;
/** 排除的文件模式 */
excludePatterns: string[];
/** 是否在编辑前提交脏文件 */
dirtyCommits: boolean;
}
/**
* Undo 配置
*/
export interface UndoConfig {
/** 最大历史记录数 */
maxHistory: number;
}
/**
* Commit message 格式配置
*/
export interface MessageFormatConfig {
/** 格式风格 */
style: 'conventional' | 'simple' | 'detailed';
/** 是否包含文件列表 */
includeFileList: boolean;
/** 最大长度 */
maxLength: number;
/** 是否使用 AI 生成 */
useAI: boolean;
/** 提交消息语言 */
language?: string;
}
/**
* 属性配置 (作者/提交者署名)
*/
export interface AttributionConfig {
/** 是否在作者名添加标记 */
attributeAuthor: boolean;
/** 是否在提交者名添加标记 */
attributeCommitter: boolean;
/** 是否添加 Co-authored-by */
coAuthoredBy: boolean;
/** 标记名称 */
markerName: string;
}
/**
* Git 仓库状态
*/
export interface GitStatus {
/** 当前分支 */
branch: string;
/** 领先远程的提交数 */
ahead: number;
/** 落后远程的提交数 */
behind: number;
/** 暂存的文件 */
staged: FileChange[];
/** 未暂存的修改 */
unstaged: FileChange[];
/** 未跟踪的文件 */
untracked: string[];
/** 是否有冲突 */
hasConflicts: boolean;
/** 冲突文件 */
conflicts: string[];
/** 是否为脏状态 */
isDirty: boolean;
}
/**
* 文件变更
*/
export interface FileChange {
/** 文件路径 */
path: string;
/** 变更状态 */
status: ChangeStatus;
/** 添加行数 */
additions: number;
/** 删除行数 */
deletions: number;
/** 原路径 (重命名时) */
oldPath?: string;
}
/**
* 变更状态
*/
export type ChangeStatus =
| 'added'
| 'modified'
| 'deleted'
| 'renamed'
| 'copied'
| 'untracked';
/**
* 提交信息
*/
export interface CommitInfo {
/** 完整哈希 */
hash: string;
/** 短哈希 */
shortHash: string;
/** 提交消息 */
message: string;
/** 作者名 */
author: string;
/** 作者邮箱 */
email: string;
/** 提交时间 */
date: Date;
/** 涉及的文件 */
files: string[];
/** 是否由 AI 生成 */
isAIGenerated: boolean;
}
/**
* Diff 结果
*/
export interface DiffResult {
/** 文件差异列表 */
files: FileDiff[];
/** 统计信息 */
stats: DiffStats;
}
/**
* 文件差异
*/
export interface FileDiff {
/** 文件路径 */
path: string;
/** 原路径 */
oldPath?: string;
/** 变更状态 */
status: ChangeStatus;
/** 差异块 */
hunks: DiffHunk[];
/** 是否为二进制文件 */
binary: boolean;
}
/**
* 差异块
*/
export interface DiffHunk {
/** 旧文件起始行 */
oldStart: number;
/** 旧文件行数 */
oldLines: number;
/** 新文件起始行 */
newStart: number;
/** 新文件行数 */
newLines: number;
/** 差异内容 */
content: string;
}
/**
* 差异统计
*/
export interface DiffStats {
/** 变更文件数 */
filesChanged: number;
/** 插入行数 */
insertions: number;
/** 删除行数 */
deletions: number;
}
/**
* Undo 历史条目
*/
export interface UndoEntry {
/** 条目 ID */
id: string;
/** 时间戳 */
timestamp: number;
/** 提交哈希 */
commitHash: string;
/** 提交消息 */
message: string;
/** 涉及的文件 */
files: string[];
/** 是否可以撤销 */
canUndo: boolean;
}
/**
* Undo 操作结果
*/
export interface UndoResult {
/** 是否成功 */
success: boolean;
/** 结果消息 */
message: string;
/** 撤销的提交哈希 */
commitHash?: string;
/** 恢复的文件 */
restoredFiles?: string[];
}
/**
* 提交选项
*/
export interface CommitOptions {
/** 提交消息 (可选,不填则自动生成) */
message?: string;
/** 要提交的文件 (可选,不填则提交所有暂存文件) */
files?: string[];
/** 是否提交所有变更 */
all?: boolean;
/** 是否为 AI 编辑 */
aiEdits?: boolean;
/** 上下文信息 (用于生成消息) */
context?: string;
}
/**
* 提交结果
*/
export interface CommitResult {
/** 是否成功 */
success: boolean;
/** 提交哈希 */
hash?: string;
/** 短哈希 */
shortHash?: string;
/** 提交消息 */
message?: string;
/** 错误信息 */
error?: string;
}
/**
* 分支信息
*/
export interface BranchInfo {
/** 分支名 */
name: string;
/** 是否为当前分支 */
current: boolean;
/** 远程名称 */
remote?: string;
/** 上游分支 */
upstream?: string;
/** 领先远程的提交数 */
ahead: number;
/** 落后远程的提交数 */
behind: number;
}
/**
* Git 事件类型
*/
export type GitEventType =
| 'commit'
| 'undo'
| 'file_staged'
| 'file_unstaged'
| 'branch_switch'
| 'pull'
| 'push';
/**
* Git 事件
*/
export interface GitEvent {
type: GitEventType;
timestamp: number;
data: unknown;
}
/**
* Git 事件监听器
*/
export type GitEventListener = (event: GitEvent) => void;
/**
* 默认配置
*/
export const DEFAULT_GIT_CONFIG: GitConfig = {
enabled: true,
autoCommit: {
enabled: true,
mode: 'batch',
batchDelay: 3000,
excludePatterns: [
'*.log',
'node_modules/**',
'.git/**',
'.ai-assistant/**',
'*.tmp',
'*.swp',
],
dirtyCommits: true,
},
undo: {
maxHistory: 50,
},
messageFormat: {
style: 'conventional',
includeFileList: true,
maxLength: 72,
useAI: false,
},
attribution: {
attributeAuthor: true,
attributeCommitter: true,
coAuthoredBy: false,
markerName: 'ai-assistant',
},
};
+194
View File
@@ -0,0 +1,194 @@
/**
* Undo 管理器
*
* 参考 aider 的 undo 实现
* 仅支持撤销 AI 生成的提交,且需要通过多项安全检查
*/
import type { GitRepo } from './repo.js';
import type { UndoConfig, UndoEntry, UndoResult, CommitInfo } from './types.js';
export class UndoManager {
private repo: GitRepo;
private config: UndoConfig;
/** Undo 历史 */
private history: UndoEntry[] = [];
constructor(repo: GitRepo, config: UndoConfig) {
this.repo = repo;
this.config = config;
}
/**
* 记录提交到 undo 历史
*/
recordCommit(commitHash: string, shortHash: string, message: string, files: string[]): void {
const entry: UndoEntry = {
id: `undo-${Date.now()}`,
timestamp: Date.now(),
commitHash: shortHash,
message,
files,
canUndo: true,
};
this.history.push(entry);
// 限制历史记录数量
if (this.history.length > this.config.maxHistory) {
this.history = this.history.slice(-this.config.maxHistory);
}
}
/**
* 执行 undo 操作
*
* 安全检查清单(参考 aider):
* 1. 提交是否由 AI 生成
* 2. 提交是否已推送到远程
* 3. 是否为合并提交
* 4. 文件是否有未提交的更改
*/
async undo(): Promise<UndoResult> {
// 1. 检查是否有可撤销的提交
if (this.history.length === 0) {
return {
success: false,
message: 'Nothing to undo. No AI commits in this session.',
};
}
// 2. 获取最后一条记录
const lastEntry = this.history[this.history.length - 1];
// 3. 获取 HEAD 提交
const headShortHash = await this.repo.getHeadShortHash();
if (!headShortHash) {
return {
success: false,
message: 'Unable to get HEAD commit.',
};
}
// 4. 验证最后的提交是否与记录匹配
if (headShortHash !== lastEntry.commitHash) {
return {
success: false,
message: `The last commit (${headShortHash}) does not match the recorded AI commit (${lastEntry.commitHash}). The repository may have changed since the AI edit.`,
};
}
// 5. 验证是否为 AI 生成的提交
if (!this.repo.isAICommit(headShortHash)) {
return {
success: false,
message: `The commit ${headShortHash} was not made by AI in this session.`,
};
}
// 6. 检查提交是否已推送到远程
const headHash = await this.repo.getHeadCommit();
if (headHash) {
const isPushed = await this.repo.isCommitPushed(headHash);
if (isPushed) {
return {
success: false,
message: 'The commit has already been pushed to remote. Undo is not safe.',
};
}
}
// 7. 检查文件是否有未提交的更改
for (const file of lastEntry.files) {
const isDirty = await this.repo.isDirty(file);
if (isDirty) {
return {
success: false,
message: `The file ${file} has uncommitted changes. Please commit or stash them first.`,
};
}
}
// 8. 执行撤销操作
try {
// 恢复文件到上一个版本
const restoredFiles: string[] = [];
for (const file of lastEntry.files) {
try {
await this.repo.checkoutFile('HEAD~1', file);
restoredFiles.push(file);
} catch {
// 文件可能在前一个提交中不存在(新建的文件)
// 这种情况下跳过
}
}
// 软重置 HEAD
await this.repo.resetSoft('HEAD~1');
// 从历史中移除
this.history.pop();
// 从 AI 提交记录中移除
this.repo.removeAICommitHash(lastEntry.commitHash);
return {
success: true,
message: `Undone: ${lastEntry.commitHash} - ${lastEntry.message}`,
commitHash: lastEntry.commitHash,
restoredFiles,
};
} catch (error) {
return {
success: false,
message: `Undo failed: ${error instanceof Error ? error.message : String(error)}`,
};
}
}
/**
* 预览将要撤销的内容
*/
getUndoPreview(): UndoEntry | null {
if (this.history.length === 0) {
return null;
}
return this.history[this.history.length - 1];
}
/**
* 获取 undo 历史
*/
getHistory(): UndoEntry[] {
return [...this.history];
}
/**
* 清除历史
*/
clearHistory(): void {
this.history = [];
}
/**
* 检查是否可以 undo
*/
async canUndo(): Promise<{ canUndo: boolean; reason?: string }> {
if (this.history.length === 0) {
return { canUndo: false, reason: 'No AI commits to undo' };
}
const lastEntry = this.history[this.history.length - 1];
const headShortHash = await this.repo.getHeadShortHash();
if (headShortHash !== lastEntry.commitHash) {
return {
canUndo: false,
reason: 'Repository has changed since last AI commit',
};
}
return { canUndo: true };
}
}
+232
View File
@@ -0,0 +1,232 @@
/**
* Hook 配置加载器
*
* 从项目配置文件加载 hook 配置
* 支持 .ai-assistant.json, .ai-assistant.jsonc, ai-assistant.config.json 等格式
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import type { HookConfig, ShellCommandConfig, FileHookConfig } from './types.js';
// 支持的配置文件名
const CONFIG_FILE_NAMES = [
'.ai-assistant.json',
'.ai-assistant.jsonc',
'ai-assistant.config.json',
'.ai-assistantrc',
'.ai-assistantrc.json',
];
/**
* 完整的配置文件结构
*/
export interface ProjectConfig {
/** Hook 配置 */
hooks?: HookConfig;
/** 插件列表 */
plugins?: string[];
/** 其他配置... */
[key: string]: unknown;
}
/**
* 移除 JSON 中的注释(支持 JSONC 格式)
*/
function stripJsonComments(jsonString: string): string {
// 移除单行注释 // ...
let result = jsonString.replace(/\/\/.*$/gm, '');
// 移除多行注释 /* ... */
result = result.replace(/\/\*[\s\S]*?\*\//g, '');
return result;
}
/**
* 解析 JSON 文件(支持 JSONC)
*/
async function parseJsonFile(filePath: string): Promise<ProjectConfig | null> {
try {
const content = await fs.readFile(filePath, 'utf-8');
const cleanContent = stripJsonComments(content);
return JSON.parse(cleanContent);
} catch {
return null;
}
}
/**
* 在目录中查找配置文件
*/
async function findConfigFile(directory: string): Promise<string | null> {
for (const fileName of CONFIG_FILE_NAMES) {
const filePath = path.join(directory, fileName);
try {
await fs.access(filePath);
return filePath;
} catch {
// 文件不存在,继续查找
}
}
return null;
}
/**
* 验证 ShellCommandConfig
*/
function validateShellCommandConfig(config: unknown): config is ShellCommandConfig {
if (typeof config !== 'object' || config === null) return false;
const obj = config as Record<string, unknown>;
// command 必须是非空字符串数组
if (!Array.isArray(obj.command) || obj.command.length === 0) return false;
if (!obj.command.every((c) => typeof c === 'string')) return false;
// environment 如果存在,必须是对象
if (obj.environment !== undefined) {
if (typeof obj.environment !== 'object' || obj.environment === null) return false;
const env = obj.environment as Record<string, unknown>;
if (!Object.values(env).every((v) => typeof v === 'string')) return false;
}
// timeout 如果存在,必须是正数
if (obj.timeout !== undefined) {
if (typeof obj.timeout !== 'number' || obj.timeout <= 0) return false;
}
// cwd 如果存在,必须是字符串
if (obj.cwd !== undefined) {
if (typeof obj.cwd !== 'string') return false;
}
return true;
}
/**
* 验证 FileHookConfig
*/
function validateFileHookConfig(config: unknown): config is FileHookConfig {
if (typeof config !== 'object' || config === null) return false;
const obj = config as Record<string, unknown>;
for (const [pattern, commands] of Object.entries(obj)) {
// pattern 必须是非空字符串
if (typeof pattern !== 'string' || pattern.length === 0) return false;
// commands 必须是 ShellCommandConfig 数组
if (!Array.isArray(commands)) return false;
if (!commands.every(validateShellCommandConfig)) return false;
}
return true;
}
/**
* 验证 HookConfig
*/
function validateHookConfig(config: unknown): config is HookConfig {
if (typeof config !== 'object' || config === null) return false;
const obj = config as Record<string, unknown>;
// file_edited
if (obj.file_edited !== undefined) {
if (!validateFileHookConfig(obj.file_edited)) return false;
}
// file_created
if (obj.file_created !== undefined) {
if (!validateFileHookConfig(obj.file_created)) return false;
}
// file_deleted
if (obj.file_deleted !== undefined) {
if (!validateFileHookConfig(obj.file_deleted)) return false;
}
// session_completed
if (obj.session_completed !== undefined) {
if (!Array.isArray(obj.session_completed)) return false;
if (!obj.session_completed.every(validateShellCommandConfig)) return false;
}
return true;
}
/**
* 加载项目配置
*/
export async function loadProjectConfig(directory: string): Promise<ProjectConfig | null> {
const configPath = await findConfigFile(directory);
if (!configPath) return null;
const config = await parseJsonFile(configPath);
return config;
}
/**
* 加载 Hook 配置
*/
export async function loadHookConfig(directory: string): Promise<HookConfig | null> {
const projectConfig = await loadProjectConfig(directory);
if (!projectConfig?.hooks) return null;
// 验证配置
if (!validateHookConfig(projectConfig.hooks)) {
console.warn('Invalid hook configuration in project config file');
return null;
}
return projectConfig.hooks;
}
/**
* 加载插件列表
*/
export async function loadPluginList(directory: string): Promise<string[]> {
const projectConfig = await loadProjectConfig(directory);
if (!projectConfig?.plugins) return [];
// 验证插件列表
if (!Array.isArray(projectConfig.plugins)) return [];
if (!projectConfig.plugins.every((p) => typeof p === 'string')) return [];
return projectConfig.plugins;
}
/**
* 创建默认配置文件
*/
export async function createDefaultConfig(directory: string): Promise<void> {
const configPath = path.join(directory, '.ai-assistant.json');
const defaultConfig: ProjectConfig = {
hooks: {
file_edited: {
'*.ts': [
{
command: ['npx', 'tsc', '--noEmit'],
timeout: 30000,
},
],
'*.{js,jsx,ts,tsx}': [
{
command: ['npx', 'eslint', '--fix'],
timeout: 30000,
},
],
},
file_created: {},
file_deleted: {},
session_completed: [],
},
plugins: [],
};
await fs.writeFile(configPath, JSON.stringify(defaultConfig, null, 2));
}
/**
* 获取配置文件路径(如果存在)
*/
export async function getConfigFilePath(directory: string): Promise<string | null> {
return findConfigFile(directory);
}
+48
View File
@@ -0,0 +1,48 @@
/**
* Hook 系统模块
*
* 提供工具执行前后的 hook 功能,支持自定义命令执行
* 参考 open-code 的实现
*/
// Hook 管理器
export {
HookManager,
getHookManager,
initHookManager,
resetHookManager,
} from './manager.js';
// 配置加载
export {
loadProjectConfig,
loadHookConfig,
loadPluginList,
createDefaultConfig,
getConfigFilePath,
type ProjectConfig,
} from './config-loader.js';
// 类型导出
export type {
HookType,
HookConfig,
HookEvent,
HookEventListener,
ShellCommandConfig,
FileHookConfig,
Hooks,
Plugin,
PluginInput,
ToolExecuteBeforeInput,
ToolExecuteBeforeOutput,
ToolExecuteAfterInput,
ToolExecuteAfterOutput,
SessionStartInput,
SessionEndInput,
MessageBeforeInput,
MessageBeforeOutput,
MessageAfterInput,
FileChangeInput,
FileChangeOutput,
} from './types.js';
+495
View File
@@ -0,0 +1,495 @@
/**
* Hook 管理器
*
* 负责 hook 的注册、触发和管理
*/
import { spawn } from 'child_process';
import { minimatch } from 'minimatch';
import type {
Hooks,
HookType,
HookConfig,
HookEvent,
HookEventListener,
ShellCommandConfig,
FileHookConfig,
ToolExecuteBeforeInput,
ToolExecuteBeforeOutput,
ToolExecuteAfterInput,
ToolExecuteAfterOutput,
SessionStartInput,
SessionEndInput,
MessageBeforeInput,
MessageBeforeOutput,
MessageAfterInput,
FileChangeInput,
FileChangeOutput,
Plugin,
PluginInput,
} from './types.js';
/**
* Hook 管理器
*/
export class HookManager {
/** 已注册的 hooks */
private hooks: Hooks[] = [];
/** 配置型 hooks(从配置文件加载) */
private configHooks: HookConfig | null = null;
/** 事件监听器 */
private eventListeners: HookEventListener[] = [];
/** 当前工作目录 */
private workdir: string;
/** 会话 ID */
private sessionId: string;
constructor(workdir: string, sessionId?: string) {
this.workdir = workdir;
this.sessionId = sessionId || 'default';
}
/**
* 注册插件
*/
async registerPlugin(plugin: Plugin): Promise<void> {
const input: PluginInput = {
workdir: this.workdir,
sessionId: this.sessionId,
};
try {
const hooks = await plugin(input);
this.hooks.push(hooks);
} catch (error) {
console.error('Failed to register plugin:', error);
}
}
/**
* 注册 hooks 对象
*/
registerHooks(hooks: Hooks): void {
this.hooks.push(hooks);
}
/**
* 设置配置型 hooks
*/
setConfigHooks(config: HookConfig): void {
this.configHooks = config;
}
/**
* 添加事件监听器
*/
addEventListener(listener: HookEventListener): void {
this.eventListeners.push(listener);
}
/**
* 移除事件监听器
*/
removeEventListener(listener: HookEventListener): void {
const index = this.eventListeners.indexOf(listener);
if (index !== -1) {
this.eventListeners.splice(index, 1);
}
}
/**
* 发送事件
*/
private emitEvent(type: HookType, data: unknown): void {
const event: HookEvent = {
type,
timestamp: Date.now(),
data,
};
for (const listener of this.eventListeners) {
try {
listener(event);
} catch (error) {
console.error('Event listener error:', error);
}
}
}
/**
* 触发工具执行前 hook
*/
async triggerToolExecuteBefore(
input: ToolExecuteBeforeInput
): Promise<ToolExecuteBeforeOutput> {
const output: ToolExecuteBeforeOutput = {
args: { ...input.args },
};
for (const hook of this.hooks) {
if (hook['tool.execute.before']) {
try {
await hook['tool.execute.before'](input, output);
} catch (error) {
console.error('Hook tool.execute.before error:', error);
}
}
}
this.emitEvent('tool.execute.before', { input, output });
return output;
}
/**
* 触发工具执行后 hook
*/
async triggerToolExecuteAfter(
input: ToolExecuteAfterInput,
result: ToolExecuteAfterOutput['result']
): Promise<ToolExecuteAfterOutput> {
const output: ToolExecuteAfterOutput = {
result: { ...result },
};
for (const hook of this.hooks) {
if (hook['tool.execute.after']) {
try {
await hook['tool.execute.after'](input, output);
} catch (error) {
console.error('Hook tool.execute.after error:', error);
}
}
}
this.emitEvent('tool.execute.after', { input, output });
return output;
}
/**
* 触发会话开始 hook
*/
async triggerSessionStart(input: SessionStartInput): Promise<void> {
for (const hook of this.hooks) {
if (hook['session.start']) {
try {
await hook['session.start'](input);
} catch (error) {
console.error('Hook session.start error:', error);
}
}
}
this.emitEvent('session.start', input);
}
/**
* 触发会话结束 hook
*/
async triggerSessionEnd(input: SessionEndInput): Promise<void> {
for (const hook of this.hooks) {
if (hook['session.end']) {
try {
await hook['session.end'](input);
} catch (error) {
console.error('Hook session.end error:', error);
}
}
}
// 执行配置型 session_completed hooks
if (this.configHooks?.session_completed) {
await this.executeShellCommands(this.configHooks.session_completed);
}
this.emitEvent('session.end', input);
}
/**
* 触发消息前 hook
*/
async triggerMessageBefore(
input: MessageBeforeInput
): Promise<MessageBeforeOutput> {
const output: MessageBeforeOutput = {
content: input.content,
};
for (const hook of this.hooks) {
if (hook['message.before']) {
try {
await hook['message.before'](input, output);
} catch (error) {
console.error('Hook message.before error:', error);
}
}
}
this.emitEvent('message.before', { input, output });
return output;
}
/**
* 触发消息后 hook
*/
async triggerMessageAfter(input: MessageAfterInput): Promise<void> {
for (const hook of this.hooks) {
if (hook['message.after']) {
try {
await hook['message.after'](input);
} catch (error) {
console.error('Hook message.after error:', error);
}
}
}
this.emitEvent('message.after', input);
}
/**
* 触发文件编辑 hook
*/
async triggerFileEdited(input: FileChangeInput): Promise<FileChangeOutput> {
const output: FileChangeOutput = {};
// 执行插件 hooks
for (const hook of this.hooks) {
if (hook['file.edited']) {
try {
await hook['file.edited'](input, output);
} catch (error) {
console.error('Hook file.edited error:', error);
}
}
}
// 执行配置型 hooks
if (this.configHooks?.file_edited) {
const results = await this.executeFileHooks(
input.path,
this.configHooks.file_edited
);
output.commandResults = results;
}
this.emitEvent('file.edited', { input, output });
return output;
}
/**
* 触发文件创建 hook
*/
async triggerFileCreated(input: FileChangeInput): Promise<FileChangeOutput> {
const output: FileChangeOutput = {};
// 执行插件 hooks
for (const hook of this.hooks) {
if (hook['file.created']) {
try {
await hook['file.created'](input, output);
} catch (error) {
console.error('Hook file.created error:', error);
}
}
}
// 执行配置型 hooks
if (this.configHooks?.file_created) {
const results = await this.executeFileHooks(
input.path,
this.configHooks.file_created
);
output.commandResults = results;
}
this.emitEvent('file.created', { input, output });
return output;
}
/**
* 触发文件删除 hook
*/
async triggerFileDeleted(input: FileChangeInput): Promise<FileChangeOutput> {
const output: FileChangeOutput = {};
// 执行插件 hooks
for (const hook of this.hooks) {
if (hook['file.deleted']) {
try {
await hook['file.deleted'](input, output);
} catch (error) {
console.error('Hook file.deleted error:', error);
}
}
}
// 执行配置型 hooks
if (this.configHooks?.file_deleted) {
const results = await this.executeFileHooks(
input.path,
this.configHooks.file_deleted
);
output.commandResults = results;
}
this.emitEvent('file.deleted', { input, output });
return output;
}
/**
* 执行文件 hooks
* 根据文件路径匹配 glob 模式并执行对应命令
*/
private async executeFileHooks(
filePath: string,
config: FileHookConfig
): Promise<FileChangeOutput['commandResults']> {
const results: FileChangeOutput['commandResults'] = [];
for (const [pattern, commands] of Object.entries(config)) {
// 使用 minimatch 进行 glob 匹配
if (minimatch(filePath, pattern, { matchBase: true })) {
const commandResults = await this.executeShellCommands(commands, {
FILE_PATH: filePath,
});
results.push(...commandResults);
}
}
return results;
}
/**
* 执行 shell 命令列表
*/
private async executeShellCommands(
commands: ShellCommandConfig[],
extraEnv?: Record<string, string>
): Promise<Array<{ command: string[]; success: boolean; output?: string; error?: string }>> {
const results: Array<{
command: string[];
success: boolean;
output?: string;
error?: string;
}> = [];
for (const cmdConfig of commands) {
const result = await this.executeShellCommand(cmdConfig, extraEnv);
results.push(result);
}
return results;
}
/**
* 执行单个 shell 命令
*/
private executeShellCommand(
config: ShellCommandConfig,
extraEnv?: Record<string, string>
): Promise<{ command: string[]; success: boolean; output?: string; error?: string }> {
return new Promise((resolve) => {
const [cmd, ...args] = config.command;
const timeout = config.timeout || 30000;
const cwd = config.cwd || this.workdir;
const env = {
...process.env,
...config.environment,
...extraEnv,
};
let stdout = '';
let stderr = '';
const child = spawn(cmd, args, {
cwd,
env,
shell: true,
});
const timer = setTimeout(() => {
child.kill('SIGTERM');
resolve({
command: config.command,
success: false,
error: `Command timed out after ${timeout}ms`,
});
}, timeout);
child.stdout?.on('data', (data) => {
stdout += data.toString();
});
child.stderr?.on('data', (data) => {
stderr += data.toString();
});
child.on('close', (code) => {
clearTimeout(timer);
resolve({
command: config.command,
success: code === 0,
output: stdout.trim() || undefined,
error: stderr.trim() || undefined,
});
});
child.on('error', (error) => {
clearTimeout(timer);
resolve({
command: config.command,
success: false,
error: error.message,
});
});
});
}
/**
* 获取所有已注册的 hooks 数量
*/
getHookCount(): number {
return this.hooks.length;
}
/**
* 清空所有 hooks
*/
clear(): void {
this.hooks = [];
this.configHooks = null;
this.eventListeners = [];
}
}
// 全局 Hook 管理器实例
let globalHookManager: HookManager | null = null;
/**
* 获取全局 Hook 管理器
*/
export function getHookManager(): HookManager | null {
return globalHookManager;
}
/**
* 初始化全局 Hook 管理器
*/
export function initHookManager(workdir: string, sessionId?: string): HookManager {
globalHookManager = new HookManager(workdir, sessionId);
return globalHookManager;
}
/**
* 重置全局 Hook 管理器
*/
export function resetHookManager(): void {
if (globalHookManager) {
globalHookManager.clear();
globalHookManager = null;
}
}
+215
View File
@@ -0,0 +1,215 @@
/**
* Hook 系统类型定义
*
* 参考 open-code 的 hook 实现
*/
import type { Tool, ToolResult } from '../types/index.js';
/**
* Hook 类型枚举
*/
export type HookType =
| 'tool.execute.before' // 工具执行前
| 'tool.execute.after' // 工具执行后
| 'session.start' // 会话开始
| 'session.end' // 会话结束
| 'message.before' // 消息发送前
| 'message.after' // 消息接收后
| 'file.edited' // 文件被编辑后
| 'file.created' // 文件被创建后
| 'file.deleted'; // 文件被删除后
/**
* 工具执行前 Hook 的输入
*/
export interface ToolExecuteBeforeInput {
tool: string;
sessionId: string;
callId: string;
args: Record<string, unknown>;
}
/**
* 工具执行前 Hook 的输出(可修改)
*/
export interface ToolExecuteBeforeOutput {
args: Record<string, unknown>;
/** 设为 true 可阻止工具执行 */
skip?: boolean;
/** 跳过时返回的结果 */
skipResult?: ToolResult;
}
/**
* 工具执行后 Hook 的输入
*/
export interface ToolExecuteAfterInput {
tool: string;
sessionId: string;
callId: string;
args: Record<string, unknown>;
duration: number; // 执行时长(毫秒)
}
/**
* 工具执行后 Hook 的输出(可修改)
*/
export interface ToolExecuteAfterOutput {
result: ToolResult;
}
/**
* 会话开始 Hook 的输入
*/
export interface SessionStartInput {
sessionId: string;
workdir: string;
}
/**
* 会话结束 Hook 的输入
*/
export interface SessionEndInput {
sessionId: string;
messageCount: number;
duration: number; // 会话时长(毫秒)
}
/**
* 消息前 Hook 的输入
*/
export interface MessageBeforeInput {
sessionId: string;
content: string;
}
/**
* 消息前 Hook 的输出(可修改)
*/
export interface MessageBeforeOutput {
content: string;
/** 设为 true 可阻止消息发送 */
skip?: boolean;
}
/**
* 消息后 Hook 的输入
*/
export interface MessageAfterInput {
sessionId: string;
content: string;
toolCalls: number;
}
/**
* 文件变更 Hook 的输入
*/
export interface FileChangeInput {
path: string;
tool: string;
sessionId: string;
}
/**
* 文件变更 Hook 的输出
*/
export interface FileChangeOutput {
/** 执行的命令结果 */
commandResults?: Array<{
command: string[];
success: boolean;
output?: string;
error?: string;
}>;
}
/**
* Shell 命令配置
*/
export interface ShellCommandConfig {
/** 命令数组,第一个元素是命令,后面是参数 */
command: string[];
/** 环境变量 */
environment?: Record<string, string>;
/** 超时时间(毫秒),默认 30000 */
timeout?: number;
/** 工作目录,默认使用当前目录 */
cwd?: string;
}
/**
* 文件 Hook 配置
* 支持 glob 模式匹配文件
*/
export interface FileHookConfig {
/** glob 模式 -> 命令配置列表 */
[pattern: string]: ShellCommandConfig[];
}
/**
* Hook 配置
*/
export interface HookConfig {
/** 文件编辑后执行的 hook */
file_edited?: FileHookConfig;
/** 文件创建后执行的 hook */
file_created?: FileHookConfig;
/** 文件删除后执行的 hook */
file_deleted?: FileHookConfig;
/** 会话完成后执行的命令 */
session_completed?: ShellCommandConfig[];
}
/**
* Hook 函数类型
*/
export type HookFunction<Input, Output> = (
input: Input,
output: Output
) => Promise<void>;
/**
* Hook 定义接口
*/
export interface Hooks {
'tool.execute.before'?: HookFunction<ToolExecuteBeforeInput, ToolExecuteBeforeOutput>;
'tool.execute.after'?: HookFunction<ToolExecuteAfterInput, ToolExecuteAfterOutput>;
'session.start'?: (input: SessionStartInput) => Promise<void>;
'session.end'?: (input: SessionEndInput) => Promise<void>;
'message.before'?: HookFunction<MessageBeforeInput, MessageBeforeOutput>;
'message.after'?: (input: MessageAfterInput) => Promise<void>;
'file.edited'?: HookFunction<FileChangeInput, FileChangeOutput>;
'file.created'?: HookFunction<FileChangeInput, FileChangeOutput>;
'file.deleted'?: HookFunction<FileChangeInput, FileChangeOutput>;
}
/**
* 插件输入
*/
export interface PluginInput {
/** 当前工作目录 */
workdir: string;
/** 会话 ID */
sessionId?: string;
}
/**
* 插件定义
* 一个插件是一个函数,接收 PluginInput 返回 Hooks
*/
export type Plugin = (input: PluginInput) => Promise<Hooks>;
/**
* Hook 事件
*/
export interface HookEvent {
type: HookType;
timestamp: number;
data: unknown;
}
/**
* Hook 事件监听器
*/
export type HookEventListener = (event: HookEvent) => void;
+430
View File
@@ -0,0 +1,430 @@
#!/usr/bin/env node
import { Command } from 'commander';
import { Agent } from './core/agent.js';
import { TerminalUI } from './ui/terminal.js';
import { loadConfig, initConfig } from './utils/config.js';
import { toolRegistry, todoManager, initTaskContext, updateTaskDescription, updateSkillDescription } from './tools/index.js';
import { getPermissionManager, promptPermission } from './permission/index.js';
import { SessionManager } from './session/index.js';
import { agentRegistry } from './agent/index.js';
import { initLSP, shutdownLSP } from './lsp/index.js';
import { getCommandRegistry } from './commands/index.js';
import { getSkillRegistry } from './skills/index.js';
import {
printServerList,
installServer,
installAllServers,
showServerInfo,
} from './lsp/cli.js';
import {
getMCPManager,
loadMCPConfig,
createMCPToolAdapter,
} from './mcp/index.js';
const program = new Command();
// MCP 管理器实例
let mcpInitialized = false;
/**
* 初始化 MCP 系统
* 加载配置、连接服务器、注册工具
*/
async function initMCP(workdir: string): Promise<void> {
if (mcpInitialized) {
return;
}
const mcpConfig = loadMCPConfig(workdir);
// 如果没有 MCP 配置,跳过初始化
if (!mcpConfig.mcp || Object.keys(mcpConfig.mcp).length === 0) {
return;
}
const mcpManager = getMCPManager();
// 监听工具变化事件
mcpManager.on('tools:changed', () => {
registerMCPTools(mcpManager);
});
// 监听服务器事件(用于日志)
mcpManager.on('server:connected', (name) => {
console.log(`🔌 MCP 服务器已连接: ${name}`);
});
mcpManager.on('server:disconnected', (name) => {
console.log(`🔌 MCP 服务器已断开: ${name}`);
});
mcpManager.on('server:error', (name, error) => {
console.error(`❌ MCP 服务器 ${name} 错误:`, error);
});
try {
await mcpManager.initialize(mcpConfig);
registerMCPTools(mcpManager);
mcpInitialized = true;
// 显示 MCP 状态
const statuses = mcpManager.getServerStatuses();
const connected = statuses.filter((s) => s.status === 'connected');
if (connected.length > 0) {
const totalTools = connected.reduce((sum, s) => sum + s.toolCount, 0);
console.log(
`🔌 MCP: ${connected.length} 个服务器已连接,${totalTools} 个工具可用`
);
}
} catch (error) {
console.error(
'❌ MCP 初始化失败:',
error instanceof Error ? error.message : String(error)
);
}
}
/**
* 将 MCP 工具注册到工具注册表
*/
function registerMCPTools(
mcpManager: ReturnType<typeof getMCPManager>
): void {
const adapter = createMCPToolAdapter(mcpManager);
const mcpTools = mcpManager.getTools();
const adaptedTools = adapter.adaptTools(mcpTools);
// 注册到工具注册表
toolRegistry.registerAll(adaptedTools);
}
/**
* 关闭 MCP 系统
*/
async function shutdownMCP(): Promise<void> {
if (mcpInitialized) {
const mcpManager = getMCPManager();
await mcpManager.shutdown();
mcpInitialized = false;
}
}
program
.name('ai-assist')
.description('AI Terminal Assistant - 终端中的 AI 编程助手')
.version('1.0.0');
// 初始化命令
program
.command('init')
.description('初始化配置(设置 API Key 等)')
.action(async () => {
await initConfig();
});
// LSP 命令组
const lspCommand = program
.command('lsp')
.description('语言服务器管理');
lspCommand
.command('list')
.description('列出所有语言服务器及其安装状态')
.action(() => {
printServerList();
});
lspCommand
.command('install [servers...]')
.description('安装指定的语言服务器')
.option('-a, --all', '安装所有语言服务器')
.action(async (servers: string[], options: { all?: boolean }) => {
if (options.all) {
await installAllServers();
} else if (servers.length === 0) {
console.log('用法: ai-assist lsp install <server> [server2] ...');
console.log(' ai-assist lsp install --all');
console.log('\n运行 "ai-assist lsp list" 查看可用的服务器');
} else {
for (const server of servers) {
await installServer(server);
}
}
});
lspCommand
.command('info <server>')
.description('显示语言服务器详细信息')
.action((server: string) => {
showServerInfo(server);
});
// MCP 命令组
const mcpCommand = program.command('mcp').description('MCP 服务器管理');
mcpCommand
.command('list')
.description('列出所有 MCP 服务器及其状态')
.action(async () => {
const mcpConfig = loadMCPConfig(process.cwd());
if (!mcpConfig.mcp || Object.keys(mcpConfig.mcp).length === 0) {
console.log('没有配置 MCP 服务器');
console.log('\n配置方法:');
console.log(' 在 ~/.ai-assist/config.json 或 .ai-assist/config.json 中添加 mcp 配置');
console.log('\n示例:');
console.log(' {');
console.log(' "mcp": {');
console.log(' "filesystem": {');
console.log(' "type": "local",');
console.log(' "command": ["npx", "-y", "@anthropic-ai/mcp-server-filesystem", "/path/to/dir"]');
console.log(' }');
console.log(' }');
console.log(' }');
return;
}
const mcpManager = getMCPManager();
// 尝试连接以获取工具数量
try {
if (!mcpManager.isInitialized()) {
await mcpManager.initialize(mcpConfig);
}
} catch {
// 忽略连接错误,仍然显示配置的服务器
}
const statuses = mcpManager.getServerStatuses();
console.log('\nMCP 服务器列表:\n');
const statusIcons: Record<string, string> = {
connected: '✅',
connecting: '🔄',
disconnected: '⭕',
disabled: '🚫',
error: '❌',
};
for (const status of statuses) {
const icon = statusIcons[status.status] || '❓';
const toolInfo = status.toolCount > 0 ? ` (${status.toolCount} 个工具)` : '';
const errorInfo = status.error ? ` - ${status.error}` : '';
console.log(
` ${icon} ${status.name} [${status.type}] - ${status.status}${toolInfo}${errorInfo}`
);
}
console.log('');
// 关闭连接
await mcpManager.shutdown();
});
mcpCommand
.command('tools [server]')
.description('列出 MCP 服务器提供的工具')
.action(async (server?: string) => {
const mcpConfig = loadMCPConfig(process.cwd());
if (!mcpConfig.mcp || Object.keys(mcpConfig.mcp).length === 0) {
console.log('没有配置 MCP 服务器');
return;
}
const mcpManager = getMCPManager();
try {
if (!mcpManager.isInitialized()) {
await mcpManager.initialize(mcpConfig);
}
const tools = mcpManager.getTools();
if (tools.length === 0) {
console.log('没有可用的 MCP 工具');
return;
}
// 按服务器分组
const toolsByServer = new Map<string, typeof tools>();
for (const tool of tools) {
if (server && tool.server !== server) {
continue;
}
const serverTools = toolsByServer.get(tool.server) || [];
serverTools.push(tool);
toolsByServer.set(tool.server, serverTools);
}
if (toolsByServer.size === 0) {
console.log(server ? `服务器 "${server}" 没有提供工具` : '没有可用的工具');
return;
}
console.log('\nMCP 工具列表:\n');
for (const [serverName, serverTools] of toolsByServer) {
console.log(`📦 ${serverName}:`);
for (const tool of serverTools) {
console.log(` ${tool.name}`);
if (tool.description) {
console.log(` ${tool.description.substring(0, 80)}${tool.description.length > 80 ? '...' : ''}`);
}
}
console.log('');
}
} catch (error) {
console.error(
'获取工具列表失败:',
error instanceof Error ? error.message : String(error)
);
} finally {
await mcpManager.shutdown();
}
});
mcpCommand
.command('test <server>')
.description('测试 MCP 服务器连接')
.action(async (server: string) => {
const mcpConfig = loadMCPConfig(process.cwd());
if (!mcpConfig.mcp?.[server]) {
console.log(`❌ 未找到服务器配置: ${server}`);
return;
}
console.log(`🔄 正在连接 ${server}...`);
const mcpManager = getMCPManager();
try {
await mcpManager.initialize({
mcp: { [server]: mcpConfig.mcp[server] },
tools: mcpConfig.tools,
});
const status = mcpManager.getServerStatus(server);
if (status?.status === 'connected') {
console.log(`✅ 连接成功!`);
console.log(` 工具数量: ${status.toolCount}`);
const tools = mcpManager.getTools();
if (tools.length > 0) {
console.log(' 可用工具:');
for (const tool of tools) {
console.log(` - ${tool.originalName}`);
}
}
} else {
console.log(`❌ 连接失败: ${status?.error || '未知错误'}`);
}
} catch (error) {
console.error(
`❌ 连接失败:`,
error instanceof Error ? error.message : String(error)
);
} finally {
await mcpManager.shutdown();
}
});
// 初始化权限系统
function setupPermissions(): void {
const permissionManager = getPermissionManager();
permissionManager.setAskCallback(promptPermission);
}
// 单次查询命令
program
.command('ask <question>')
.description('单次提问(不进入交互模式)')
.action(async (question: string) => {
setupPermissions();
const config = loadConfig();
const agent = new Agent(config);
// 设置工具注册表(支持动态工具发现)
agent.setRegistry(toolRegistry);
try {
await agent.chat(question, (text) => {
process.stdout.write(text);
});
console.log('');
} catch (error) {
console.error(
'错误:',
error instanceof Error ? error.message : String(error)
);
process.exit(1);
}
});
// 默认:交互模式
program.action(async () => {
setupPermissions();
const config = loadConfig();
const agent = new Agent(config);
// 初始化 LSP 系统
initLSP(process.cwd());
// 初始化 MCP 系统(加载外部工具服务器)
await initMCP(process.cwd());
// 设置工具注册表(支持动态工具发现)
agent.setRegistry(toolRegistry);
// 初始化会话管理器(支持会话持久化)
const sessionManager = new SessionManager();
await sessionManager.init(process.cwd());
agent.setSessionManager(sessionManager);
// 初始化 todoManager(让 todo 工具可以访问会话)
todoManager.setSessionManager(sessionManager);
// 初始化 Agent 注册表(加载预设和用户配置)
await agentRegistry.init(process.cwd());
// 初始化 Task 工具上下文
initTaskContext(config, sessionManager);
updateTaskDescription();
// 初始化 Skill 注册表
const skillRegistry = getSkillRegistry();
await skillRegistry.initialize(process.cwd());
updateSkillDescription();
// 初始化 Command 注册表
const commandRegistry = getCommandRegistry();
await commandRegistry.initialize(process.cwd());
// 显示会话恢复信息
const session = sessionManager.getSession();
if (session && session.messages.length > 0) {
console.log(`\n📂 已恢复会话 (${session.messages.length} 条消息)`);
}
// 启动终端 UI
const ui = new TerminalUI(agent);
// 优雅退出
process.on('SIGINT', async () => {
console.log('\n\n👋 再见!');
await shutdownMCP();
await shutdownLSP();
await sessionManager.close();
ui.close();
process.exit(0);
});
await ui.start();
});
program.parse();
+284
View File
@@ -0,0 +1,284 @@
/**
* LSP CLI 命令
* 提供语言服务器的查询和安装功能
*/
import { execSync, spawnSync } from 'child_process';
import { getUniqueServers, type ServerConfig, type InstallConfig } from './server.js';
import type { LanguageId } from './language.js';
// 服务器状态
export interface ServerStatus {
id: string;
displayName: string;
description: string;
command: string;
installed: boolean;
languages: LanguageId[];
install: InstallConfig;
}
/**
* 检查命令是否存在
*/
function commandExists(command: string): boolean {
try {
execSync(`which ${command}`, { stdio: 'ignore' });
return true;
} catch {
return false;
}
}
/**
* 检查包管理器是否可用
*/
function hasPackageManager(manager: string): boolean {
return commandExists(manager);
}
/**
* 获取所有服务器状态
*/
export function listServers(): ServerStatus[] {
const servers = getUniqueServers();
return servers.map((server) => ({
id: server.id,
displayName: server.config.displayName,
description: server.config.description,
command: server.config.command,
installed: commandExists(server.config.command),
languages: server.languages,
install: server.config.install,
}));
}
/**
* 打印服务器列表
*/
export function printServerList(): void {
const servers = listServers();
console.log('\n语言服务器状态:\n');
console.log(' 状态 | 服务器 | 支持语言');
console.log(' ------+--------------------------------+------------------');
for (const server of servers) {
const status = server.installed ? ' ✓ ' : ' ✗ ';
const statusColor = server.installed ? '\x1b[32m' : '\x1b[31m';
const reset = '\x1b[0m';
const name = server.displayName.padEnd(30);
const langs = server.languages.slice(0, 3).join(', ') + (server.languages.length > 3 ? '...' : '');
console.log(` ${statusColor}${status}${reset} | ${name} | ${langs}`);
}
const installed = servers.filter((s) => s.installed).length;
console.log(`\n 已安装: ${installed}/${servers.length}\n`);
}
/**
* 获取安装命令
*/
function getInstallCommand(install: InstallConfig): { command: string; description: string } | null {
// 优先使用 npm
if (install.npm && hasPackageManager('npm')) {
return {
command: `npm install -g ${install.npm}`,
description: 'npm',
};
}
// pip
if (install.pip && hasPackageManager('pip3')) {
return {
command: `pip3 install ${install.pip}`,
description: 'pip3',
};
}
if (install.pip && hasPackageManager('pip')) {
return {
command: `pip install ${install.pip}`,
description: 'pip',
};
}
// go
if (install.go && hasPackageManager('go')) {
return {
command: `go install ${install.go}`,
description: 'go install',
};
}
// rustup
if (install.rustup && hasPackageManager('rustup')) {
return {
command: `rustup component add ${install.rustup}`,
description: 'rustup',
};
}
// cargo
if (install.cargo && hasPackageManager('cargo')) {
return {
command: `cargo install ${install.cargo}`,
description: 'cargo',
};
}
// brew
if (install.brew && hasPackageManager('brew')) {
return {
command: `brew install ${install.brew}`,
description: 'Homebrew',
};
}
// gem
if (install.gem && hasPackageManager('gem')) {
return {
command: `gem install ${install.gem}`,
description: 'RubyGems',
};
}
// custom
if (install.custom) {
return {
command: install.custom,
description: '自定义命令',
};
}
return null;
}
/**
* 安装服务器
*/
export async function installServer(serverId: string): Promise<boolean> {
const servers = listServers();
const server = servers.find(
(s) => s.id === serverId || s.displayName.toLowerCase() === serverId.toLowerCase()
);
if (!server) {
console.error(`\x1b[31m错误: 未找到服务器 "${serverId}"\x1b[0m`);
console.log('\n可用的服务器:');
servers.forEach((s) => console.log(` - ${s.displayName} (${s.id})`));
return false;
}
if (server.installed) {
console.log(`\x1b[32m✓ ${server.displayName} 已安装\x1b[0m`);
return true;
}
const installCmd = getInstallCommand(server.install);
if (!installCmd) {
console.error(`\x1b[31m无法自动安装 ${server.displayName}\x1b[0m`);
if (server.install.manual) {
console.log('\n手动安装说明:');
console.log(server.install.manual);
}
return false;
}
console.log(`\n正在安装 ${server.displayName}...`);
console.log(`使用 ${installCmd.description}: ${installCmd.command}\n`);
try {
const result = spawnSync(installCmd.command, {
shell: true,
stdio: 'inherit',
});
if (result.status === 0) {
console.log(`\n\x1b[32m✓ ${server.displayName} 安装成功\x1b[0m`);
return true;
} else {
console.error(`\n\x1b[31m✗ ${server.displayName} 安装失败\x1b[0m`);
if (server.install.manual) {
console.log('\n手动安装说明:');
console.log(server.install.manual);
}
return false;
}
} catch (error) {
console.error(`\n\x1b[31m安装出错: ${error}\x1b[0m`);
return false;
}
}
/**
* 安装所有服务器
*/
export async function installAllServers(): Promise<void> {
const servers = listServers();
const notInstalled = servers.filter((s) => !s.installed);
if (notInstalled.length === 0) {
console.log('\x1b[32m所有语言服务器都已安装\x1b[0m');
return;
}
console.log(`\n将安装 ${notInstalled.length} 个语言服务器:\n`);
notInstalled.forEach((s) => console.log(` - ${s.displayName}`));
console.log('');
let success = 0;
let failed = 0;
for (const server of notInstalled) {
const result = await installServer(server.id);
if (result) {
success++;
} else {
failed++;
}
console.log('');
}
console.log(`\n安装完成: ${success} 成功, ${failed} 失败`);
}
/**
* 显示服务器详细信息
*/
export function showServerInfo(serverId: string): void {
const servers = listServers();
const server = servers.find(
(s) => s.id === serverId || s.displayName.toLowerCase() === serverId.toLowerCase()
);
if (!server) {
console.error(`\x1b[31m错误: 未找到服务器 "${serverId}"\x1b[0m`);
return;
}
const status = server.installed ? '\x1b[32m已安装\x1b[0m' : '\x1b[31m未安装\x1b[0m';
console.log(`\n${server.displayName}`);
console.log('='.repeat(40));
console.log(`状态: ${status}`);
console.log(`命令: ${server.command}`);
console.log(`描述: ${server.description}`);
console.log(`支持语言: ${server.languages.join(', ')}`);
if (!server.installed) {
const installCmd = getInstallCommand(server.install);
if (installCmd) {
console.log(`\n安装命令 (${installCmd.description}):`);
console.log(` ${installCmd.command}`);
}
if (server.install.manual) {
console.log('\n手动安装说明:');
console.log(server.install.manual);
}
}
console.log('');
}
+409
View File
@@ -0,0 +1,409 @@
/**
* LSP 客户端
* 负责与语言服务器通信
*/
import { spawn, ChildProcess } from 'child_process';
import {
createMessageConnection,
StreamMessageReader,
StreamMessageWriter,
MessageConnection,
} from 'vscode-jsonrpc/node.js';
import {
InitializeParams,
Diagnostic,
DiagnosticSeverity,
PublishDiagnosticsParams,
} from 'vscode-languageserver-protocol';
import * as fs from 'fs/promises';
import * as path from 'path';
import { getLanguageId, type LanguageId } from './language.js';
import { getServerConfig, type ServerConfig } from './server.js';
// 诊断信息接口
export interface FileDiagnostic {
file: string;
line: number;
column: number;
endLine?: number;
endColumn?: number;
severity: 'error' | 'warning' | 'info' | 'hint';
message: string;
source?: string;
code?: string | number;
}
// 客户端状态
interface ClientState {
process: ChildProcess;
connection: MessageConnection;
languageId: LanguageId;
config: ServerConfig;
initialized: boolean;
openDocuments: Set<string>;
documentVersions: Map<string, number>;
diagnostics: Map<string, Diagnostic[]>;
rootUri: string;
}
/**
* LSP 客户端管理器
* 管理多个语言服务器的生命周期
*/
export class LSPClientManager {
private clients: Map<LanguageId, ClientState> = new Map();
private rootPath: string;
constructor(rootPath?: string) {
this.rootPath = rootPath || process.cwd();
}
/**
* 设置工作区根目录
*/
setRootPath(rootPath: string): void {
this.rootPath = rootPath;
}
/**
* 获取或启动语言服务器
*/
async getClient(languageId: LanguageId): Promise<ClientState | undefined> {
// 如果已有客户端,直接返回
if (this.clients.has(languageId)) {
return this.clients.get(languageId);
}
// 获取服务器配置
const config = getServerConfig(languageId);
if (!config) {
return undefined;
}
// 启动语言服务器
try {
const client = await this.startServer(languageId, config);
this.clients.set(languageId, client);
return client;
} catch {
// 静默忽略启动失败(如语言服务器未安装)
return undefined;
}
}
/**
* 启动语言服务器
*/
private async startServer(languageId: LanguageId, config: ServerConfig): Promise<ClientState> {
// 检查命令是否可用
const commandExists = await this.checkCommand(config.command);
if (!commandExists) {
throw new Error(`语言服务器命令不存在: ${config.command}`);
}
// 启动进程
const serverProcess = spawn(config.command, config.args, {
env: { ...process.env, ...config.env },
stdio: ['pipe', 'pipe', 'pipe'],
});
if (!serverProcess.stdin || !serverProcess.stdout) {
throw new Error('无法创建进程管道');
}
// 创建 JSON-RPC 连接
const connection = createMessageConnection(
new StreamMessageReader(serverProcess.stdout),
new StreamMessageWriter(serverProcess.stdin)
);
const rootUri = `file://${this.rootPath}`;
const state: ClientState = {
process: serverProcess,
connection,
languageId,
config,
initialized: false,
openDocuments: new Set(),
documentVersions: new Map(),
diagnostics: new Map(),
rootUri,
};
// 监听诊断通知(使用字符串方法名)
connection.onNotification('textDocument/publishDiagnostics', (params: PublishDiagnosticsParams) => {
const filePath = params.uri.replace('file://', '');
state.diagnostics.set(filePath, params.diagnostics);
});
// 监听进程错误
serverProcess.on('error', (error) => {
console.error(`语言服务器错误 (${languageId}):`, error);
});
serverProcess.on('exit', () => {
this.clients.delete(languageId);
});
// 启动连接
connection.listen();
// 初始化服务器
await this.initializeServer(state, config);
return state;
}
/**
* 初始化语言服务器
*/
private async initializeServer(state: ClientState, config: ServerConfig): Promise<void> {
const initParams: InitializeParams = {
processId: process.pid,
rootUri: state.rootUri,
capabilities: {
textDocument: {
synchronization: {
dynamicRegistration: false,
willSave: false,
willSaveWaitUntil: false,
didSave: true,
},
publishDiagnostics: {
relatedInformation: true,
tagSupport: { valueSet: [1, 2] },
},
},
workspace: {
workspaceFolders: true,
},
},
workspaceFolders: [
{
uri: state.rootUri,
name: path.basename(this.rootPath),
},
],
initializationOptions: config.initializationOptions,
};
// 使用字符串方法名发送请求
await state.connection.sendRequest('initialize', initParams);
await state.connection.sendNotification('initialized', {});
state.initialized = true;
}
/**
* 通知文件变更(打开或更新)
* @returns 是否是首次启动服务器(需要更长等待时间)
*/
async touchFile(filePath: string, isNew: boolean = false): Promise<boolean> {
const absolutePath = path.isAbsolute(filePath)
? filePath
: path.join(this.rootPath, filePath);
const languageId = getLanguageId(absolutePath);
if (!languageId) {
return false;
}
// 检查是否是首次启动
const wasRunning = this.clients.has(languageId);
const client = await this.getClient(languageId);
if (!client || !client.initialized) {
return false;
}
const isFirstStart = !wasRunning;
const uri = `file://${absolutePath}`;
const content = await fs.readFile(absolutePath, 'utf-8');
if (!client.openDocuments.has(absolutePath)) {
// 打开文档(使用字符串方法名)
await client.connection.sendNotification('textDocument/didOpen', {
textDocument: {
uri,
languageId,
version: 1,
text: content,
},
});
client.openDocuments.add(absolutePath);
client.documentVersions.set(absolutePath, 1);
} else {
// 更新文档
const version = (client.documentVersions.get(absolutePath) || 0) + 1;
client.documentVersions.set(absolutePath, version);
await client.connection.sendNotification('textDocument/didChange', {
textDocument: {
uri,
version,
},
contentChanges: [{ text: content }],
});
}
// 等待诊断结果(给服务器一些时间处理)
await this.waitForDiagnostics(100);
return isFirstStart;
}
/**
* 关闭文档
*/
async closeFile(filePath: string): Promise<void> {
const absolutePath = path.isAbsolute(filePath)
? filePath
: path.join(this.rootPath, filePath);
const languageId = getLanguageId(absolutePath);
if (!languageId) {
return;
}
const client = this.clients.get(languageId);
if (!client || !client.openDocuments.has(absolutePath)) {
return;
}
const uri = `file://${absolutePath}`;
await client.connection.sendNotification('textDocument/didClose', {
textDocument: { uri },
});
client.openDocuments.delete(absolutePath);
client.documentVersions.delete(absolutePath);
client.diagnostics.delete(absolutePath);
}
/**
* 获取文件诊断信息
*/
getDiagnostics(filePath?: string): Map<string, FileDiagnostic[]> {
const result = new Map<string, FileDiagnostic[]>();
for (const client of this.clients.values()) {
for (const [file, diagnostics] of client.diagnostics) {
if (filePath && file !== filePath) {
continue;
}
const fileDiagnostics = diagnostics.map((d) => this.convertDiagnostic(file, d));
if (fileDiagnostics.length > 0) {
const existing = result.get(file) || [];
result.set(file, [...existing, ...fileDiagnostics]);
}
}
}
return result;
}
/**
* 获取单个文件的诊断信息
*/
getFileDiagnostics(filePath: string): FileDiagnostic[] {
const absolutePath = path.isAbsolute(filePath)
? filePath
: path.join(this.rootPath, filePath);
for (const client of this.clients.values()) {
const diagnostics = client.diagnostics.get(absolutePath);
if (diagnostics) {
return diagnostics.map((d) => this.convertDiagnostic(absolutePath, d));
}
}
return [];
}
/**
* 转换诊断信息格式
*/
private convertDiagnostic(file: string, diagnostic: Diagnostic): FileDiagnostic {
return {
file,
line: diagnostic.range.start.line + 1,
column: diagnostic.range.start.character + 1,
endLine: diagnostic.range.end.line + 1,
endColumn: diagnostic.range.end.character + 1,
severity: this.convertSeverity(diagnostic.severity),
message: diagnostic.message,
source: diagnostic.source,
code: diagnostic.code as string | number | undefined,
};
}
/**
* 转换严重性
*/
private convertSeverity(severity?: DiagnosticSeverity): FileDiagnostic['severity'] {
switch (severity) {
case DiagnosticSeverity.Error:
return 'error';
case DiagnosticSeverity.Warning:
return 'warning';
case DiagnosticSeverity.Information:
return 'info';
case DiagnosticSeverity.Hint:
return 'hint';
default:
return 'error';
}
}
/**
* 检查命令是否存在
*/
private async checkCommand(command: string): Promise<boolean> {
try {
const { execSync } = await import('child_process');
execSync(`which ${command}`, { stdio: 'ignore' });
return true;
} catch {
return false;
}
}
/**
* 等待诊断结果
*/
private waitForDiagnostics(ms: number): Promise<void> {
return new Promise((resolve) => setTimeout(resolve, ms));
}
/**
* 关闭所有客户端
*/
async shutdown(): Promise<void> {
for (const [languageId, client] of this.clients) {
try {
client.connection.dispose();
client.process.kill();
} catch (error) {
console.error(`关闭语言服务器失败 (${languageId}):`, error);
}
}
this.clients.clear();
}
/**
* 检查服务器是否运行中
*/
isServerRunning(languageId: LanguageId): boolean {
return this.clients.has(languageId);
}
/**
* 获取运行中的服务器列表
*/
getRunningServers(): LanguageId[] {
return Array.from(this.clients.keys());
}
}
+132
View File
@@ -0,0 +1,132 @@
/**
* LSP 模块入口
* 提供简化的 API 供其他模块使用
*/
import { LSPClientManager, type FileDiagnostic } from './client.js';
import { getLanguageId, isLanguageSupported } from './language.js';
// 全局 LSP 管理器实例
let lspManager: LSPClientManager | null = null;
/**
* 初始化 LSP 系统
*/
export function initLSP(rootPath?: string): void {
if (lspManager) {
return;
}
lspManager = new LSPClientManager(rootPath);
}
/**
* 获取 LSP 管理器实例
*/
export function getLSPManager(): LSPClientManager | null {
return lspManager;
}
/**
* 通知 LSP 文件已变更
* 用于编辑/写入文件后通知语言服务器
* @returns 是否是首次启动服务器(需要更长等待时间)
*/
export async function touchFile(filePath: string, isNew: boolean = false): Promise<boolean> {
if (!lspManager) {
return false;
}
if (!isLanguageSupported(filePath)) {
return false;
}
try {
return await lspManager.touchFile(filePath, isNew);
} catch (error) {
// 静默失败,LSP 是增强功能
console.error('LSP touchFile 错误:', error);
return false;
}
}
/**
* 获取所有文件的诊断信息
*/
export function getDiagnostics(): Map<string, FileDiagnostic[]> {
if (!lspManager) {
return new Map();
}
return lspManager.getDiagnostics();
}
/**
* 获取单个文件的诊断信息
*/
export function getFileDiagnostics(filePath: string): FileDiagnostic[] {
if (!lspManager) {
return [];
}
return lspManager.getFileDiagnostics(filePath);
}
/**
* 格式化诊断信息为字符串
* 用于返回给 AI
*/
export function formatDiagnostics(diagnostics: FileDiagnostic[]): string {
if (diagnostics.length === 0) {
return '';
}
const lines: string[] = [];
for (const d of diagnostics) {
const location = `${d.line}:${d.column}`;
const severity = d.severity.toUpperCase();
const code = d.code ? ` [${d.code}]` : '';
const source = d.source ? ` (${d.source})` : '';
lines.push(` ${location} ${severity}${code}: ${d.message}${source}`);
}
return lines.join('\n');
}
/**
* 获取文件诊断并格式化为 AI 可读的格式
* 只显示错误和警告,忽略 hint 和 info
*/
export async function getFormattedFileDiagnostics(filePath: string): Promise<string> {
const diagnostics = getFileDiagnostics(filePath);
// 只关注错误和警告,忽略 hint 和 info
const errors = diagnostics.filter(d => d.severity === 'error');
const warnings = diagnostics.filter(d => d.severity === 'warning');
const relevantDiagnostics = [...errors, ...warnings];
// 没有错误和警告就不显示
if (relevantDiagnostics.length === 0) {
return '';
}
let result = `\n<file_diagnostics file="${filePath}">\n`;
result += `发现 ${errors.length} 个错误, ${warnings.length} 个警告:\n`;
result += formatDiagnostics(relevantDiagnostics);
result += '\n</file_diagnostics>';
return result;
}
/**
* 关闭 LSP 系统
*/
export async function shutdownLSP(): Promise<void> {
if (lspManager) {
await lspManager.shutdown();
lspManager = null;
}
}
// 导出类型
export type { FileDiagnostic } from './client.js';
export { getLanguageId, isLanguageSupported, getSupportedExtensions } from './language.js';
export { getServerConfig, hasServerConfig, getSupportedLanguages } from './server.js';
export { LSPClientManager } from './client.js';
+138
View File
@@ -0,0 +1,138 @@
/**
* 文件扩展名到 LSP languageId 的映射
*/
// 语言 ID 定义
export type LanguageId =
| 'typescript'
| 'javascript'
| 'typescriptreact'
| 'javascriptreact'
| 'python'
| 'go'
| 'rust'
| 'java'
| 'c'
| 'cpp'
| 'csharp'
| 'php'
| 'ruby'
| 'swift'
| 'kotlin'
| 'scala'
| 'html'
| 'css'
| 'scss'
| 'less'
| 'json'
| 'yaml'
| 'markdown'
| 'vue'
| 'svelte';
// 扩展名到语言 ID 的映射
const extensionToLanguageId: Record<string, LanguageId> = {
// TypeScript/JavaScript
'.ts': 'typescript',
'.tsx': 'typescriptreact',
'.js': 'javascript',
'.jsx': 'javascriptreact',
'.mjs': 'javascript',
'.cjs': 'javascript',
'.mts': 'typescript',
'.cts': 'typescript',
// Python
'.py': 'python',
'.pyi': 'python',
'.pyw': 'python',
// Go
'.go': 'go',
// Rust
'.rs': 'rust',
// Java
'.java': 'java',
// C/C++
'.c': 'c',
'.h': 'c',
'.cpp': 'cpp',
'.cc': 'cpp',
'.cxx': 'cpp',
'.hpp': 'cpp',
'.hh': 'cpp',
'.hxx': 'cpp',
// C#
'.cs': 'csharp',
// PHP
'.php': 'php',
// Ruby
'.rb': 'ruby',
'.rake': 'ruby',
// Swift
'.swift': 'swift',
// Kotlin
'.kt': 'kotlin',
'.kts': 'kotlin',
// Scala
'.scala': 'scala',
'.sc': 'scala',
// Web
'.html': 'html',
'.htm': 'html',
'.css': 'css',
'.scss': 'scss',
'.less': 'less',
'.vue': 'vue',
'.svelte': 'svelte',
// Data formats
'.json': 'json',
'.yaml': 'yaml',
'.yml': 'yaml',
// Markdown
'.md': 'markdown',
'.markdown': 'markdown',
};
/**
* 根据文件路径获取语言 ID
*/
export function getLanguageId(filePath: string): LanguageId | undefined {
const ext = getExtension(filePath);
return extensionToLanguageId[ext];
}
/**
* 获取文件扩展名(小写)
*/
function getExtension(filePath: string): string {
const lastDot = filePath.lastIndexOf('.');
if (lastDot === -1) return '';
return filePath.slice(lastDot).toLowerCase();
}
/**
* 检查文件是否支持 LSP
*/
export function isLanguageSupported(filePath: string): boolean {
return getLanguageId(filePath) !== undefined;
}
/**
* 获取所有支持的扩展名
*/
export function getSupportedExtensions(): string[] {
return Object.keys(extensionToLanguageId);
}
+336
View File
@@ -0,0 +1,336 @@
/**
* 语言服务器定义
* 定义各种语言的 LSP 服务器配置
*/
import type { LanguageId } from './language.js';
// 安装命令配置
export interface InstallConfig {
/** npm 全局安装包名 */
npm?: string;
/** pip 安装包名 */
pip?: string;
/** go install 路径 */
go?: string;
/** cargo install 包名 */
cargo?: string;
/** brew 安装包名 */
brew?: string;
/** gem 安装包名 */
gem?: string;
/** rustup component */
rustup?: string;
/** 自定义安装命令 */
custom?: string;
/** 安装说明(当无法自动安装时) */
manual?: string;
}
// 服务器配置接口
export interface ServerConfig {
/** 服务器命令 */
command: string;
/** 命令参数 */
args: string[];
/** 环境变量 */
env?: Record<string, string>;
/** 初始化选项 */
initializationOptions?: Record<string, unknown>;
/** 安装配置 */
install: InstallConfig;
/** 显示名称 */
displayName: string;
/** 描述 */
description: string;
}
// 语言服务器定义
const serverConfigs: Partial<Record<LanguageId, ServerConfig>> = {
// TypeScript/JavaScript - 使用 typescript-language-server
typescript: {
command: 'typescript-language-server',
args: ['--stdio'],
initializationOptions: {
preferences: {
includeInlayParameterNameHints: 'all',
includeInlayPropertyDeclarationTypeHints: true,
includeInlayFunctionLikeReturnTypeHints: true,
},
},
install: {
npm: 'typescript-language-server typescript',
},
displayName: 'TypeScript',
description: 'TypeScript/JavaScript 语言服务器',
},
javascript: {
command: 'typescript-language-server',
args: ['--stdio'],
install: {
npm: 'typescript-language-server typescript',
},
displayName: 'JavaScript',
description: 'JavaScript 语言服务器(共用 TypeScript)',
},
typescriptreact: {
command: 'typescript-language-server',
args: ['--stdio'],
install: {
npm: 'typescript-language-server typescript',
},
displayName: 'TypeScript React',
description: 'TSX 语言服务器(共用 TypeScript)',
},
javascriptreact: {
command: 'typescript-language-server',
args: ['--stdio'],
install: {
npm: 'typescript-language-server typescript',
},
displayName: 'JavaScript React',
description: 'JSX 语言服务器(共用 TypeScript)',
},
// Python - 使用 pyright
python: {
command: 'pyright-langserver',
args: ['--stdio'],
install: {
npm: 'pyright',
pip: 'pyright',
},
displayName: 'Python',
description: 'Python 语言服务器 (Pyright)',
},
// Go - 使用 gopls
go: {
command: 'gopls',
args: ['serve'],
install: {
go: 'golang.org/x/tools/gopls@latest',
},
displayName: 'Go',
description: 'Go 语言服务器 (gopls)',
},
// Rust - 使用 rust-analyzer
rust: {
command: 'rust-analyzer',
args: [],
install: {
rustup: 'rust-analyzer',
brew: 'rust-analyzer',
},
displayName: 'Rust',
description: 'Rust 语言服务器 (rust-analyzer)',
},
// C/C++ - 使用 clangd
c: {
command: 'clangd',
args: ['--background-index'],
install: {
brew: 'llvm',
manual: 'Ubuntu: sudo apt install clangd\nmacOS: brew install llvm',
},
displayName: 'C',
description: 'C 语言服务器 (clangd)',
},
cpp: {
command: 'clangd',
args: ['--background-index'],
install: {
brew: 'llvm',
manual: 'Ubuntu: sudo apt install clangd\nmacOS: brew install llvm',
},
displayName: 'C++',
description: 'C++ 语言服务器 (clangd)',
},
// Java - 使用 jdtls
java: {
command: 'jdtls',
args: [],
install: {
brew: 'jdtls',
manual: '请访问 https://github.com/eclipse/eclipse.jdt.ls 获取安装说明',
},
displayName: 'Java',
description: 'Java 语言服务器 (Eclipse JDT.LS)',
},
// C# - 使用 OmniSharp
csharp: {
command: 'omnisharp',
args: ['-lsp'],
install: {
brew: 'omnisharp/omnisharp-roslyn/omnisharp-mono',
manual: '请访问 https://github.com/OmniSharp/omnisharp-roslyn 获取安装说明',
},
displayName: 'C#',
description: 'C# 语言服务器 (OmniSharp)',
},
// PHP - 使用 intelephense
php: {
command: 'intelephense',
args: ['--stdio'],
install: {
npm: 'intelephense',
},
displayName: 'PHP',
description: 'PHP 语言服务器 (Intelephense)',
},
// Ruby - 使用 solargraph
ruby: {
command: 'solargraph',
args: ['stdio'],
install: {
gem: 'solargraph',
},
displayName: 'Ruby',
description: 'Ruby 语言服务器 (Solargraph)',
},
// Vue - 使用 vue-language-server
vue: {
command: 'vue-language-server',
args: ['--stdio'],
install: {
npm: '@vue/language-server',
},
displayName: 'Vue',
description: 'Vue 语言服务器',
},
// Svelte - 使用 svelte-language-server
svelte: {
command: 'svelteserver',
args: ['--stdio'],
install: {
npm: 'svelte-language-server',
},
displayName: 'Svelte',
description: 'Svelte 语言服务器',
},
// HTML - 使用 vscode-html-language-server
html: {
command: 'vscode-html-language-server',
args: ['--stdio'],
install: {
npm: 'vscode-langservers-extracted',
},
displayName: 'HTML',
description: 'HTML 语言服务器',
},
// CSS/SCSS/Less - 使用 vscode-css-language-server
css: {
command: 'vscode-css-language-server',
args: ['--stdio'],
install: {
npm: 'vscode-langservers-extracted',
},
displayName: 'CSS',
description: 'CSS 语言服务器',
},
scss: {
command: 'vscode-css-language-server',
args: ['--stdio'],
install: {
npm: 'vscode-langservers-extracted',
},
displayName: 'SCSS',
description: 'SCSS 语言服务器(共用 CSS)',
},
less: {
command: 'vscode-css-language-server',
args: ['--stdio'],
install: {
npm: 'vscode-langservers-extracted',
},
displayName: 'Less',
description: 'Less 语言服务器(共用 CSS)',
},
// JSON - 使用 vscode-json-language-server
json: {
command: 'vscode-json-language-server',
args: ['--stdio'],
install: {
npm: 'vscode-langservers-extracted',
},
displayName: 'JSON',
description: 'JSON 语言服务器',
},
// YAML - 使用 yaml-language-server
yaml: {
command: 'yaml-language-server',
args: ['--stdio'],
install: {
npm: 'yaml-language-server',
},
displayName: 'YAML',
description: 'YAML 语言服务器',
},
};
/**
* 获取语言服务器配置
*/
export function getServerConfig(languageId: LanguageId): ServerConfig | undefined {
return serverConfigs[languageId];
}
/**
* 检查语言是否有可用的服务器配置
*/
export function hasServerConfig(languageId: LanguageId): boolean {
return languageId in serverConfigs;
}
/**
* 获取所有支持的语言 ID
*/
export function getSupportedLanguages(): LanguageId[] {
return Object.keys(serverConfigs) as LanguageId[];
}
/**
* 获取所有服务器配置(用于 CLI)
*/
export function getAllServerConfigs(): Partial<Record<LanguageId, ServerConfig>> {
return serverConfigs;
}
/**
* 获取唯一的服务器列表(去重相同命令的服务器)
*/
export function getUniqueServers(): Array<{ id: string; config: ServerConfig; languages: LanguageId[] }> {
const commandMap = new Map<string, { config: ServerConfig; languages: LanguageId[] }>();
for (const [langId, config] of Object.entries(serverConfigs)) {
if (!config) continue;
const existing = commandMap.get(config.command);
if (existing) {
existing.languages.push(langId as LanguageId);
} else {
commandMap.set(config.command, {
config,
languages: [langId as LanguageId],
});
}
}
return Array.from(commandMap.entries()).map(([command, data]) => ({
id: command,
config: data.config,
languages: data.languages,
}));
}
+265
View File
@@ -0,0 +1,265 @@
/**
* MCP 客户端
* 管理与单个 MCP 服务器的连接
*/
import type {
Transport,
MCPServerConfig,
MCPTool,
MCPServerStatus,
MCPServerStatusType,
MCPToolCallResult,
InitializeParams,
InitializeResult,
ListToolsResult,
CallToolParams,
CallToolResult,
ServerCapabilities,
MCPContent,
} from './types.js';
import { createTransport } from './transports/index.js';
/** MCP 协议版本 */
const PROTOCOL_VERSION = '2024-11-05';
/** 客户端信息 */
const CLIENT_INFO = {
name: 'ai-terminal-assistant',
version: '1.0.0',
};
/**
* MCP 客户端类
*/
export class MCPClient {
private transport: Transport;
private serverCapabilities?: ServerCapabilities;
private serverInfo?: { name: string; version: string };
private tools: MCPTool[] = [];
private _status: MCPServerStatusType = 'disconnected';
private _error?: string;
private lastConnected?: Date;
constructor(
private name: string,
private config: MCPServerConfig
) {
this.transport = createTransport(name, config);
}
/**
* 连接到 MCP 服务器
*/
async connect(): Promise<void> {
if (this._status === 'connected') {
return;
}
this._status = 'connecting';
this._error = undefined;
try {
// 启动传输层
await this.transport.start();
// 监听连接关闭
this.transport.onClose((error) => {
this._status = 'disconnected';
if (error) {
this._error = error.message;
}
});
// 监听服务器通知
this.transport.onNotification((method, params) => {
this.handleNotification(method, params);
});
// 初始化握手
await this.initialize();
// 获取工具列表
await this.refreshTools();
this._status = 'connected';
this.lastConnected = new Date();
} catch (error) {
this._status = 'error';
this._error = error instanceof Error ? error.message : String(error);
throw error;
}
}
/**
* 断开连接
*/
async disconnect(): Promise<void> {
if (this._status === 'disconnected') {
return;
}
try {
await this.transport.close();
} finally {
this._status = 'disconnected';
this.tools = [];
}
}
/**
* MCP 初始化握手
*/
private async initialize(): Promise<void> {
const params: InitializeParams = {
protocolVersion: PROTOCOL_VERSION,
capabilities: {
roots: { listChanged: true },
},
clientInfo: CLIENT_INFO,
};
const result = (await this.transport.request(
'initialize',
params
)) as InitializeResult;
this.serverCapabilities = result.capabilities;
this.serverInfo = result.serverInfo;
// 发送 initialized 通知
await this.transport.notify('notifications/initialized');
}
/**
* 刷新工具列表
*/
async refreshTools(): Promise<void> {
const result = (await this.transport.request(
'tools/list',
{}
)) as ListToolsResult;
this.tools = result.tools.map((tool) => ({
server: this.name,
name: `${this.name}-${tool.name}`,
originalName: tool.name,
description: tool.description || '',
inputSchema: tool.inputSchema,
outputSchema: tool.outputSchema,
}));
}
/**
* 调用工具
*/
async callTool(
toolName: string,
args: Record<string, unknown>
): Promise<MCPToolCallResult> {
// 获取原始工具名
const tool = this.tools.find((t) => t.name === toolName);
if (!tool) {
return {
success: false,
content: [{ type: 'text', text: `Tool not found: ${toolName}` }],
isError: true,
};
}
try {
const params: CallToolParams = {
name: tool.originalName,
arguments: args,
};
const result = (await this.transport.request(
'tools/call',
params
)) as CallToolResult;
return {
success: !result.isError,
content: result.content as MCPContent[],
isError: result.isError,
};
} catch (error) {
return {
success: false,
content: [
{
type: 'text',
text: error instanceof Error ? error.message : String(error),
},
],
isError: true,
};
}
}
/**
* 获取工具列表
*/
getTools(): MCPTool[] {
return [...this.tools];
}
/**
* 获取状态
*/
getStatus(): MCPServerStatus {
return {
name: this.name,
type: this.config.type,
status: this._status,
toolCount: this.tools.length,
error: this._error,
lastConnected: this.lastConnected,
};
}
/**
* 获取服务器信息
*/
getServerInfo(): { name: string; version: string } | undefined {
return this.serverInfo;
}
/**
* 获取服务器能力
*/
getServerCapabilities(): ServerCapabilities | undefined {
return this.serverCapabilities;
}
/**
* 处理服务器通知
*/
private handleNotification(method: string, params: unknown): void {
switch (method) {
case 'notifications/tools/list_changed':
// 工具列表变化,刷新
this.refreshTools().catch((error) => {
console.error(
`[MCP:${this.name}] Failed to refresh tools:`,
error
);
});
break;
case 'notifications/resources/list_changed':
// 资源列表变化(暂不处理)
break;
case 'notifications/prompts/list_changed':
// 提示列表变化(暂不处理)
break;
default:
console.debug(
`[MCP:${this.name}] Unknown notification: ${method}`,
params
);
}
}
}
+342
View File
@@ -0,0 +1,342 @@
/**
* MCP 配置加载和验证
*/
import * as fs from 'fs';
import * as path from 'path';
import * as yaml from 'yaml';
import type {
MCPConfig,
MCPServerConfig,
LocalMCPServer,
RemoteMCPServer,
} from './types.js';
/** 默认配置值 */
const DEFAULTS = {
timeout: 30000,
enabled: true,
};
/**
* 解析环境变量引用
* 支持 {env:VAR_NAME} 语法
*/
export function resolveEnvVariables(value: string): string {
return value.replace(/\{env:([^}]+)\}/g, (_, varName) => {
const envValue = process.env[varName];
if (envValue === undefined) {
console.warn(`警告: 环境变量 ${varName} 未设置`);
return '';
}
return envValue;
});
}
/**
* 递归解析对象中的环境变量
*/
export function resolveEnvInObject<T>(obj: T): T {
if (typeof obj === 'string') {
return resolveEnvVariables(obj) as T;
}
if (Array.isArray(obj)) {
return obj.map((item) => resolveEnvInObject(item)) as T;
}
if (obj !== null && typeof obj === 'object') {
const result: Record<string, unknown> = {};
for (const [key, value] of Object.entries(obj)) {
result[key] = resolveEnvInObject(value);
}
return result as T;
}
return obj;
}
/**
* 验证本地服务器配置
*/
function validateLocalServer(
name: string,
config: LocalMCPServer
): string | null {
if (!config.command || !Array.isArray(config.command)) {
return `服务器 "${name}": command 必须是字符串数组`;
}
if (config.command.length === 0) {
return `服务器 "${name}": command 不能为空`;
}
if (config.timeout !== undefined && typeof config.timeout !== 'number') {
return `服务器 "${name}": timeout 必须是数字`;
}
return null;
}
/**
* 验证远程服务器配置
*/
function validateRemoteServer(
name: string,
config: RemoteMCPServer
): string | null {
if (!config.url || typeof config.url !== 'string') {
return `服务器 "${name}": url 是必需的`;
}
try {
new URL(config.url);
} catch {
return `服务器 "${name}": url 格式无效`;
}
if (config.timeout !== undefined && typeof config.timeout !== 'number') {
return `服务器 "${name}": timeout 必须是数字`;
}
return null;
}
/**
* 验证服务器配置
*/
function validateServerConfig(
name: string,
config: MCPServerConfig
): string | null {
if (!config.type) {
return `服务器 "${name}": type 是必需的 (local 或 remote)`;
}
if (config.type === 'local') {
return validateLocalServer(name, config as LocalMCPServer);
}
if (config.type === 'remote') {
return validateRemoteServer(name, config as RemoteMCPServer);
}
return `服务器 "${name}": type 必须是 "local" 或 "remote"`;
}
/**
* 验证 MCP 配置
*/
export function validateMCPConfig(config: MCPConfig): string[] {
const errors: string[] = [];
if (!config.mcp) {
return errors; // 没有 MCP 配置是合法的
}
if (typeof config.mcp !== 'object' || Array.isArray(config.mcp)) {
errors.push('mcp 配置必须是对象格式');
return errors;
}
for (const [name, serverConfig] of Object.entries(config.mcp)) {
// 验证服务器名称
if (!/^[a-zA-Z][a-zA-Z0-9_-]*$/.test(name)) {
errors.push(
`服务器名称 "${name}" 无效: 必须以字母开头,只能包含字母、数字、下划线和连字符`
);
continue;
}
const error = validateServerConfig(name, serverConfig);
if (error) {
errors.push(error);
}
}
// 验证 tools 配置
if (config.tools) {
if (typeof config.tools !== 'object' || Array.isArray(config.tools)) {
errors.push('tools 配置必须是对象格式');
} else {
for (const [pattern, enabled] of Object.entries(config.tools)) {
if (typeof enabled !== 'boolean') {
errors.push(`tools["${pattern}"] 的值必须是布尔值`);
}
}
}
}
return errors;
}
/**
* 规范化服务器配置(添加默认值)
*/
function normalizeServerConfig(config: MCPServerConfig): MCPServerConfig {
return {
...config,
enabled: config.enabled ?? DEFAULTS.enabled,
timeout: config.timeout ?? DEFAULTS.timeout,
};
}
/**
* 规范化 MCP 配置
*/
export function normalizeMCPConfig(config: MCPConfig): MCPConfig {
if (!config.mcp) {
return config;
}
const normalizedMcp: Record<string, MCPServerConfig> = {};
for (const [name, serverConfig] of Object.entries(config.mcp)) {
normalizedMcp[name] = normalizeServerConfig(serverConfig);
}
return {
...config,
mcp: normalizedMcp,
};
}
/**
* 从文件加载配置
*/
function loadConfigFile(filePath: string): MCPConfig | null {
if (!fs.existsSync(filePath)) {
return null;
}
try {
const content = fs.readFileSync(filePath, 'utf-8');
const ext = path.extname(filePath).toLowerCase();
if (ext === '.json' || ext === '.jsonc') {
// 移除 JSONC 注释
const jsonContent = content.replace(
/\/\/.*$|\/\*[\s\S]*?\*\//gm,
''
);
return JSON.parse(jsonContent);
}
if (ext === '.yaml' || ext === '.yml') {
return yaml.parse(content);
}
console.warn(`不支持的配置文件格式: ${ext}`);
return null;
} catch (error) {
console.error(`加载配置文件失败 ${filePath}:`, error);
return null;
}
}
/**
* 合并配置(后者覆盖前者)
*/
function mergeConfigs(base: MCPConfig, override: MCPConfig): MCPConfig {
return {
mcp: {
...base.mcp,
...override.mcp,
},
tools: {
...base.tools,
...override.tools,
},
};
}
/**
* 加载 MCP 配置
* 按优先级从低到高加载:
* 1. 用户级: ~/.ai-assist/config.{json,yaml}
* 2. 项目级: .ai-assist/config.{json,yaml}
*/
export function loadMCPConfig(workdir?: string): MCPConfig {
const cwd = workdir || process.cwd();
const homeDir = process.env.HOME || process.env.USERPROFILE || '';
// 配置文件搜索路径
const searchPaths = [
// 用户级配置
path.join(homeDir, '.ai-assist', 'config.json'),
path.join(homeDir, '.ai-assist', 'config.jsonc'),
path.join(homeDir, '.ai-assist', 'config.yaml'),
path.join(homeDir, '.ai-assist', 'config.yml'),
// 项目级配置
path.join(cwd, '.ai-assist', 'config.json'),
path.join(cwd, '.ai-assist', 'config.jsonc'),
path.join(cwd, '.ai-assist', 'config.yaml'),
path.join(cwd, '.ai-assist', 'config.yml'),
];
let mergedConfig: MCPConfig = {};
for (const configPath of searchPaths) {
const config = loadConfigFile(configPath);
if (config) {
mergedConfig = mergeConfigs(mergedConfig, config);
}
}
// 验证配置
const errors = validateMCPConfig(mergedConfig);
if (errors.length > 0) {
console.error('MCP 配置验证失败:');
for (const error of errors) {
console.error(` - ${error}`);
}
// 返回空配置而不是无效配置
return {};
}
// 规范化配置
const normalizedConfig = normalizeMCPConfig(mergedConfig);
// 解析环境变量
return resolveEnvInObject(normalizedConfig);
}
/**
* 获取已启用的服务器列表
*/
export function getEnabledServers(
config: MCPConfig
): Array<{ name: string; config: MCPServerConfig }> {
if (!config.mcp) {
return [];
}
return Object.entries(config.mcp)
.filter(([, serverConfig]) => serverConfig.enabled !== false)
.map(([name, serverConfig]) => ({ name, config: serverConfig }));
}
/**
* 检查工具是否被配置启用
* 支持通配符匹配,如 "server-*"
*/
export function isToolEnabled(
toolName: string,
toolsConfig?: Record<string, boolean>
): boolean {
if (!toolsConfig) {
return true; // 默认启用
}
// 精确匹配优先
if (toolName in toolsConfig) {
return toolsConfig[toolName];
}
// 通配符匹配(从最具体到最通用)
const patterns = Object.keys(toolsConfig)
.filter((pattern) => pattern.includes('*'))
.sort((a, b) => {
// 更具体的模式(* 出现位置更靠后)优先
const aPos = a.indexOf('*');
const bPos = b.indexOf('*');
return bPos - aPos;
});
for (const pattern of patterns) {
const regex = new RegExp(
'^' + pattern.replace(/\*/g, '.*').replace(/\?/g, '.') + '$'
);
if (regex.test(toolName)) {
return toolsConfig[pattern];
}
}
return true; // 默认启用
}
+47
View File
@@ -0,0 +1,47 @@
/**
* MCP 模块导出入口
*/
// 类型导出
export type {
MCPConfig,
MCPServerConfig,
LocalMCPServer,
RemoteMCPServer,
OAuthConfig,
MCPTool,
MCPServerStatus,
MCPServerStatusType,
MCPToolCallResult,
MCPContent,
MCPTextContent,
MCPImageContent,
MCPResourceContent,
Transport,
ServerCapabilities,
ClientCapabilities,
MCPManagerEvents,
} from './types.js';
// 配置相关
export {
loadMCPConfig,
validateMCPConfig,
normalizeMCPConfig,
getEnabledServers,
isToolEnabled,
resolveEnvVariables,
resolveEnvInObject,
} from './config.js';
// 客户端
export { MCPClient } from './client.js';
// 管理器
export { MCPManager, getMCPManager, resetMCPManager } from './manager.js';
// 工具适配器
export { MCPToolAdapter, createMCPToolAdapter } from './tool-adapter.js';
// 传输层
export { createTransport, StdioTransport } from './transports/index.js';
+319
View File
@@ -0,0 +1,319 @@
/**
* MCP Manager
* 管理多个 MCP 服务器的生命周期
*/
import { EventEmitter } from 'events';
import type {
MCPConfig,
MCPServerConfig,
MCPTool,
MCPServerStatus,
MCPToolCallResult,
MCPManagerEvents,
} from './types.js';
import { MCPClient } from './client.js';
import { getEnabledServers, isToolEnabled } from './config.js';
/**
* MCP Manager 类
* 负责管理所有 MCP 服务器连接
*/
export class MCPManager extends EventEmitter {
private clients: Map<string, MCPClient> = new Map();
private config: MCPConfig = {};
private initialized = false;
/**
* 初始化所有已启用的服务器
*/
async initialize(config: MCPConfig): Promise<void> {
if (this.initialized) {
await this.shutdown();
}
this.config = config;
const servers = getEnabledServers(config);
// 并行连接所有服务器
const connectPromises = servers.map(async ({ name, config: serverConfig }) => {
try {
await this.connectServer(name, serverConfig);
} catch (error) {
console.error(
`[MCP] Failed to connect to ${name}:`,
error instanceof Error ? error.message : String(error)
);
}
});
await Promise.all(connectPromises);
this.initialized = true;
}
/**
* 连接单个服务器
*/
private async connectServer(
name: string,
config: MCPServerConfig
): Promise<void> {
const client = new MCPClient(name, config);
try {
await client.connect();
this.clients.set(name, client);
this.emit('server:connected', name);
} catch (error) {
this.emit('server:error', name, error instanceof Error ? error : new Error(String(error)));
throw error;
}
}
/**
* 关闭所有连接
*/
async shutdown(): Promise<void> {
const disconnectPromises = Array.from(this.clients.entries()).map(
async ([name, client]) => {
try {
await client.disconnect();
this.emit('server:disconnected', name);
} catch (error) {
console.error(`[MCP] Error disconnecting ${name}:`, error);
}
}
);
await Promise.all(disconnectPromises);
this.clients.clear();
this.initialized = false;
}
/**
* 重连指定服务器
*/
async reconnect(serverName: string): Promise<void> {
const existingClient = this.clients.get(serverName);
if (existingClient) {
await existingClient.disconnect();
this.clients.delete(serverName);
this.emit('server:disconnected', serverName);
}
const serverConfig = this.config.mcp?.[serverName];
if (!serverConfig) {
throw new Error(`Server not found: ${serverName}`);
}
await this.connectServer(serverName, serverConfig);
this.emit('tools:changed');
}
/**
* 获取所有可用工具
*/
getTools(): MCPTool[] {
const allTools: MCPTool[] = [];
for (const client of this.clients.values()) {
const status = client.getStatus();
if (status.status === 'connected') {
const tools = client.getTools();
// 根据配置过滤工具
for (const tool of tools) {
if (isToolEnabled(tool.name, this.config.tools)) {
allTools.push(tool);
}
}
}
}
return allTools;
}
/**
* 获取指定工具
*/
getTool(name: string): MCPTool | undefined {
for (const client of this.clients.values()) {
const tools = client.getTools();
const tool = tools.find((t) => t.name === name);
if (tool && isToolEnabled(tool.name, this.config.tools)) {
return tool;
}
}
return undefined;
}
/**
* 调用工具
*/
async callTool(
name: string,
args: Record<string, unknown>
): Promise<MCPToolCallResult> {
// 从工具名中解析服务器名
const dashIndex = name.indexOf('-');
if (dashIndex === -1) {
return {
success: false,
content: [{ type: 'text', text: `Invalid tool name format: ${name}` }],
isError: true,
};
}
const serverName = name.substring(0, dashIndex);
const client = this.clients.get(serverName);
if (!client) {
return {
success: false,
content: [{ type: 'text', text: `Server not connected: ${serverName}` }],
isError: true,
};
}
// 检查工具是否启用
if (!isToolEnabled(name, this.config.tools)) {
return {
success: false,
content: [{ type: 'text', text: `Tool is disabled: ${name}` }],
isError: true,
};
}
return await client.callTool(name, args);
}
/**
* 获取所有服务器状态
*/
getServerStatuses(): MCPServerStatus[] {
const statuses: MCPServerStatus[] = [];
// 已连接的服务器
for (const client of this.clients.values()) {
statuses.push(client.getStatus());
}
// 未连接但已配置的服务器
if (this.config.mcp) {
for (const [name, serverConfig] of Object.entries(this.config.mcp)) {
if (!this.clients.has(name)) {
statuses.push({
name,
type: serverConfig.type,
status: serverConfig.enabled === false ? 'disabled' : 'disconnected',
toolCount: 0,
});
}
}
}
return statuses;
}
/**
* 获取指定服务器状态
*/
getServerStatus(name: string): MCPServerStatus | undefined {
const client = this.clients.get(name);
if (client) {
return client.getStatus();
}
const serverConfig = this.config.mcp?.[name];
if (serverConfig) {
return {
name,
type: serverConfig.type,
status: serverConfig.enabled === false ? 'disabled' : 'disconnected',
toolCount: 0,
};
}
return undefined;
}
/**
* 启用/禁用服务器
*/
async setServerEnabled(serverName: string, enabled: boolean): Promise<void> {
const serverConfig = this.config.mcp?.[serverName];
if (!serverConfig) {
throw new Error(`Server not found: ${serverName}`);
}
if (enabled) {
// 启用并连接
serverConfig.enabled = true;
if (!this.clients.has(serverName)) {
await this.connectServer(serverName, serverConfig);
this.emit('tools:changed');
}
} else {
// 禁用并断开
serverConfig.enabled = false;
const client = this.clients.get(serverName);
if (client) {
await client.disconnect();
this.clients.delete(serverName);
this.emit('server:disconnected', serverName);
this.emit('tools:changed');
}
}
}
/**
* 获取客户端实例(用于测试)
*/
getClient(name: string): MCPClient | undefined {
return this.clients.get(name);
}
/**
* 是否已初始化
*/
isInitialized(): boolean {
return this.initialized;
}
// 类型安全的事件方法
override on<K extends keyof MCPManagerEvents>(
event: K,
listener: MCPManagerEvents[K]
): this {
return super.on(event, listener);
}
override emit<K extends keyof MCPManagerEvents>(
event: K,
...args: Parameters<MCPManagerEvents[K]>
): boolean {
return super.emit(event, ...args);
}
}
// 单例实例
let mcpManager: MCPManager | null = null;
/**
* 获取 MCP Manager 单例
*/
export function getMCPManager(): MCPManager {
if (!mcpManager) {
mcpManager = new MCPManager();
}
return mcpManager;
}
/**
* 重置 MCP Manager(用于测试)
*/
export function resetMCPManager(): void {
if (mcpManager) {
mcpManager.shutdown().catch(console.error);
mcpManager = null;
}
}
+166
View File
@@ -0,0 +1,166 @@
/**
* MCP 工具适配器
* 将 MCP 工具转换为内部工具格式
*/
import type { ToolParameter, ToolResult } from '../types/index.js';
import type { ToolWithMetadata, ToolCategory } from '../tools/types.js';
import type { MCPTool, MCPToolCallResult, MCPContent } from './types.js';
import type { MCPManager } from './manager.js';
/**
* 将 JSON Schema 类型转换为内部参数类型
*/
function convertJsonSchemaType(
schemaType: string | string[] | undefined
): 'string' | 'number' | 'boolean' | 'object' | 'array' {
if (Array.isArray(schemaType)) {
// 取第一个非 null 类型
const type = schemaType.find((t) => t !== 'null');
return convertJsonSchemaType(type);
}
switch (schemaType) {
case 'string':
return 'string';
case 'number':
case 'integer':
return 'number';
case 'boolean':
return 'boolean';
case 'array':
return 'array';
case 'object':
default:
return 'object';
}
}
/**
* 将 JSON Schema 转换为内部参数定义
*/
function convertInputSchema(
schema: MCPTool['inputSchema']
): Record<string, ToolParameter> {
const parameters: Record<string, ToolParameter> = {};
if (!schema || typeof schema !== 'object') {
return parameters;
}
const properties = schema.properties as Record<string, {
type?: string | string[];
description?: string;
}> | undefined;
const required = (schema.required as string[]) || [];
if (properties) {
for (const [name, prop] of Object.entries(properties)) {
parameters[name] = {
type: convertJsonSchemaType(prop.type),
description: prop.description || '',
required: required.includes(name),
};
}
}
return parameters;
}
/**
* 将 MCP 内容转换为字符串输出
*/
function contentToString(content: MCPContent[]): string {
const parts: string[] = [];
for (const item of content) {
switch (item.type) {
case 'text':
parts.push(item.text);
break;
case 'image':
parts.push(`[Image: ${item.mimeType}]`);
break;
case 'resource':
if (item.resource.text) {
parts.push(item.resource.text);
} else {
parts.push(`[Resource: ${item.resource.uri}]`);
}
break;
}
}
return parts.join('\n');
}
/**
* 将 MCP 调用结果转换为内部 ToolResult
*/
function convertToolResult(result: MCPToolCallResult): ToolResult {
const output = contentToString(result.content);
if (result.isError || !result.success) {
return {
success: false,
output: '',
error: output || 'Tool execution failed',
};
}
return {
success: true,
output,
};
}
/**
* MCP 工具适配器类
*/
export class MCPToolAdapter {
constructor(private manager: MCPManager) {}
/**
* 将 MCP 工具适配为内部工具格式
*/
adaptToInternalTool(mcpTool: MCPTool): ToolWithMetadata {
const parameters = convertInputSchema(mcpTool.inputSchema);
return {
name: mcpTool.name,
description: mcpTool.description || `MCP tool: ${mcpTool.originalName}`,
parameters,
execute: async (params: Record<string, unknown>): Promise<ToolResult> => {
const result = await this.manager.callTool(mcpTool.name, params);
return convertToolResult(result);
},
metadata: {
name: mcpTool.name,
category: 'agent' as ToolCategory, // MCP 工具归类为 agent
description: mcpTool.description || `MCP tool from ${mcpTool.server}`,
keywords: [
'mcp',
mcpTool.server,
mcpTool.originalName,
...mcpTool.originalName.split(/[_-]/),
],
deferLoading: false, // MCP 工具不延迟加载
},
};
}
/**
* 批量适配工具
*/
adaptTools(mcpTools: MCPTool[]): ToolWithMetadata[] {
return mcpTools.map((tool) => this.adaptToInternalTool(tool));
}
}
/**
* 创建 MCP 工具适配器
*/
export function createMCPToolAdapter(manager: MCPManager): MCPToolAdapter {
return new MCPToolAdapter(manager);
}
+27
View File
@@ -0,0 +1,27 @@
/**
* 传输层工厂
*/
import type { Transport, MCPServerConfig, LocalMCPServer } from '../types.js';
import { StdioTransport } from './stdio.js';
/**
* 创建传输层实例
*/
export function createTransport(
serverName: string,
config: MCPServerConfig
): Transport {
if (config.type === 'local') {
return new StdioTransport(config as LocalMCPServer, serverName);
}
if (config.type === 'remote') {
// TODO: 实现 HTTP 传输
throw new Error('Remote transport not yet implemented');
}
throw new Error(`Unknown transport type: ${(config as MCPServerConfig).type}`);
}
export { StdioTransport } from './stdio.js';
+278
View File
@@ -0,0 +1,278 @@
/**
* stdio 传输实现
* 通过子进程的 stdin/stdout 与 MCP 服务器通信
*/
import { spawn, type ChildProcess } from 'child_process';
import { EventEmitter } from 'events';
import type {
Transport,
JSONRPCRequest,
JSONRPCResponse,
JSONRPCNotification,
LocalMCPServer,
} from '../types.js';
/** 默认请求超时 */
const DEFAULT_TIMEOUT = 30000;
/** 待处理请求 */
interface PendingRequest {
resolve: (result: unknown) => void;
reject: (error: Error) => void;
timer: NodeJS.Timeout;
}
/**
* stdio 传输层实现
*/
export class StdioTransport extends EventEmitter implements Transport {
private process: ChildProcess | null = null;
private nextId = 1;
private pendingRequests = new Map<string | number, PendingRequest>();
private buffer = '';
private notificationHandler?: (method: string, params: unknown) => void;
private closeHandler?: (error?: Error) => void;
private isClosing = false;
constructor(
private config: LocalMCPServer,
private serverName: string
) {
super();
}
/**
* 启动传输(创建子进程)
*/
async start(): Promise<void> {
if (this.process) {
throw new Error('Transport already started');
}
const [command, ...args] = this.config.command;
return new Promise((resolve, reject) => {
try {
this.process = spawn(command, args, {
cwd: this.config.cwd,
env: {
...process.env,
...this.config.env,
},
stdio: ['pipe', 'pipe', 'pipe'],
});
// 处理 stdout(JSON-RPC 消息)
this.process.stdout?.on('data', (data: Buffer) => {
this.handleData(data);
});
// 处理 stderr(日志输出)
this.process.stderr?.on('data', (data: Buffer) => {
const message = data.toString().trim();
if (message) {
console.debug(`[MCP:${this.serverName}] ${message}`);
}
});
// 处理进程错误
this.process.on('error', (error) => {
if (!this.isClosing) {
this.handleClose(error);
}
reject(error);
});
// 处理进程退出
this.process.on('exit', (code, signal) => {
if (!this.isClosing) {
const error =
code !== 0
? new Error(
`Process exited with code ${code}${signal ? ` (signal: ${signal})` : ''}`
)
: undefined;
this.handleClose(error);
}
});
// 给进程一点时间启动
setTimeout(() => {
if (this.process && !this.process.killed) {
resolve();
}
}, 100);
} catch (error) {
reject(error);
}
});
}
/**
* 关闭传输
*/
async close(): Promise<void> {
this.isClosing = true;
// 拒绝所有待处理请求
for (const [id, pending] of this.pendingRequests) {
clearTimeout(pending.timer);
pending.reject(new Error('Transport closed'));
this.pendingRequests.delete(id);
}
// 终止子进程
if (this.process && !this.process.killed) {
return new Promise((resolve) => {
const timeout = setTimeout(() => {
this.process?.kill('SIGKILL');
resolve();
}, 5000);
this.process!.on('exit', () => {
clearTimeout(timeout);
resolve();
});
this.process!.kill('SIGTERM');
});
}
}
/**
* 发送请求并等待响应
*/
async request(method: string, params?: unknown): Promise<unknown> {
if (!this.process || this.process.killed) {
throw new Error('Transport not connected');
}
const id = this.nextId++;
const request: JSONRPCRequest = {
jsonrpc: '2.0',
id,
method,
params,
};
return new Promise((resolve, reject) => {
const timeout = this.config.timeout ?? DEFAULT_TIMEOUT;
const timer = setTimeout(() => {
this.pendingRequests.delete(id);
reject(new Error(`Request timeout after ${timeout}ms: ${method}`));
}, timeout);
this.pendingRequests.set(id, { resolve, reject, timer });
this.send(request);
});
}
/**
* 发送通知(无响应)
*/
async notify(method: string, params?: unknown): Promise<void> {
if (!this.process || this.process.killed) {
throw new Error('Transport not connected');
}
const notification: JSONRPCNotification = {
jsonrpc: '2.0',
method,
params,
};
this.send(notification);
}
/**
* 监听服务器通知
*/
onNotification(handler: (method: string, params: unknown) => void): void {
this.notificationHandler = handler;
}
/**
* 监听连接关闭
*/
onClose(handler: (error?: Error) => void): void {
this.closeHandler = handler;
}
/**
* 发送 JSON-RPC 消息
*/
private send(message: JSONRPCRequest | JSONRPCNotification): void {
const data = JSON.stringify(message) + '\n';
this.process?.stdin?.write(data);
}
/**
* 处理接收到的数据
*/
private handleData(data: Buffer): void {
this.buffer += data.toString();
// 按行分割处理
const lines = this.buffer.split('\n');
this.buffer = lines.pop() || ''; // 保留不完整的行
for (const line of lines) {
if (line.trim()) {
try {
this.handleMessage(JSON.parse(line));
} catch (error) {
console.error(
`[MCP:${this.serverName}] Failed to parse message:`,
line
);
}
}
}
}
/**
* 处理 JSON-RPC 消息
*/
private handleMessage(message: JSONRPCResponse | JSONRPCNotification): void {
// 响应消息
if ('id' in message && message.id !== undefined) {
const pending = this.pendingRequests.get(message.id);
if (pending) {
clearTimeout(pending.timer);
this.pendingRequests.delete(message.id);
if ('error' in message && message.error) {
pending.reject(
new Error(
`${message.error.message} (code: ${message.error.code})`
)
);
} else {
pending.resolve(message.result);
}
}
return;
}
// 通知消息
if ('method' in message) {
this.notificationHandler?.(message.method, message.params);
}
}
/**
* 处理连接关闭
*/
private handleClose(error?: Error): void {
// 拒绝所有待处理请求
for (const [id, pending] of this.pendingRequests) {
clearTimeout(pending.timer);
pending.reject(error || new Error('Connection closed'));
this.pendingRequests.delete(id);
}
this.closeHandler?.(error);
this.process = null;
}
}
+276
View File
@@ -0,0 +1,276 @@
/**
* MCP (Model Context Protocol) 类型定义
*/
import type { JSONSchema7 } from 'json-schema';
// ============================================================================
// 配置类型
// ============================================================================
/** 本地 MCP 服务器配置 (stdio 传输) */
export interface LocalMCPServer {
type: 'local';
/** 启动命令和参数,如 ["npx", "-y", "@anthropic/mcp-server-filesystem", "/path"] */
command: string[];
/** 环境变量,支持 {env:VAR} 语法引用系统环境变量 */
env?: Record<string, string>;
/** 工作目录 */
cwd?: string;
/** 是否启用,默认 true */
enabled?: boolean;
/** 超时毫秒数,默认 30000 */
timeout?: number;
}
/** 远程 MCP 服务器配置 (HTTP/SSE 传输) */
export interface RemoteMCPServer {
type: 'remote';
/** MCP 端点 URL */
url: string;
/** 自定义请求头,支持 {env:VAR} 语法 */
headers?: Record<string, string>;
/** OAuth 配置,空对象 {} 表示自动发现 (RFC 7591) */
oauth?: OAuthConfig | Record<string, never>;
/** 是否启用,默认 true */
enabled?: boolean;
/** 超时毫秒数,默认 30000 */
timeout?: number;
}
/** OAuth 配置 */
export interface OAuthConfig {
clientId: string;
clientSecret: string;
scope?: string;
/** Token 端点,可选,自动发现时不需要 */
tokenEndpoint?: string;
}
/** MCP 服务器配置(本地或远程) */
export type MCPServerConfig = LocalMCPServer | RemoteMCPServer;
/** MCP 配置根节点 */
export interface MCPConfig {
/** MCP 服务器配置,key 为服务器名称 */
mcp?: Record<string, MCPServerConfig>;
/** 工具启用/禁用配置,支持通配符如 "server-*" */
tools?: Record<string, boolean>;
}
// ============================================================================
// 运行时类型
// ============================================================================
/** MCP 工具定义 */
export interface MCPTool {
/** 来源服务器名称 */
server: string;
/** 完整工具名: {server}-{originalName} */
name: string;
/** MCP 服务器中的原始名称 */
originalName: string;
/** 工具描述 */
description: string;
/** 输入参数 JSON Schema */
inputSchema: JSONSchema7;
/** 输出结果 JSON Schema(可选) */
outputSchema?: JSONSchema7;
}
/** 服务器运行状态 */
export type MCPServerStatusType =
| 'connected'
| 'connecting'
| 'disconnected'
| 'auth_required'
| 'disabled'
| 'error';
/** 服务器状态信息 */
export interface MCPServerStatus {
/** 服务器名称 */
name: string;
/** 服务器类型 */
type: 'local' | 'remote';
/** 当前状态 */
status: MCPServerStatusType;
/** 工具数量 */
toolCount: number;
/** 错误信息(当 status 为 error 时) */
error?: string;
/** 最后连接时间 */
lastConnected?: Date;
}
/** 工具调用结果 */
export interface MCPToolCallResult {
/** 是否成功 */
success: boolean;
/** 结果内容 */
content: MCPContent[];
/** 是否为错误结果 */
isError?: boolean;
}
/** MCP 内容类型 */
export type MCPContent =
| MCPTextContent
| MCPImageContent
| MCPResourceContent;
/** 文本内容 */
export interface MCPTextContent {
type: 'text';
text: string;
}
/** 图片内容 */
export interface MCPImageContent {
type: 'image';
data: string;
mimeType: string;
}
/** 资源内容 */
export interface MCPResourceContent {
type: 'resource';
resource: {
uri: string;
mimeType?: string;
text?: string;
blob?: string;
};
}
// ============================================================================
// 传输层类型
// ============================================================================
/** JSON-RPC 请求 */
export interface JSONRPCRequest {
jsonrpc: '2.0';
id: string | number;
method: string;
params?: unknown;
}
/** JSON-RPC 响应 */
export interface JSONRPCResponse {
jsonrpc: '2.0';
id: string | number;
result?: unknown;
error?: JSONRPCError;
}
/** JSON-RPC 通知 */
export interface JSONRPCNotification {
jsonrpc: '2.0';
method: string;
params?: unknown;
}
/** JSON-RPC 错误 */
export interface JSONRPCError {
code: number;
message: string;
data?: unknown;
}
/** 传输层接口 */
export interface Transport {
/** 启动传输 */
start(): Promise<void>;
/** 关闭传输 */
close(): Promise<void>;
/** 发送请求并等待响应 */
request(method: string, params?: unknown): Promise<unknown>;
/** 发送通知(无响应) */
notify(method: string, params?: unknown): Promise<void>;
/** 监听服务器通知 */
onNotification(handler: (method: string, params: unknown) => void): void;
/** 监听连接关闭 */
onClose(handler: (error?: Error) => void): void;
}
// ============================================================================
// MCP 协议消息类型
// ============================================================================
/** 服务器能力 */
export interface ServerCapabilities {
tools?: {
listChanged?: boolean;
};
resources?: {
subscribe?: boolean;
listChanged?: boolean;
};
prompts?: {
listChanged?: boolean;
};
}
/** 客户端能力 */
export interface ClientCapabilities {
roots?: {
listChanged?: boolean;
};
sampling?: Record<string, never>;
}
/** 初始化请求参数 */
export interface InitializeParams {
protocolVersion: string;
capabilities: ClientCapabilities;
clientInfo: {
name: string;
version: string;
};
}
/** 初始化响应结果 */
export interface InitializeResult {
protocolVersion: string;
capabilities: ServerCapabilities;
serverInfo: {
name: string;
version: string;
};
instructions?: string;
}
/** 工具列表响应 */
export interface ListToolsResult {
tools: Array<{
name: string;
description?: string;
inputSchema: JSONSchema7;
outputSchema?: JSONSchema7;
}>;
nextCursor?: string;
}
/** 工具调用参数 */
export interface CallToolParams {
name: string;
arguments?: Record<string, unknown>;
}
/** 工具调用响应 */
export interface CallToolResult {
content: MCPContent[];
isError?: boolean;
}
// ============================================================================
// 事件类型
// ============================================================================
/** MCP Manager 事件 */
export interface MCPManagerEvents {
'server:connected': (serverName: string) => void;
'server:disconnected': (serverName: string, error?: Error) => void;
'server:error': (serverName: string, error: Error) => void;
'tools:changed': () => void;
}
+225
View File
@@ -0,0 +1,225 @@
import { Parser, Language, type Node } from 'web-tree-sitter';
import * as path from 'path';
import { fileURLToPath } from 'url';
const __filename = fileURLToPath(import.meta.url);
const __dirname = path.dirname(__filename);
// 解析后的命令结构
export interface ParsedCommand {
name: string; // 命令名,如 "git"
subcommand?: string; // 子命令,如 "push"
args: string[]; // 参数列表
text: string; // 原始命令文本
}
// 解析结果
export interface ParseResult {
commands: ParsedCommand[]; // 所有解析出的命令
success: boolean;
error?: string;
}
// 单例解析器
let parserInstance: Parser | null = null;
let bashLanguage: Language | null = null;
let initPromise: Promise<void> | null = null;
/**
* 获取 wasm 文件路径
*/
function getWasmPath(filename: string): string {
// 从 node_modules 加载
const nodeModulesPath = path.resolve(__dirname, '../../node_modules');
if (filename === 'tree-sitter.wasm') {
return path.join(nodeModulesPath, 'web-tree-sitter', filename);
} else if (filename === 'tree-sitter-bash.wasm') {
return path.join(nodeModulesPath, 'tree-sitter-bash', filename);
}
throw new Error(`Unknown wasm file: ${filename}`);
}
/**
* 初始化解析器(懒加载,只初始化一次)
*/
async function initParser(): Promise<void> {
if (parserInstance && bashLanguage) {
return;
}
if (initPromise) {
return initPromise;
}
initPromise = (async () => {
try {
// 初始化 tree-sitter
await Parser.init({
locateFile: (scriptName: string) => {
return getWasmPath(scriptName);
},
});
// 创建解析器实例
parserInstance = new Parser();
// 加载 bash 语言
const bashWasmPath = getWasmPath('tree-sitter-bash.wasm');
bashLanguage = await Language.load(bashWasmPath);
parserInstance.setLanguage(bashLanguage);
} catch (error) {
initPromise = null;
throw error;
}
})();
return initPromise;
}
/**
* 从语法树节点中提取命令信息
*/
function extractCommandFromNode(node: Node): ParsedCommand {
const parts: string[] = [];
for (let i = 0; i < node.childCount; i++) {
const child = node.child(i);
if (!child) continue;
// 提取命令名和参数
if (
child.type === 'command_name' ||
child.type === 'word' ||
child.type === 'string' ||
child.type === 'raw_string' ||
child.type === 'concatenation' ||
child.type === 'simple_expansion' ||
child.type === 'expansion'
) {
// 对于字符串类型,提取内部文本(去掉引号)
if (child.type === 'string' || child.type === 'raw_string') {
const text = child.text;
if ((text.startsWith('"') && text.endsWith('"')) ||
(text.startsWith("'") && text.endsWith("'"))) {
parts.push(text.slice(1, -1));
} else {
parts.push(text);
}
} else {
parts.push(child.text);
}
}
}
const name = parts[0] || '';
// 找到第一个非 flag 参数作为子命令
let subcommand: string | undefined;
const args: string[] = [];
for (let i = 1; i < parts.length; i++) {
const part = parts[i];
if (!subcommand && !part.startsWith('-')) {
subcommand = part;
} else {
args.push(part);
}
}
return {
name,
subcommand,
args,
text: node.text,
};
}
/**
* 递归查找所有命令节点
*/
function findCommandNodes(node: Node): Node[] {
const commands: Node[] = [];
if (node.type === 'command') {
commands.push(node);
}
for (let i = 0; i < node.childCount; i++) {
const child = node.child(i);
if (child) {
commands.push(...findCommandNodes(child));
}
}
return commands;
}
/**
* 解析 bash 命令字符串
*/
export async function parseBashCommand(command: string): Promise<ParseResult> {
try {
await initParser();
if (!parserInstance) {
return {
commands: [],
success: false,
error: 'Parser not initialized',
};
}
const tree = parserInstance.parse(command);
if (!tree) {
return {
commands: [],
success: false,
error: 'Failed to parse command',
};
}
const commandNodes = findCommandNodes(tree.rootNode);
const commands = commandNodes.map(extractCommandFromNode);
return {
commands,
success: true,
};
} catch (error) {
return {
commands: [],
success: false,
error: error instanceof Error ? error.message : String(error),
};
}
}
/**
* 简单解析(用于降级,当 tree-sitter 不可用时)
*/
export function parseCommandSimple(command: string): ParsedCommand {
const parts = command.trim().split(/\s+/);
const name = parts[0] || '';
let subcommand: string | undefined;
const args: string[] = [];
for (let i = 1; i < parts.length; i++) {
const part = parts[i];
if (!subcommand && !part.startsWith('-')) {
subcommand = part;
} else {
args.push(part);
}
}
return { name, subcommand, args, text: command };
}
/**
* 检查解析器是否已初始化
*/
export function isParserInitialized(): boolean {
return parserInstance !== null && bashLanguage !== null;
}
@@ -0,0 +1,34 @@
import type { PermissionCheckResult, PermissionContext } from '../types.js';
/**
* 权限检查器基础接口
* 所有工具的权限检查器都应该实现此接口
*/
export interface PermissionChecker {
/**
* 检查器名称
*/
readonly name: string;
/**
* 检查权限
* @param ctx 权限检查上下文
* @returns 权限检查结果
*/
check(ctx: PermissionContext): Promise<PermissionCheckResult>;
/**
* 清除会话级别的临时权限
*/
clearSessionPermissions(): void;
}
/**
* 权限检查器配置基类
*/
export interface BasePermissionConfig {
/**
* 默认权限动作
*/
default: 'allow' | 'deny' | 'ask';
}
@@ -0,0 +1,305 @@
import * as fs from 'fs';
import * as path from 'path';
import * as os from 'os';
import type {
BashPermissionConfig,
PermissionContext,
PermissionCheckResult,
PermissionDecision,
PermissionRule,
} from '../types.js';
import type { PermissionChecker } from './base.js';
import { matchRulesAsync } from '../wildcard.js';
import { parseBashCommand } from '../bash-parser.js';
const CONFIG_DIR = path.join(os.homedir(), '.ai-terminal-assistant');
const PERMISSION_FILE = path.join(CONFIG_DIR, 'permissions.json');
// 默认权限配置
const DEFAULT_CONFIG: BashPermissionConfig = {
rules: [
// 默认允许的安全命令
{ pattern: 'ls *', action: 'allow' },
{ pattern: 'cat *', action: 'allow' },
{ pattern: 'head *', action: 'allow' },
{ pattern: 'tail *', action: 'allow' },
{ pattern: 'grep *', action: 'allow' },
{ pattern: 'find *', action: 'allow' },
{ pattern: 'echo *', action: 'allow' },
{ pattern: 'pwd', action: 'allow' },
{ pattern: 'which *', action: 'allow' },
{ pattern: 'type *', action: 'allow' },
{ pattern: 'git status', action: 'allow' },
{ pattern: 'git log *', action: 'allow' },
{ pattern: 'git diff *', action: 'allow' },
{ pattern: 'git branch *', action: 'allow' },
{ pattern: 'npm list *', action: 'allow' },
{ pattern: 'npm run *', action: 'ask' },
// 需要确认的命令
{ pattern: 'git push *', action: 'ask' },
{ pattern: 'git commit *', action: 'ask' },
{ pattern: 'git checkout *', action: 'ask' },
{ pattern: 'git reset *', action: 'ask' },
{ pattern: 'npm install *', action: 'ask' },
{ pattern: 'npm uninstall *', action: 'ask' },
{ pattern: 'yarn *', action: 'ask' },
{ pattern: 'pnpm *', action: 'ask' },
// 危险命令 - 默认拒绝
{ pattern: 'rm -rf *', action: 'deny' },
{ pattern: 'rm -r *', action: 'ask' },
{ pattern: 'sudo *', action: 'deny' },
{ pattern: 'chmod 777 *', action: 'deny' },
{ pattern: 'mkfs *', action: 'deny' },
{ pattern: 'dd *', action: 'deny' },
{ pattern: '> /dev/*', action: 'deny' },
],
externalDirectory: 'ask',
default: 'ask',
};
/**
* Bash 命令权限检查器
* 使用 tree-sitter 解析命令并检查权限
*/
export class BashPermissionChecker implements PermissionChecker {
readonly name = 'bash';
private config: BashPermissionConfig;
private projectRoot: string;
private askCallback?: (ctx: PermissionContext) => Promise<PermissionDecision>;
private sessionPermissions = new Map<string, 'allow' | 'deny'>();
constructor(projectRoot: string = process.cwd()) {
this.projectRoot = path.resolve(projectRoot);
this.config = this.loadConfig();
}
/**
* 设置权限询问回调
*/
setAskCallback(callback: (ctx: PermissionContext) => Promise<PermissionDecision>): void {
this.askCallback = callback;
}
/**
* 加载权限配置
*/
private loadConfig(): BashPermissionConfig {
try {
if (fs.existsSync(PERMISSION_FILE)) {
const content = fs.readFileSync(PERMISSION_FILE, 'utf-8');
const stored = JSON.parse(content) as Partial<BashPermissionConfig>;
return {
rules: stored.rules || DEFAULT_CONFIG.rules,
externalDirectory: stored.externalDirectory || DEFAULT_CONFIG.externalDirectory,
default: stored.default || DEFAULT_CONFIG.default,
};
}
} catch {
// 忽略加载错误
}
return { ...DEFAULT_CONFIG };
}
/**
* 保存权限配置
*/
saveConfig(): void {
if (!fs.existsSync(CONFIG_DIR)) {
fs.mkdirSync(CONFIG_DIR, { recursive: true });
}
fs.writeFileSync(PERMISSION_FILE, JSON.stringify(this.config, null, 2));
}
/**
* 添加权限规则
*/
addRule(rule: PermissionRule): void {
const existingIndex = this.config.rules.findIndex(r => r.pattern === rule.pattern);
if (existingIndex >= 0) {
this.config.rules[existingIndex] = rule;
} else {
this.config.rules.unshift(rule);
}
this.saveConfig();
}
/**
* 检查路径是否在项目目录内
*/
private isInProjectDirectory(targetPath: string): boolean {
const resolved = path.resolve(targetPath);
return resolved.startsWith(this.projectRoot + path.sep) || resolved === this.projectRoot;
}
/**
* 从解析后的命令中提取可能的外部路径
*/
private async extractPathsFromCommand(command: string, workdir: string): Promise<string[]> {
const externalPaths: string[] = [];
const parseResult = await parseBashCommand(command);
const pathCommands = new Set(['cd', 'rm', 'cp', 'mv', 'mkdir', 'touch', 'chmod', 'chown', 'cat', 'ls']);
for (const cmd of parseResult.commands) {
if (!pathCommands.has(cmd.name)) continue;
const pathsToCheck = [cmd.subcommand, ...cmd.args].filter(Boolean) as string[];
for (const arg of pathsToCheck) {
if (arg.startsWith('-')) continue;
let resolved: string | null = null;
if (arg.startsWith('/')) {
resolved = arg;
} else if (arg.startsWith('~')) {
resolved = arg.replace('~', os.homedir());
} else if (arg.includes('..') || arg.includes('/')) {
resolved = path.resolve(workdir, arg);
}
if (resolved && !this.isInProjectDirectory(resolved)) {
externalPaths.push(resolved);
}
}
}
return [...new Set(externalPaths)];
}
/**
* 检查命令权限
*/
async check(ctx: PermissionContext): Promise<PermissionCheckResult> {
const { command, workdir } = ctx;
// 1. 使用 tree-sitter 解析命令并检查权限
const parseResult = await matchRulesAsync(command, this.config.rules, this.config.default);
// 2. 检查会话级别的临时权限
for (const pattern of parseResult.askPatterns) {
const sessionPerm = this.sessionPermissions.get(pattern);
if (sessionPerm === 'deny') {
return {
allowed: false,
action: 'deny',
reason: `本次会话已拒绝此类命令: ${pattern}`,
};
}
if (sessionPerm === 'allow') {
const allAllowed = parseResult.askPatterns.every(p => this.sessionPermissions.get(p) === 'allow');
if (allAllowed && parseResult.action === 'ask') {
return {
allowed: true,
action: 'allow',
reason: '本次会话已允许此类命令',
};
}
}
}
// 3. 检查外部目录访问
const externalPaths = await this.extractPathsFromCommand(command, workdir);
if (externalPaths.length > 0) {
const extAction = this.config.externalDirectory;
if (extAction === 'deny') {
return {
allowed: false,
action: 'deny',
reason: `命令访问项目目录外的路径: ${externalPaths.join(', ')}`,
};
}
if (extAction === 'ask') {
if (!this.askCallback) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason: `命令访问项目目录外的路径`,
patterns: externalPaths,
};
}
const decision = await this.askCallback({
...ctx,
externalPaths,
});
if (!decision.allow) {
return {
allowed: false,
action: 'deny',
reason: '用户拒绝访问外部目录',
};
}
}
}
// 4. 根据解析结果处理权限
const { action, matchedPattern, allCommands, askPatterns } = parseResult;
if (action === 'allow') {
return {
allowed: true,
action: 'allow',
reason: matchedPattern ? `匹配规则: ${matchedPattern}` : '默认允许',
};
}
if (action === 'deny') {
return {
allowed: false,
action: 'deny',
reason: matchedPattern
? `匹配规则: ${matchedPattern}`
: `包含被拒绝的命令: ${allCommands.map(c => c.name).join(', ')}`,
};
}
// action === 'ask'
if (!this.askCallback) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
patterns: askPatterns,
};
}
const decision = await this.askCallback({
...ctx,
patterns: askPatterns,
});
if (decision.remember) {
for (const pattern of askPatterns) {
this.sessionPermissions.set(pattern, decision.allow ? 'allow' : 'deny');
}
}
return {
allowed: decision.allow,
action: decision.allow ? 'allow' : 'deny',
reason: decision.allow ? '用户允许' : '用户拒绝',
};
}
/**
* 清除会话权限
*/
clearSessionPermissions(): void {
this.sessionPermissions.clear();
}
/**
* 获取当前配置
*/
getConfig(): BashPermissionConfig {
return { ...this.config };
}
}
@@ -0,0 +1,343 @@
import * as fs from 'fs';
import * as path from 'path';
import * as os from 'os';
import type {
FilePermissionConfig,
FilePermissionContext,
PermissionCheckResult,
PermissionDecision,
PermissionContext,
} from '../types.js';
import type { PermissionChecker } from './base.js';
import { promptFilePermission } from '../file-prompt.js';
const CONFIG_DIR = path.join(os.homedir(), '.ai-terminal-assistant');
const FILE_PERMISSION_FILE = path.join(CONFIG_DIR, 'file-permissions.json');
// 默认文件权限配置
const DEFAULT_CONFIG: FilePermissionConfig = {
externalDirectory: 'ask',
operations: {
read: 'allow', // 读取默认允许
write: 'ask', // 写入需要确认
edit: 'ask', // 编辑需要确认
list: 'allow', // 列目录默认允许
search: 'allow', // 搜索默认允许
grep: 'allow', // 内容搜索默认允许
info: 'allow', // 获取文件信息默认允许
move: 'ask', // 移动需要确认
copy: 'ask', // 复制需要确认
delete: 'ask', // 删除需要确认
mkdir: 'ask', // 创建目录需要确认
},
sensitivePaths: [
// 系统关键路径 - 拒绝
{ pattern: '/etc/*', action: 'deny' },
{ pattern: '/usr/*', action: 'deny' },
{ pattern: '/bin/*', action: 'deny' },
{ pattern: '/sbin/*', action: 'deny' },
{ pattern: '/System/*', action: 'deny' },
{ pattern: '/var/*', action: 'deny' },
{ pattern: 'C:\\Windows\\*', action: 'deny' },
{ pattern: 'C:\\Program Files\\*', action: 'deny' },
// 用户敏感文件 - 需要确认
{ pattern: '*/.ssh/*', action: 'ask' },
{ pattern: '*/.gnupg/*', action: 'ask' },
{ pattern: '*/.aws/*', action: 'ask' },
{ pattern: '*/.kube/*', action: 'ask' },
{ pattern: '*/.env', action: 'ask' },
{ pattern: '*/.env.*', action: 'ask' },
{ pattern: '*/credentials*', action: 'ask' },
{ pattern: '*/secrets*', action: 'ask' },
{ pattern: '*/.git/config', action: 'ask' },
],
};
/**
* 文件操作权限检查器
*/
export class FilePermissionChecker implements PermissionChecker {
readonly name = 'file';
private config: FilePermissionConfig;
private projectRoot: string;
private askCallback?: (ctx: PermissionContext) => Promise<PermissionDecision>;
private sessionPermissions = new Map<string, 'allow' | 'deny'>();
constructor(projectRoot: string = process.cwd()) {
this.projectRoot = path.resolve(projectRoot);
this.config = this.loadConfig();
}
/**
* 设置权限询问回调
*/
setAskCallback(callback: (ctx: PermissionContext) => Promise<PermissionDecision>): void {
this.askCallback = callback;
}
/**
* 加载权限配置
*/
private loadConfig(): FilePermissionConfig {
try {
if (fs.existsSync(FILE_PERMISSION_FILE)) {
const content = fs.readFileSync(FILE_PERMISSION_FILE, 'utf-8');
const stored = JSON.parse(content) as Partial<FilePermissionConfig>;
return {
externalDirectory: stored.externalDirectory || DEFAULT_CONFIG.externalDirectory,
operations: { ...DEFAULT_CONFIG.operations, ...stored.operations },
sensitivePaths: stored.sensitivePaths || DEFAULT_CONFIG.sensitivePaths,
};
}
} catch {
// 忽略加载错误
}
return { ...DEFAULT_CONFIG, operations: { ...DEFAULT_CONFIG.operations } };
}
/**
* 保存权限配置
*/
saveConfig(): void {
if (!fs.existsSync(CONFIG_DIR)) {
fs.mkdirSync(CONFIG_DIR, { recursive: true });
}
fs.writeFileSync(FILE_PERMISSION_FILE, JSON.stringify(this.config, null, 2));
}
/**
* 检查路径是否在项目目录内
*/
private isInProjectDirectory(targetPath: string): boolean {
const resolved = path.resolve(targetPath);
return resolved.startsWith(this.projectRoot + path.sep) || resolved === this.projectRoot;
}
/**
* 将通配符模式转换为正则表达式
*/
private patternToRegex(pattern: string): RegExp {
const escaped = pattern
.replace(/[.+^${}()|[\]\\]/g, '\\$&')
.replace(/\*/g, '.*')
.replace(/\?/g, '.');
return new RegExp(`^${escaped}$`, 'i');
}
/**
* 检查路径是否匹配敏感路径规则
*/
private matchSensitivePath(targetPath: string): { matched: boolean; action?: 'allow' | 'deny' | 'ask' } {
const normalizedPath = targetPath.replace(/\\/g, '/');
for (const rule of this.config.sensitivePaths) {
const regex = this.patternToRegex(rule.pattern);
if (regex.test(normalizedPath)) {
return { matched: true, action: rule.action };
}
}
return { matched: false };
}
/**
* 展开路径中的 ~ 符号
*/
private expandTilde(targetPath: string): string {
if (targetPath.startsWith('~/')) {
return path.join(os.homedir(), targetPath.slice(2));
}
if (targetPath === '~') {
return os.homedir();
}
return targetPath;
}
/**
* 检查文件操作权限
*/
async checkFilePermission(ctx: FilePermissionContext): Promise<PermissionCheckResult> {
const { operation, path: targetPath, workdir } = ctx;
// 展开 ~ 并解析绝对路径
const expandedPath = this.expandTilde(targetPath);
const absolutePath = path.isAbsolute(expandedPath)
? expandedPath
: path.resolve(workdir, expandedPath);
// 生成会话权限 key
const sessionKey = `${operation}:${absolutePath}`;
// 1. 检查会话级别的临时权限
const sessionPerm = this.sessionPermissions.get(sessionKey);
if (sessionPerm === 'deny') {
return {
allowed: false,
action: 'deny',
reason: `本次会话已拒绝此操作: ${operation} ${absolutePath}`,
};
}
if (sessionPerm === 'allow') {
return {
allowed: true,
action: 'allow',
reason: '本次会话已允许此操作',
};
}
// 2. 检查敏感路径规则
const sensitiveMatch = this.matchSensitivePath(absolutePath);
if (sensitiveMatch.matched && sensitiveMatch.action) {
if (sensitiveMatch.action === 'deny') {
return {
allowed: false,
action: 'deny',
reason: `路径匹配敏感路径规则,禁止访问: ${absolutePath}`,
};
}
if (sensitiveMatch.action === 'ask') {
return this.handleAskPermission(ctx, absolutePath, sessionKey, '敏感路径需要确认');
}
// allow 则继续检查
}
// 3. 检查外部目录访问
if (!this.isInProjectDirectory(absolutePath)) {
const extAction = this.config.externalDirectory;
if (extAction === 'deny') {
return {
allowed: false,
action: 'deny',
reason: `禁止访问项目目录外的路径: ${absolutePath}`,
};
}
if (extAction === 'ask') {
return this.handleAskPermission(ctx, absolutePath, sessionKey, '访问项目外部路径需要确认');
}
}
// 4. 根据操作类型的默认策略处理
const operationAction = this.config.operations[operation];
if (operationAction === 'allow') {
return {
allowed: true,
action: 'allow',
reason: `${operation} 操作默认允许`,
};
}
if (operationAction === 'deny') {
return {
allowed: false,
action: 'deny',
reason: `${operation} 操作默认拒绝`,
};
}
// operationAction === 'ask'
return this.handleAskPermission(ctx, absolutePath, sessionKey, `${operation} 操作需要确认`);
}
/**
* 处理需要询问的权限
*/
private async handleAskPermission(
ctx: FilePermissionContext,
absolutePath: string,
sessionKey: string,
reason: string
): Promise<PermissionCheckResult> {
// 对于 write/edit 操作,如果有内容信息,使用 diff 显示
if ((ctx.operation === 'write' || ctx.operation === 'edit') && ctx.newContent !== undefined) {
// 更新 ctx 中的路径为绝对路径
const ctxWithAbsPath: FilePermissionContext = {
...ctx,
path: absolutePath,
};
const decision = await promptFilePermission(ctxWithAbsPath);
if (decision.remember) {
this.sessionPermissions.set(sessionKey, decision.allow ? 'allow' : 'deny');
}
return {
allowed: decision.allow,
action: decision.allow ? 'allow' : 'deny',
reason: decision.allow ? '用户允许' : '用户拒绝',
};
}
// 其他操作使用原有的回调
if (!this.askCallback) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason,
patterns: [absolutePath],
};
}
// 构造兼容的 PermissionContext
const permCtx: PermissionContext = {
command: `${ctx.operation} ${ctx.path}`,
workdir: ctx.workdir,
patterns: [ctx.operation],
externalPaths: this.isInProjectDirectory(absolutePath) ? undefined : [absolutePath],
};
const decision = await this.askCallback(permCtx);
if (decision.remember) {
this.sessionPermissions.set(sessionKey, decision.allow ? 'allow' : 'deny');
}
return {
allowed: decision.allow,
action: decision.allow ? 'allow' : 'deny',
reason: decision.allow ? '用户允许' : '用户拒绝',
};
}
/**
* 实现 PermissionChecker 接口的 check 方法
* 从 PermissionContext 提取文件操作信息
*/
async check(ctx: PermissionContext): Promise<PermissionCheckResult> {
// 从 command 解析操作类型和路径
// 格式: "operation path" 如 "read /path/to/file"
const parts = ctx.command.split(' ');
const operation = parts[0] as FilePermissionContext['operation'];
const targetPath = parts.slice(1).join(' ');
return this.checkFilePermission({
operation,
path: targetPath,
workdir: ctx.workdir,
});
}
/**
* 清除会话权限
*/
clearSessionPermissions(): void {
this.sessionPermissions.clear();
}
/**
* 获取当前配置
*/
getConfig(): FilePermissionConfig {
return {
...this.config,
operations: { ...this.config.operations },
sensitivePaths: [...this.config.sensitivePaths],
};
}
}
@@ -0,0 +1,236 @@
import type {
GitPermissionConfig,
GitPermissionContext,
GitOperation,
PermissionCheckResult,
PermissionDecision,
PermissionContext,
} from '../types.js';
import type { PermissionChecker } from './base.js';
// 读操作列表
const READ_OPERATIONS: GitOperation[] = [
'status',
'diff',
'log',
'branch_list',
'show',
];
// 写操作列表
const WRITE_OPERATIONS: GitOperation[] = [
'add',
'commit',
'push',
'pull',
'checkout',
'branch_create',
'branch_delete',
'stash',
'stash_pop',
'merge',
'rebase',
];
// 危险操作(需要 force 参数的操作)
const DANGEROUS_WHEN_FORCED: GitOperation[] = [
'push',
'reset',
'checkout',
'rebase',
];
// 默认 Git 权限配置
const DEFAULT_CONFIG: GitPermissionConfig = {
readOperations: 'allow',
writeOperations: 'ask',
dangerousOperations: 'ask',
};
/**
* Git 操作权限检查器
* 控制 Git 仓库操作的权限
*/
export class GitPermissionChecker implements PermissionChecker {
readonly name = 'git';
private config: GitPermissionConfig;
private askCallback?: (ctx: PermissionContext) => Promise<PermissionDecision>;
private sessionPermissions = new Map<string, 'allow' | 'deny'>();
constructor() {
this.config = { ...DEFAULT_CONFIG };
}
/**
* 设置权限询问回调
*/
setAskCallback(callback: (ctx: PermissionContext) => Promise<PermissionDecision>): void {
this.askCallback = callback;
}
/**
* 判断操作类型
*/
private getOperationType(operation: GitOperation, force?: boolean): 'read' | 'write' | 'dangerous' {
// 强制操作属于危险操作
if (force && DANGEROUS_WHEN_FORCED.includes(operation)) {
return 'dangerous';
}
// reset 总是危险操作
if (operation === 'reset') {
return 'dangerous';
}
if (READ_OPERATIONS.includes(operation)) {
return 'read';
}
return 'write';
}
/**
* 获取操作的描述
*/
private getOperationDescription(ctx: GitPermissionContext): string {
const { operation, target, remote, force, message } = ctx;
const parts: string[] = [`git ${operation.replace('_', ' ')}`];
if (force) {
parts.push('--force');
}
if (target) {
parts.push(target);
}
if (remote) {
parts.push(`(remote: ${remote})`);
}
if (message) {
parts.push(`"${message.substring(0, 50)}${message.length > 50 ? '...' : ''}"`);
}
return parts.join(' ');
}
/**
* 检查 Git 操作权限
*/
async checkGitPermission(ctx: GitPermissionContext): Promise<PermissionCheckResult> {
const { operation, force } = ctx;
const operationType = this.getOperationType(operation, force);
const description = this.getOperationDescription(ctx);
// 1. 检查会话级别的临时权限
const sessionKey = `git_${operationType}`;
const sessionPerm = this.sessionPermissions.get(sessionKey);
if (sessionPerm === 'allow') {
return {
allowed: true,
action: 'allow',
reason: `本次会话已允许 Git ${operationType === 'read' ? '读' : operationType === 'write' ? '写' : '危险'}操作`,
};
}
if (sessionPerm === 'deny') {
return {
allowed: false,
action: 'deny',
reason: `本次会话已拒绝 Git ${operationType === 'read' ? '读' : operationType === 'write' ? '写' : '危险'}操作`,
};
}
// 2. 根据操作类型确定权限策略
let action = this.config.writeOperations;
if (operationType === 'read') {
action = this.config.readOperations;
} else if (operationType === 'dangerous') {
action = this.config.dangerousOperations;
}
// 3. 处理权限决策
if (action === 'allow') {
return {
allowed: true,
action: 'allow',
reason: operationType === 'read' ? '读操作默认允许' : '配置允许',
};
}
if (action === 'deny') {
return {
allowed: false,
action: 'deny',
reason: operationType === 'dangerous' ? '危险操作默认拒绝' : '配置拒绝',
};
}
// action === 'ask'
if (!this.askCallback) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason: description,
};
}
// 调用回调询问用户
const decision = await this.askCallback({
command: description,
workdir: process.cwd(),
});
if (decision.remember) {
this.sessionPermissions.set(sessionKey, decision.allow ? 'allow' : 'deny');
}
return {
allowed: decision.allow,
action: decision.allow ? 'allow' : 'deny',
reason: decision.allow ? '用户允许' : '用户拒绝',
};
}
/**
* 实现 PermissionChecker 接口的 check 方法
*/
async check(ctx: PermissionContext): Promise<PermissionCheckResult> {
// 从 command 中解析 Git 操作
const match = ctx.command.match(/^git[_\s]+(\w+)/);
if (!match) {
return {
allowed: false,
action: 'deny',
reason: '无法解析 Git 操作',
};
}
const operation = match[1] as GitOperation;
return this.checkGitPermission({ operation });
}
/**
* 清除会话权限
*/
clearSessionPermissions(): void {
this.sessionPermissions.clear();
}
/**
* 获取当前配置
*/
getConfig(): GitPermissionConfig {
return { ...this.config };
}
/**
* 更新配置
*/
setConfig(config: Partial<GitPermissionConfig>): void {
this.config = { ...this.config, ...config };
}
}
@@ -0,0 +1,5 @@
export type { PermissionChecker, BasePermissionConfig } from './base.js';
export { BashPermissionChecker } from './bash.js';
export { FilePermissionChecker } from './file.js';
export { WebPermissionChecker } from './web.js';
export { GitPermissionChecker } from './git.js';
@@ -0,0 +1,159 @@
import type {
WebPermissionConfig,
WebPermissionContext,
PermissionCheckResult,
PermissionDecision,
PermissionContext,
} from '../types.js';
import type { PermissionChecker } from './base.js';
// 默认 Web 权限配置
const DEFAULT_CONFIG: WebPermissionConfig = {
default: 'ask', // 默认需要确认
allowAdvancedSearch: true,
allowedTopics: [], // 空数组表示允许所有主题
};
/**
* Web 搜索权限检查器
* 控制网络搜索操作的权限
*/
export class WebPermissionChecker implements PermissionChecker {
readonly name = 'web';
private config: WebPermissionConfig;
private askCallback?: (ctx: PermissionContext) => Promise<PermissionDecision>;
private sessionPermissions = new Map<string, 'allow' | 'deny'>();
constructor() {
this.config = { ...DEFAULT_CONFIG };
}
/**
* 设置权限询问回调
*/
setAskCallback(callback: (ctx: PermissionContext) => Promise<PermissionDecision>): void {
this.askCallback = callback;
}
/**
* 检查 Web 搜索权限
*/
async checkWebPermission(ctx: WebPermissionContext): Promise<PermissionCheckResult> {
const { query, searchDepth, topic } = ctx;
// 1. 检查深度搜索权限
if (searchDepth === 'advanced' && !this.config.allowAdvancedSearch) {
return {
allowed: false,
action: 'deny',
reason: '不允许深度搜索',
};
}
// 2. 检查主题限制
if (this.config.allowedTopics.length > 0 && topic) {
if (!this.config.allowedTopics.includes(topic)) {
return {
allowed: false,
action: 'deny',
reason: `不允许搜索主题: ${topic}`,
};
}
}
// 3. 检查会话级别的临时权限
const sessionKey = `web_search`;
const sessionPerm = this.sessionPermissions.get(sessionKey);
if (sessionPerm === 'allow') {
return {
allowed: true,
action: 'allow',
reason: '本次会话已允许网络搜索',
};
}
if (sessionPerm === 'deny') {
return {
allowed: false,
action: 'deny',
reason: '本次会话已拒绝网络搜索',
};
}
// 4. 根据默认策略处理
const action = this.config.default;
if (action === 'allow') {
return {
allowed: true,
action: 'allow',
reason: '默认允许网络搜索',
};
}
if (action === 'deny') {
return {
allowed: false,
action: 'deny',
reason: '默认拒绝网络搜索',
};
}
// action === 'ask'
if (!this.askCallback) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason: `搜索: ${query}`,
};
}
// 调用回调询问用户
const decision = await this.askCallback({
command: `web_search: ${query}`,
workdir: process.cwd(),
});
if (decision.remember) {
this.sessionPermissions.set(sessionKey, decision.allow ? 'allow' : 'deny');
}
return {
allowed: decision.allow,
action: decision.allow ? 'allow' : 'deny',
reason: decision.allow ? '用户允许' : '用户拒绝',
};
}
/**
* 实现 PermissionChecker 接口的 check 方法
* 从通用 PermissionContext 中提取 Web 搜索信息
*/
async check(ctx: PermissionContext): Promise<PermissionCheckResult> {
// 从 command 中提取搜索查询
const query = ctx.command.replace(/^web_search:\s*/, '');
return this.checkWebPermission({ query });
}
/**
* 清除会话权限
*/
clearSessionPermissions(): void {
this.sessionPermissions.clear();
}
/**
* 获取当前配置
*/
getConfig(): WebPermissionConfig {
return { ...this.config };
}
/**
* 更新配置
*/
setConfig(config: Partial<WebPermissionConfig>): void {
this.config = { ...this.config, ...config };
}
}
+186
View File
@@ -0,0 +1,186 @@
/**
* 文件操作确认提示
* 显示 diff 对比并让用户确认
*/
import * as readline from 'readline';
import * as fs from 'fs/promises';
import chalk from 'chalk';
import type { FilePermissionContext, PermissionDecision } from './types.js';
import { computeDiff, formatDiff, countChanges, formatEditDiff } from '../utils/diff.js';
/**
* 显示文件写入的 diff 并请求确认
*/
export async function promptFileWrite(ctx: FilePermissionContext): Promise<PermissionDecision> {
const { path: filePath, newContent } = ctx;
if (!newContent) {
// 没有内容,使用简单确认
return promptSimpleConfirm(ctx);
}
// 读取原文件内容
let oldContent: string | null = null;
try {
oldContent = await fs.readFile(filePath, 'utf-8');
} catch {
// 文件不存在,是新文件
}
// 如果内容相同,直接允许
if (oldContent === newContent) {
return { allow: true, remember: false };
}
// 计算 diff
const diff = computeDiff(oldContent, newContent);
const changes = countChanges(diff);
// 显示 diff
console.log('');
console.log(chalk.yellow('📝 文件写入预览'));
console.log(chalk.cyan('文件: ') + chalk.white(filePath));
if (diff.isNew) {
console.log(chalk.green('状态: ') + chalk.white('新文件'));
console.log(chalk.green(`+${changes.additions} 行`));
} else {
console.log(chalk.green(`+${changes.additions} 行`) + ' / ' + chalk.red(`-${changes.deletions} 行`));
}
console.log('');
console.log(chalk.gray('─'.repeat(60)));
// 限制显示行数
const diffOutput = formatDiff(diff, filePath);
const lines = diffOutput.split('\n');
const MAX_LINES = 50;
if (lines.length > MAX_LINES) {
console.log(lines.slice(0, MAX_LINES).join('\n'));
console.log(chalk.yellow(`\n... 省略 ${lines.length - MAX_LINES} 行 ...`));
} else {
console.log(diffOutput);
}
console.log(chalk.gray('─'.repeat(60)));
console.log('');
// 询问用户确认
return promptConfirm();
}
/**
* 显示文件编辑的 diff 并请求确认
*/
export async function promptFileEdit(ctx: FilePermissionContext): Promise<PermissionDecision> {
const { path: filePath, oldContent, newContent } = ctx;
if (!oldContent || !newContent) {
// 没有内容,使用简单确认
return promptSimpleConfirm(ctx);
}
// 显示编辑 diff
console.log('');
console.log(chalk.yellow('✏️ 文件编辑预览'));
console.log(chalk.cyan('文件: ') + chalk.white(filePath));
console.log('');
console.log(chalk.gray('─'.repeat(60)));
console.log(formatEditDiff(oldContent, newContent));
console.log(chalk.gray('─'.repeat(60)));
console.log('');
// 询问用户确认
return promptConfirm();
}
/**
* 简单确认(无 diff)
*/
async function promptSimpleConfirm(ctx: FilePermissionContext): Promise<PermissionDecision> {
const rl = readline.createInterface({
input: process.stdin,
output: process.stdout,
});
return new Promise((resolve) => {
console.log('');
console.log(chalk.yellow('⚠️ 文件操作确认'));
console.log(chalk.cyan('操作: ') + chalk.white(ctx.operation));
console.log(chalk.cyan('文件: ') + chalk.white(ctx.path));
console.log('');
showConfirmOptions();
rl.question(chalk.yellow('请选择 [y/Y/n/N]: '), (answer) => {
rl.close();
resolve(parseAnswer(answer));
});
});
}
/**
* 通用确认提示
*/
async function promptConfirm(): Promise<PermissionDecision> {
const rl = readline.createInterface({
input: process.stdin,
output: process.stdout,
});
return new Promise((resolve) => {
showConfirmOptions();
rl.question(chalk.yellow('请选择 [y/Y/n/N]: '), (answer) => {
rl.close();
resolve(parseAnswer(answer));
});
});
}
/**
* 显示确认选项
*/
function showConfirmOptions(): void {
console.log(chalk.white('选择操作:'));
console.log(chalk.green(' [y] ') + '确认执行');
console.log(chalk.green(' [Y] ') + '确认执行,并记住此类操作(本次会话)');
console.log(chalk.red(' [n] ') + '拒绝执行');
console.log(chalk.red(' [N] ') + '拒绝执行,并记住此类操作(本次会话)');
console.log('');
}
/**
* 解析用户输入
*/
function parseAnswer(answer: string): PermissionDecision {
const choice = answer.trim();
switch (choice) {
case 'y':
return { allow: true, remember: false };
case 'Y':
return { allow: true, remember: true };
case 'N':
return { allow: false, remember: true };
case 'n':
default:
return { allow: false, remember: false };
}
}
/**
* 根据操作类型选择合适的确认提示
*/
export async function promptFilePermission(ctx: FilePermissionContext): Promise<PermissionDecision> {
switch (ctx.operation) {
case 'write':
return promptFileWrite(ctx);
case 'edit':
return promptFileEdit(ctx);
default:
return promptSimpleConfirm(ctx);
}
}
+24
View File
@@ -0,0 +1,24 @@
export type {
PermissionAction,
PermissionRule,
BashPermissionConfig,
PermissionContext,
PermissionCheckResult,
PermissionDecision,
FileOperation,
FilePermissionContext,
FilePermissionConfig,
} from './types.js';
export { matchPattern, matchRules, parseCommand, generateAskPattern } from './wildcard.js';
export { PermissionManager, getPermissionManager, resetPermissionManager } from './manager.js';
export { promptPermission, showPermissionDenied, showPermissionAllowed } from './prompt.js';
export { promptFilePermission, promptFileWrite, promptFileEdit } from './file-prompt.js';
// Checker pattern exports
export type { PermissionChecker, BasePermissionConfig } from './checkers/base.js';
export { BashPermissionChecker } from './checkers/bash.js';
export { FilePermissionChecker } from './checkers/file.js';
+173
View File
@@ -0,0 +1,173 @@
import type {
PermissionContext,
PermissionCheckResult,
PermissionDecision,
FilePermissionContext,
WebPermissionContext,
GitPermissionContext,
} from './types.js';
import type { PermissionChecker } from './checkers/base.js';
import { BashPermissionChecker } from './checkers/bash.js';
import { FilePermissionChecker } from './checkers/file.js';
import { WebPermissionChecker } from './checkers/web.js';
import { GitPermissionChecker } from './checkers/git.js';
/**
* 权限管理器
* 统一管理所有工具的权限检查
*/
export class PermissionManager {
private checkers = new Map<string, PermissionChecker>();
private askCallback?: (ctx: PermissionContext) => Promise<PermissionDecision>;
constructor(projectRoot: string = process.cwd()) {
// 注册默认的检查器
this.registerChecker(new BashPermissionChecker(projectRoot));
this.registerChecker(new FilePermissionChecker(projectRoot));
this.registerChecker(new WebPermissionChecker());
this.registerChecker(new GitPermissionChecker());
}
/**
* 注册权限检查器
*/
registerChecker(checker: PermissionChecker): void {
this.checkers.set(checker.name, checker);
// 如果检查器支持设置回调,传递当前的回调
if (this.askCallback && 'setAskCallback' in checker) {
(checker as BashPermissionChecker).setAskCallback(this.askCallback);
}
}
/**
* 获取权限检查器
*/
getChecker<T extends PermissionChecker>(name: string): T | undefined {
return this.checkers.get(name) as T | undefined;
}
/**
* 设置权限询问回调
*/
setAskCallback(callback: (ctx: PermissionContext) => Promise<PermissionDecision>): void {
this.askCallback = callback;
// 传递给所有支持回调的检查器
for (const checker of this.checkers.values()) {
if ('setAskCallback' in checker) {
(checker as BashPermissionChecker).setAskCallback(callback);
}
}
}
/**
* 检查权限(使用指定的检查器)
*/
async checkPermission(
checkerName: string,
ctx: PermissionContext
): Promise<PermissionCheckResult> {
const checker = this.checkers.get(checkerName);
if (!checker) {
// 未注册的检查器,默认需要确认
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason: `未找到检查器: ${checkerName}`,
};
}
return checker.check(ctx);
}
/**
* 检查 bash 命令权限(便捷方法)
*/
async checkBashPermission(ctx: PermissionContext): Promise<PermissionCheckResult> {
return this.checkPermission('bash', ctx);
}
/**
* 检查文件操作权限(便捷方法)
*/
async checkFilePermission(ctx: FilePermissionContext): Promise<PermissionCheckResult> {
const fileChecker = this.getChecker<FilePermissionChecker>('file');
if (!fileChecker) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason: '文件权限检查器未注册',
};
}
return fileChecker.checkFilePermission(ctx);
}
/**
* 检查 Web 搜索权限(便捷方法)
*/
async checkWebPermission(ctx: WebPermissionContext): Promise<PermissionCheckResult> {
const webChecker = this.getChecker<WebPermissionChecker>('web');
if (!webChecker) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason: 'Web 权限检查器未注册',
};
}
return webChecker.checkWebPermission(ctx);
}
/**
* 检查 Git 操作权限(便捷方法)
*/
async checkGitPermission(ctx: GitPermissionContext): Promise<PermissionCheckResult> {
const gitChecker = this.getChecker<GitPermissionChecker>('git');
if (!gitChecker) {
return {
allowed: false,
action: 'ask',
needsConfirmation: true,
reason: 'Git 权限检查器未注册',
};
}
return gitChecker.checkGitPermission(ctx);
}
/**
* 清除所有检查器的会话权限
*/
clearAllSessionPermissions(): void {
for (const checker of this.checkers.values()) {
checker.clearSessionPermissions();
}
}
/**
* 清除指定检查器的会话权限
*/
clearSessionPermissions(checkerName: string): void {
const checker = this.checkers.get(checkerName);
if (checker) {
checker.clearSessionPermissions();
}
}
}
// 全局实例
let globalManager: PermissionManager | null = null;
export function getPermissionManager(projectRoot?: string): PermissionManager {
if (!globalManager) {
globalManager = new PermissionManager(projectRoot);
}
return globalManager;
}
export function resetPermissionManager(): void {
globalManager = null;
}
+79
View File
@@ -0,0 +1,79 @@
import * as readline from 'readline';
import chalk from 'chalk';
import type { PermissionContext, PermissionDecision } from './types.js';
/**
* 在终端中提示用户确认权限
*/
export async function promptPermission(ctx: PermissionContext): Promise<PermissionDecision> {
const rl = readline.createInterface({
input: process.stdin,
output: process.stdout,
});
return new Promise((resolve) => {
console.log('');
console.log(chalk.yellow('⚠️ 权限确认'));
console.log(chalk.cyan('命令: ') + chalk.white(ctx.command));
console.log(chalk.cyan('目录: ') + chalk.gray(ctx.workdir));
if (ctx.externalPaths && ctx.externalPaths.length > 0) {
console.log(chalk.red('⚠️ 此命令访问项目目录外的路径:'));
ctx.externalPaths.forEach(p => {
console.log(chalk.red(' • ') + chalk.gray(p));
});
}
if (ctx.patterns && ctx.patterns.length > 0) {
console.log(chalk.gray('匹配模式: ') + ctx.patterns.join(', '));
}
console.log('');
console.log(chalk.white('选择操作:'));
console.log(chalk.green(' [y] ') + '允许执行');
console.log(chalk.green(' [Y] ') + '允许执行,并记住此类命令(本次会话)');
console.log(chalk.red(' [n] ') + '拒绝执行');
console.log(chalk.red(' [N] ') + '拒绝执行,并记住此类命令(本次会话)');
console.log('');
rl.question(chalk.yellow('请选择 [y/Y/n/N]: '), (answer) => {
rl.close();
const choice = answer.trim();
switch (choice) {
case 'y':
resolve({ allow: true, remember: false });
break;
case 'Y':
resolve({ allow: true, remember: true });
break;
case 'N':
resolve({ allow: false, remember: true });
break;
case 'n':
default:
resolve({ allow: false, remember: false });
break;
}
});
});
}
/**
* 显示权限被拒绝的消息
*/
export function showPermissionDenied(command: string, reason: string): void {
console.log('');
console.log(chalk.red('🚫 权限被拒绝'));
console.log(chalk.cyan('命令: ') + chalk.white(command));
console.log(chalk.cyan('原因: ') + chalk.gray(reason));
console.log('');
}
/**
* 显示权限允许的消息
*/
export function showPermissionAllowed(command: string): void {
console.log(chalk.green('✓ ') + chalk.gray(`执行: ${command}`));
}
+149
View File
@@ -0,0 +1,149 @@
// 权限动作类型
export type PermissionAction = 'allow' | 'deny' | 'ask';
// 单条权限规则
export interface PermissionRule {
pattern: string; // 匹配模式,如 "git push *", "rm *"
action: PermissionAction;
}
// Bash 命令权限配置
export interface BashPermissionConfig {
// 命令规则列表,按顺序匹配
rules: PermissionRule[];
// 外部目录访问策略
externalDirectory: PermissionAction;
// 默认策略(没有规则匹配时)
default: PermissionAction;
}
// 权限请求上下文
export interface PermissionContext {
command: string;
workdir: string;
patterns?: string[]; // 匹配到的模式
externalPaths?: string[]; // 访问的外部路径
}
// 文件操作类型
export type FileOperation =
| 'read' // 读取文件
| 'write' // 写入文件
| 'edit' // 编辑文件
| 'list' // 列出目录
| 'search' // 搜索文件
| 'grep' // 搜索内容
| 'info' // 获取文件信息
| 'move' // 移动/重命名
| 'copy' // 复制
| 'delete' // 删除
| 'mkdir'; // 创建目录
// 文件权限请求上下文
export interface FilePermissionContext {
operation: FileOperation;
path: string; // 目标路径
workdir: string; // 当前工作目录
// 用于 diff 显示的内容(仅 write/edit 操作)
newContent?: string; // 新内容
oldContent?: string; // 原内容(edit 操作时,要被替换的部分)
}
// 文件权限配置
export interface FilePermissionConfig {
// 外部目录访问策略
externalDirectory: PermissionAction;
// 各操作的默认策略
operations: {
read: PermissionAction;
write: PermissionAction;
edit: PermissionAction;
list: PermissionAction;
search: PermissionAction;
grep: PermissionAction;
info: PermissionAction;
move: PermissionAction;
copy: PermissionAction;
delete: PermissionAction;
mkdir: PermissionAction;
};
// 敏感路径规则(优先于操作默认策略)
sensitivePaths: {
pattern: string;
action: PermissionAction;
}[];
}
// 权限检查结果
export interface PermissionCheckResult {
allowed: boolean;
action: PermissionAction;
reason?: string;
needsConfirmation?: boolean;
patterns?: string[];
}
// 用户权限决定(用于 ask 时的回调)
export interface PermissionDecision {
allow: boolean;
remember?: boolean; // 是否记住这个决定
}
// Web 搜索权限请求上下文
export interface WebPermissionContext {
query: string; // 搜索查询
searchDepth?: 'basic' | 'advanced'; // 搜索深度
topic?: 'general' | 'news' | 'finance'; // 搜索主题
maxResults?: number; // 最大结果数
}
// Web 权限配置
export interface WebPermissionConfig {
// 默认策略
default: PermissionAction;
// 是否允许深度搜索
allowAdvancedSearch: boolean;
// 搜索主题限制(空数组表示允许所有)
allowedTopics: ('general' | 'news' | 'finance')[];
}
// Git 操作类型
export type GitOperation =
// 读操作
| 'status'
| 'diff'
| 'log'
| 'branch_list'
| 'show'
// 写操作
| 'add'
| 'commit'
| 'push'
| 'pull'
| 'checkout'
| 'branch_create'
| 'branch_delete'
| 'stash'
| 'stash_pop'
| 'reset'
| 'merge'
| 'rebase';
// Git 权限请求上下文
export interface GitPermissionContext {
operation: GitOperation;
target?: string; // 分支名、文件路径等
remote?: string; // 远程仓库名
force?: boolean; // 是否强制操作
message?: string; // 提交信息等
}
// Git 权限配置
export interface GitPermissionConfig {
// 读操作策略(默认 allow)
readOperations: PermissionAction;
// 写操作策略(默认 ask)
writeOperations: PermissionAction;
// 危险操作策略(force push, reset --hard 等,默认 ask)
dangerousOperations: PermissionAction;
}
+157
View File
@@ -0,0 +1,157 @@
import type { PermissionAction, PermissionRule } from './types.js';
import { parseBashCommand, parseCommandSimple, type ParsedCommand } from './bash-parser.js';
/**
* 将通配符模式转换为正则表达式
* 支持 * 匹配任意字符
*/
function patternToRegex(pattern: string): RegExp {
const escaped = pattern
.replace(/[.+^${}()|[\]\\]/g, '\\$&') // 转义特殊字符
.replace(/\*/g, '.*') // * -> .*
.replace(/\?/g, '.'); // ? -> .
return new RegExp(`^${escaped}$`, 'i');
}
/**
* 检查命令是否匹配模式
*/
export function matchPattern(command: string, pattern: string): boolean {
const regex = patternToRegex(pattern);
return regex.test(command);
}
/**
* 从命令中提取命令名和子命令(简单版本,用于向后兼容)
*/
export function parseCommand(command: string): { head: string; sub?: string; args: string[] } {
const parsed = parseCommandSimple(command);
return {
head: parsed.name,
sub: parsed.subcommand,
args: parsed.args,
};
}
/**
* 生成用于权限请求的模式
* 例如: "git push origin" -> "git push *"
*/
export function generateAskPattern(command: string): string {
const parsed = parseCommandSimple(command);
if (parsed.subcommand) {
return `${parsed.name} ${parsed.subcommand} *`;
}
return `${parsed.name} *`;
}
/**
* 生成用于权限请求的模式(从 ParsedCommand)
*/
export function generateAskPatternFromParsed(parsed: ParsedCommand): string {
if (parsed.subcommand) {
return `${parsed.name} ${parsed.subcommand} *`;
}
return `${parsed.name} *`;
}
/**
* 检查单个命令是否匹配规则
*/
function matchSingleCommand(
parsed: ParsedCommand,
rules: PermissionRule[],
defaultAction: PermissionAction
): { action: PermissionAction; matchedPattern?: string } {
// 构建可能的匹配字符串
const candidates = [
parsed.text, // 完整命令
parsed.subcommand
? `${parsed.name} ${parsed.subcommand}`
: parsed.name, // 命令 + 子命令
parsed.name, // 仅命令名
];
for (const rule of rules) {
for (const candidate of candidates) {
if (matchPattern(candidate, rule.pattern)) {
return { action: rule.action, matchedPattern: rule.pattern };
}
}
}
return { action: defaultAction };
}
/**
* 检查命令是否匹配规则列表,返回对应的动作(同步版本,使用简单解析)
*/
export function matchRules(
command: string,
rules: PermissionRule[],
defaultAction: PermissionAction
): { action: PermissionAction; matchedPattern?: string } {
const parsed = parseCommandSimple(command);
return matchSingleCommand(parsed, rules, defaultAction);
}
/**
* 使用 tree-sitter 解析并检查所有命令的权限(异步版本)
* 返回最严格的权限要求
*/
export async function matchRulesAsync(
command: string,
rules: PermissionRule[],
defaultAction: PermissionAction
): Promise<{
action: PermissionAction;
matchedPattern?: string;
allCommands: ParsedCommand[];
askPatterns: string[];
}> {
const result = await parseBashCommand(command);
// 如果解析失败,降级到简单解析
if (!result.success || result.commands.length === 0) {
const parsed = parseCommandSimple(command);
const match = matchSingleCommand(parsed, rules, defaultAction);
return {
...match,
allCommands: [parsed],
askPatterns: match.action === 'ask' ? [generateAskPatternFromParsed(parsed)] : [],
};
}
// 检查所有命令,收集结果
let finalAction: PermissionAction = 'allow';
let finalPattern: string | undefined;
const askPatterns: string[] = [];
for (const parsed of result.commands) {
// 跳过 cd 命令(如果通过了外部目录检查)
if (parsed.name === 'cd') {
continue;
}
const match = matchSingleCommand(parsed, rules, defaultAction);
// 权限优先级: deny > ask > allow
if (match.action === 'deny') {
finalAction = 'deny';
finalPattern = match.matchedPattern;
break; // deny 直接终止
} else if (match.action === 'ask' && finalAction === 'allow') {
finalAction = 'ask';
finalPattern = match.matchedPattern;
askPatterns.push(generateAskPatternFromParsed(parsed));
}
// allow 不改变已有的更严格权限
}
return {
action: finalAction,
matchedPattern: finalPattern,
allCommands: result.commands,
askPatterns: [...new Set(askPatterns)], // 去重
};
}
+218
View File
@@ -0,0 +1,218 @@
/**
* 磁盘缓存实现
* 使用 JSON 文件存储,支持按文件路径索引
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import { createHash } from 'crypto';
export interface CacheEntry<T> {
key: string;
value: T;
timestamp: number;
}
/**
* 磁盘缓存类
*/
export class DiskCache<T> {
private cacheDir: string;
private memoryCache: Map<string, T> = new Map();
private dirty: Set<string> = new Set();
private initialized = false;
constructor(cacheDir: string) {
this.cacheDir = cacheDir;
}
/**
* 初始化缓存目录
*/
private async ensureDir(): Promise<void> {
if (this.initialized) return;
try {
await fs.mkdir(this.cacheDir, { recursive: true });
this.initialized = true;
} catch (error) {
console.warn(`Failed to create cache directory: ${this.cacheDir}`, error);
}
}
/**
* 生成缓存文件路径
*/
private getCacheFilePath(key: string): string {
// 使用哈希避免文件名过长或包含特殊字符
const hash = createHash('md5').update(key).digest('hex');
return path.join(this.cacheDir, `${hash}.json`);
}
/**
* 获取缓存值
*/
async get(key: string): Promise<T | null> {
// 先检查内存缓存
if (this.memoryCache.has(key)) {
return this.memoryCache.get(key)!;
}
await this.ensureDir();
// 从磁盘读取
const filePath = this.getCacheFilePath(key);
try {
const content = await fs.readFile(filePath, 'utf-8');
const entry: CacheEntry<T> = JSON.parse(content);
// 验证 key 匹配
if (entry.key === key) {
this.memoryCache.set(key, entry.value);
return entry.value;
}
} catch {
// 文件不存在或解析失败
}
return null;
}
/**
* 设置缓存值
*/
async set(key: string, value: T): Promise<void> {
this.memoryCache.set(key, value);
this.dirty.add(key);
// 立即写入磁盘(可以优化为批量写入)
await this.flush(key);
}
/**
* 删除缓存
*/
async delete(key: string): Promise<void> {
this.memoryCache.delete(key);
this.dirty.delete(key);
await this.ensureDir();
const filePath = this.getCacheFilePath(key);
try {
await fs.unlink(filePath);
} catch {
// 文件不存在
}
}
/**
* 检查缓存是否存在
*/
async has(key: string): Promise<boolean> {
if (this.memoryCache.has(key)) {
return true;
}
const value = await this.get(key);
return value !== null;
}
/**
* 刷新指定 key 到磁盘
*/
private async flush(key: string): Promise<void> {
if (!this.dirty.has(key)) return;
await this.ensureDir();
const value = this.memoryCache.get(key);
if (value === undefined) return;
const entry: CacheEntry<T> = {
key,
value,
timestamp: Date.now(),
};
const filePath = this.getCacheFilePath(key);
try {
await fs.writeFile(filePath, JSON.stringify(entry), 'utf-8');
this.dirty.delete(key);
} catch (error) {
console.warn(`Failed to write cache: ${key}`, error);
}
}
/**
* 刷新所有脏数据到磁盘
*/
async flushAll(): Promise<void> {
const keys = Array.from(this.dirty);
await Promise.all(keys.map((key) => this.flush(key)));
}
/**
* 清空所有缓存
*/
async clear(): Promise<void> {
this.memoryCache.clear();
this.dirty.clear();
try {
const files = await fs.readdir(this.cacheDir);
await Promise.all(
files
.filter((f) => f.endsWith('.json'))
.map((f) => fs.unlink(path.join(this.cacheDir, f)))
);
} catch {
// 目录不存在或其他错误
}
}
/**
* 获取缓存大小(条目数)
*/
size(): number {
return this.memoryCache.size;
}
/**
* 获取所有缓存的 keys
*/
async keys(): Promise<string[]> {
await this.ensureDir();
const result: string[] = Array.from(this.memoryCache.keys());
try {
const files = await fs.readdir(this.cacheDir);
for (const file of files) {
if (!file.endsWith('.json')) continue;
const filePath = path.join(this.cacheDir, file);
try {
const content = await fs.readFile(filePath, 'utf-8');
const entry: CacheEntry<T> = JSON.parse(content);
if (!result.includes(entry.key)) {
result.push(entry.key);
}
} catch {
// 跳过无效文件
}
}
} catch {
// 目录不存在
}
return result;
}
}
/**
* 创建磁盘缓存实例
*/
export function createDiskCache<T>(cacheDir: string): DiskCache<T> {
return new DiskCache<T>(cacheDir);
}
+6
View File
@@ -0,0 +1,6 @@
/**
* 缓存模块导出
*/
export { DiskCache, createDiskCache } from './disk-cache.js';
export type { CacheEntry } from './disk-cache.js';
+26
View File
@@ -0,0 +1,26 @@
/**
* RepoMap 模块
*
* 使用 AST 分析和 PageRank 算法生成代码仓库地图
* 参考 Aider 的实现
*/
// 主类
export { RepoMap, createRepoMap } from './repomap.js';
// Tag 提取
export { TagExtractor } from './tags/index.js';
// PageRank 排序
export {
Graph,
pagerank,
distributeRanksToDefinitions,
type PageRankOptions,
} from './ranking/index.js';
// 缓存
export { DiskCache, createDiskCache } from './cache/index.js';
// 类型
export type { Tag, TagCacheEntry, RepoMapConfig, GraphEdge } from './types.js';
@@ -0,0 +1,99 @@
/**
* 图数据结构
* 用于 PageRank 算法
*/
import type { GraphEdge } from '../types.js';
export class Graph {
/** 邻接表:from -> edges[] */
private outEdges: Map<string, GraphEdge[]> = new Map();
/** 反向邻接表:to -> edges[] */
private inEdges: Map<string, GraphEdge[]> = new Map();
/** 所有节点 */
private nodes: Set<string> = new Set();
/**
* 添加边
*/
addEdge(edge: GraphEdge): void {
this.nodes.add(edge.from);
this.nodes.add(edge.to);
// 出边
if (!this.outEdges.has(edge.from)) {
this.outEdges.set(edge.from, []);
}
this.outEdges.get(edge.from)!.push(edge);
// 入边
if (!this.inEdges.has(edge.to)) {
this.inEdges.set(edge.to, []);
}
this.inEdges.get(edge.to)!.push(edge);
}
/**
* 获取所有节点
*/
getNodes(): string[] {
return Array.from(this.nodes);
}
/**
* 获取节点的出边
*/
getOutEdges(node: string): GraphEdge[] {
return this.outEdges.get(node) || [];
}
/**
* 获取节点的入边
*/
getInEdges(node: string): GraphEdge[] {
return this.inEdges.get(node) || [];
}
/**
* 获取节点的出度(考虑权重)
*/
getOutDegree(node: string): number {
const edges = this.outEdges.get(node) || [];
return edges.reduce((sum, e) => sum + e.weight, 0);
}
/**
* 获取节点的入度(考虑权重)
*/
getInDegree(node: string): number {
const edges = this.inEdges.get(node) || [];
return edges.reduce((sum, e) => sum + e.weight, 0);
}
/**
* 获取边数量
*/
getEdgeCount(): number {
let count = 0;
for (const edges of this.outEdges.values()) {
count += edges.length;
}
return count;
}
/**
* 清空图
*/
clear(): void {
this.outEdges.clear();
this.inEdges.clear();
this.nodes.clear();
}
/**
* 是否为空
*/
isEmpty(): boolean {
return this.nodes.size === 0;
}
}
@@ -0,0 +1,7 @@
/**
* 排序模块导出
*/
export { Graph } from './graph.js';
export { pagerank, distributeRanksToDefinitions } from './pagerank.js';
export type { PageRankOptions } from './pagerank.js';
@@ -0,0 +1,146 @@
/**
* PageRank 算法实现
* 基于 Aider 的实现,用于代码符号相关性排序
*/
import { Graph } from './graph.js';
export interface PageRankOptions {
/** 阻尼系数 (默认 0.85) */
damping?: number;
/** 最大迭代次数 (默认 100) */
iterations?: number;
/** 收敛阈值 (默认 1e-6) */
tolerance?: number;
/** 个性化向量:节点 -> 初始权重 */
personalization?: Map<string, number>;
}
/**
* PageRank 算法
*
* @param graph - 图结构
* @param options - 算法选项
* @returns 节点排名 Map<节点, 排名值>
*/
export function pagerank(
graph: Graph,
options: PageRankOptions = {}
): Map<string, number> {
const {
damping = 0.85,
iterations = 100,
tolerance = 1e-6,
personalization,
} = options;
const nodes = graph.getNodes();
const n = nodes.length;
if (n === 0) {
return new Map();
}
// 初始化排名
let ranks = new Map<string, number>();
const baseRank = 1 / n;
// 处理个性化向量
let persVector = new Map<string, number>();
if (personalization && personalization.size > 0) {
// 归一化个性化向量
const total = Array.from(personalization.values()).reduce((a, b) => a + b, 0);
if (total > 0) {
for (const [node, value] of personalization) {
persVector.set(node, value / total);
}
}
} else {
// 均匀分布
for (const node of nodes) {
persVector.set(node, baseRank);
}
}
// 初始排名 = 个性化向量
for (const node of nodes) {
ranks.set(node, persVector.get(node) || baseRank);
}
// 迭代计算
for (let iter = 0; iter < iterations; iter++) {
const newRanks = new Map<string, number>();
let diff = 0;
// 计算悬挂节点的贡献(没有出边的节点)
let danglingSum = 0;
for (const node of nodes) {
const outEdges = graph.getOutEdges(node);
if (outEdges.length === 0) {
danglingSum += ranks.get(node) || 0;
}
}
for (const node of nodes) {
// 基础分数:(1 - damping) * 个性化 + damping * 悬挂贡献
let rank =
(1 - damping) * (persVector.get(node) || baseRank) +
(damping * danglingSum) / n;
// 收集入边贡献
const inEdges = graph.getInEdges(node);
for (const edge of inEdges) {
const sourceRank = ranks.get(edge.from) || 0;
const outDegree = graph.getOutDegree(edge.from);
if (outDegree > 0) {
// 边权重占源节点总出度的比例
rank += damping * sourceRank * (edge.weight / outDegree);
}
}
newRanks.set(node, rank);
diff += Math.abs(rank - (ranks.get(node) || 0));
}
ranks = newRanks;
// 检查收敛
if (diff < tolerance) {
break;
}
}
return ranks;
}
/**
* 将 PageRank 排名分配到定义上
* 按照 Aider 的方式:将源节点的排名按边权重比例分配给目标定义
*
* @param graph - 图结构
* @param nodeRanks - 节点 PageRank 排名
* @returns 定义排名 Map<"file:ident", rank>
*/
export function distributeRanksToDefinitions(
graph: Graph,
nodeRanks: Map<string, number>
): Map<string, number> {
const definitionRanks = new Map<string, number>();
for (const src of graph.getNodes()) {
const srcRank = nodeRanks.get(src) || 0;
const outEdges = graph.getOutEdges(src);
const totalWeight = outEdges.reduce((sum, e) => sum + e.weight, 0);
if (totalWeight === 0) continue;
for (const edge of outEdges) {
const edgeRank = (srcRank * edge.weight) / totalWeight;
const key = `${edge.to}:${edge.ident}`;
definitionRanks.set(key, (definitionRanks.get(key) || 0) + edgeRank);
}
}
return definitionRanks;
}
+419
View File
@@ -0,0 +1,419 @@
/**
* RepoMap 主类
* 使用 AST 分析和 PageRank 算法生成代码仓库地图
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import { TagExtractor } from './tags/extractor.js';
import { pagerank, distributeRanksToDefinitions } from './ranking/pagerank.js';
import { Graph } from './ranking/graph.js';
import { DiskCache } from './cache/disk-cache.js';
import type { Tag, RepoMapConfig, TagCacheEntry } from './types.js';
/**
* RepoMap 配置默认值
*/
const defaultConfig: RepoMapConfig = {
mapTokens: 1024,
mapMulNoFiles: 8,
maxContextWindow: 128000,
refresh: 'auto',
cacheDir: '.ai-assist/tags-cache',
verbose: false,
exclude: [
'node_modules/**',
'dist/**',
'build/**',
'.git/**',
'*.test.*',
'*.spec.*',
'**/*.d.ts',
],
include: ['**/*.ts', '**/*.tsx', '**/*.js', '**/*.jsx', '**/*.py'],
};
/**
* RepoMap 类
* 生成代码仓库的上下文地图,帮助 AI 理解代码结构
*/
export class RepoMap {
private tagExtractor: TagExtractor;
private tagsCache: DiskCache<TagCacheEntry>;
private config: RepoMapConfig;
private root: string;
constructor(root: string, config: Partial<RepoMapConfig> = {}) {
this.root = root;
this.config = { ...defaultConfig, ...config };
this.tagExtractor = new TagExtractor();
this.tagsCache = new DiskCache(
path.join(this.root, this.config.cacheDir)
);
}
/**
* 获取 repo map
* @param chatFiles 当前对话中涉及的文件
* @param otherFiles 仓库中其他文件
* @param mentionedFnames 对话中提到的文件名
* @param mentionedIdents 对话中提到的标识符
*/
async getRepoMap(
chatFiles: string[],
otherFiles: string[],
mentionedFnames: Set<string> = new Set(),
mentionedIdents: Set<string> = new Set()
): Promise<string> {
if (this.config.mapTokens <= 0 || otherFiles.length === 0) {
return '';
}
let maxMapTokens = this.config.mapTokens;
// 无聊天文件时,给更大的视图
if (chatFiles.length === 0) {
maxMapTokens = Math.min(
maxMapTokens * this.config.mapMulNoFiles,
this.config.maxContextWindow - 4096
);
}
const rankedTags = await this.getRankedTags(
chatFiles,
otherFiles,
mentionedFnames,
mentionedIdents
);
// 二分搜索找到最优 token 数量
return this.fitToTokenLimit(rankedTags, maxMapTokens, new Set(chatFiles));
}
/**
* 获取排序后的 tags
*/
private async getRankedTags(
chatFnames: string[],
otherFnames: string[],
mentionedFnames: Set<string>,
mentionedIdents: Set<string>
): Promise<Tag[]> {
// ident -> files that define it
const defines = new Map<string, Set<string>>();
// ident -> files that reference it
const references = new Map<string, string[]>();
// (file:ident) -> tag objects
const definitions = new Map<string, Tag[]>();
// personalization vector for PageRank
const personalization = new Map<string, number>();
const allFnames = [...new Set([...chatFnames, ...otherFnames])];
const chatRelFnames = new Set(chatFnames.map((f) => this.getRelFname(f)));
const basePersonalize = 100 / Math.max(allFnames.length, 1);
// 收集所有文件的 tags
for (const fname of allFnames) {
const relFname = this.getRelFname(fname);
let currentPers = 0;
// 个性化权重
if (chatFnames.includes(fname)) {
currentPers += basePersonalize;
}
if (mentionedFnames.has(relFname)) {
currentPers = Math.max(currentPers, basePersonalize);
}
// 路径组件匹配
const pathParts = relFname.split('/');
const basename = pathParts[pathParts.length - 1];
const basenameNoExt = basename.replace(/\.[^.]+$/, '');
const components = new Set([...pathParts, basename, basenameNoExt]);
if ([...components].some((c) => mentionedIdents.has(c))) {
currentPers += basePersonalize;
}
if (currentPers > 0) {
personalization.set(relFname, currentPers);
}
// 提取 tags
const tags = await this.getTags(fname, relFname);
for (const tag of tags) {
if (tag.kind === 'def') {
if (!defines.has(tag.name)) defines.set(tag.name, new Set());
defines.get(tag.name)!.add(relFname);
const key = `${relFname}:${tag.name}`;
if (!definitions.has(key)) definitions.set(key, []);
definitions.get(key)!.push(tag);
} else {
if (!references.has(tag.name)) references.set(tag.name, []);
references.get(tag.name)!.push(relFname);
}
}
}
// 构建图
const graph = new Graph();
// 找到同时有定义和引用的标识符
const idents = [...defines.keys()].filter((id) => references.has(id));
for (const ident of idents) {
const definers = defines.get(ident)!;
const refs = references.get(ident)!;
// 计算权重乘数
let mul = 1.0;
if (mentionedIdents.has(ident)) mul *= 10;
if (this.isSignificantName(ident)) mul *= 10;
if (ident.startsWith('_')) mul *= 0.1;
if (definers.size > 5) mul *= 0.1;
// 统计引用次数
const refCounts = new Map<string, number>();
for (const ref of refs) {
refCounts.set(ref, (refCounts.get(ref) || 0) + 1);
}
// 添加边
for (const [referencer, numRefs] of refCounts) {
for (const definer of definers) {
let useMul = mul;
if (chatRelFnames.has(referencer)) useMul *= 50;
graph.addEdge({
from: referencer,
to: definer,
weight: useMul * Math.sqrt(numRefs),
ident,
});
}
}
}
// 运行 PageRank
const ranked = pagerank(graph, { personalization });
// 分配排名到定义
const rankedDefinitions = distributeRanksToDefinitions(graph, ranked);
// 排序
const sorted = [...rankedDefinitions.entries()].sort((a, b) => b[1] - a[1]);
// 收集排序后的 tags
const rankedTags: Tag[] = [];
for (const [key] of sorted) {
const colonIdx = key.lastIndexOf(':');
const fname = key.substring(0, colonIdx);
if (chatRelFnames.has(fname)) continue;
const tags = definitions.get(key) || [];
rankedTags.push(...tags);
}
return rankedTags;
}
/**
* 获取文件的 tags(带缓存)
*/
private async getTags(fname: string, relFname: string): Promise<Tag[]> {
const mtime = await this.getMtime(fname);
if (!mtime) return [];
// 检查缓存
const cached = await this.tagsCache.get(fname);
if (cached && cached.mtime === mtime) {
return cached.data;
}
// 提取 tags
const tags = await this.tagExtractor.getTags(fname, relFname);
// 更新缓存
await this.tagsCache.set(fname, { mtime, data: tags });
return tags;
}
/**
* 二分搜索拟合 token 限制
*/
private fitToTokenLimit(
tags: Tag[],
maxTokens: number,
chatRelFnames: Set<string>
): string {
const numTags = tags.length;
if (numTags === 0) return '';
let lower = 0;
let upper = numTags;
let bestTree = '';
let bestTokens = 0;
let middle = Math.min(Math.floor(maxTokens / 25), numTags);
while (lower <= upper) {
const tree = this.toTree(tags.slice(0, middle), chatRelFnames);
const numTokens = this.tokenCount(tree);
const pctErr = Math.abs(numTokens - maxTokens) / maxTokens;
if ((numTokens <= maxTokens && numTokens > bestTokens) || pctErr < 0.15) {
bestTree = tree;
bestTokens = numTokens;
if (pctErr < 0.15) break;
}
if (numTokens < maxTokens) {
lower = middle + 1;
} else {
upper = middle - 1;
}
middle = Math.floor((lower + upper) / 2);
}
return bestTree;
}
/**
* 转换为树形展示
*/
private toTree(tags: Tag[], chatRelFnames: Set<string>): string {
if (tags.length === 0) return '';
const output: string[] = [];
let curFname: string | null = null;
let lois: number[] = [];
let curNames: string[] = [];
// 按文件分组
const sortedTags = [...tags].sort(
(a, b) => a.relFname.localeCompare(b.relFname) || a.line - b.line
);
for (const tag of sortedTags) {
if (chatRelFnames.has(tag.relFname)) continue;
if (tag.relFname !== curFname) {
// 输出前一个文件
if (curFname && (lois.length > 0 || curNames.length > 0)) {
output.push(this.renderFileTree(curFname, lois, curNames));
}
curFname = tag.relFname;
lois = [];
curNames = [];
}
if (tag.line >= 0) {
lois.push(tag.line);
}
curNames.push(tag.name);
}
// 输出最后一个文件
if (curFname && (lois.length > 0 || curNames.length > 0)) {
output.push(this.renderFileTree(curFname, lois, curNames));
}
return output.join('\n');
}
/**
* 渲染单个文件的树形展示
*/
private renderFileTree(
relFname: string,
lois: number[],
names: string[]
): string {
const uniqueNames = [...new Set(names)];
const uniqueLines = [...new Set(lois)].sort((a, b) => a - b);
const lines = [`${relFname}:`];
// 简化版:显示符号列表
for (const name of uniqueNames.slice(0, 10)) {
lines.push(` - ${name}`);
}
if (uniqueNames.length > 10) {
lines.push(` ... and ${uniqueNames.length - 10} more`);
}
return lines.join('\n');
}
/**
* 判断是否是有意义的名称
*/
private isSignificantName(name: string): boolean {
if (name.length < 8) return false;
const hasAlpha = /[a-zA-Z]/.test(name);
const isSnake = name.includes('_') && hasAlpha;
const isKebab = name.includes('-') && hasAlpha;
const isCamel = /[a-z]/.test(name) && /[A-Z]/.test(name);
return isSnake || isKebab || isCamel;
}
/**
* 获取相对路径
*/
private getRelFname(fname: string): string {
if (fname.startsWith(this.root)) {
return fname.slice(this.root.length + 1);
}
return fname;
}
/**
* 获取文件修改时间
*/
private async getMtime(fname: string): Promise<number | null> {
try {
const stat = await fs.stat(fname);
return stat.mtimeMs;
} catch {
return null;
}
}
/**
* 估算 token 数量
* 简化估算:约 4 字符一个 token
*/
private tokenCount(text: string): number {
return Math.ceil(text.length / 4);
}
/**
* 刷新缓存到磁盘
*/
async flushCache(): Promise<void> {
await this.tagsCache.flushAll();
}
/**
* 清空缓存
*/
async clearCache(): Promise<void> {
await this.tagsCache.clear();
}
}
/**
* 创建 RepoMap 实例
*/
export function createRepoMap(
root: string,
config?: Partial<RepoMapConfig>
): RepoMap {
return new RepoMap(root, config);
}
+465
View File
@@ -0,0 +1,465 @@
/**
* Tag 提取器
* 使用 Tree-sitter 解析代码并提取符号定义和引用
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import { fileURLToPath } from 'url';
import type { Tag } from '../types.js';
import { getLanguageFromFilename } from '../types.js';
// Tree-sitter 类型(web-tree-sitter)
interface TreeSitterParser {
parse(input: string): TreeSitterTree;
setLanguage(language: TreeSitterLanguage): void;
getLanguage(): TreeSitterLanguage;
}
interface TreeSitterTree {
rootNode: TreeSitterNode;
}
interface TreeSitterNode {
text: string;
startPosition: { row: number; column: number };
endPosition: { row: number; column: number };
type: string;
childCount: number;
namedChildCount: number;
children: TreeSitterNode[];
namedChildren: TreeSitterNode[];
}
interface TreeSitterLanguage {
query(source: string): TreeSitterQuery;
}
interface TreeSitterQuery {
captures(node: TreeSitterNode): TreeSitterCapture[];
}
interface TreeSitterCapture {
node: TreeSitterNode;
name: string;
}
// 动态导入 web-tree-sitter
let ParserClass: any = null;
let treeSitterInitialized = false;
/**
* 初始化 Tree-sitter
*/
async function initTreeSitter(): Promise<void> {
if (treeSitterInitialized) return;
try {
const TreeSitter = await import('web-tree-sitter');
// Parser 是命名导出的类,init 是其静态方法
await TreeSitter.Parser.init();
ParserClass = TreeSitter.Parser;
treeSitterInitialized = true;
} catch (error) {
console.warn('Failed to initialize tree-sitter:', error);
throw new Error('Tree-sitter initialization failed');
}
}
/**
* Tag 提取器类
*/
export class TagExtractor {
private parsers: Map<string, TreeSitterParser> = new Map();
private languages: Map<string, TreeSitterLanguage> = new Map();
private queries: Map<string, TreeSitterQuery> = new Map();
private initialized = false;
private queriesDir: string;
constructor() {
// 获取查询文件目录
const __filename = fileURLToPath(import.meta.url);
const __dirname = path.dirname(__filename);
this.queriesDir = path.join(__dirname, 'queries');
}
/**
* 初始化提取器
*/
async initialize(): Promise<void> {
if (this.initialized) return;
try {
await initTreeSitter();
this.initialized = true;
} catch (error) {
// Tree-sitter 初始化失败,使用回退方案
console.warn('Tree-sitter not available, using regex fallback');
this.initialized = true;
}
}
/**
* 获取文件的所有 tags
*/
async getTags(fname: string, relFname: string): Promise<Tag[]> {
await this.initialize();
const lang = getLanguageFromFilename(fname);
if (!lang) return [];
let code: string;
try {
code = await fs.readFile(fname, 'utf-8');
} catch {
return [];
}
// 尝试使用 Tree-sitter
if (ParserClass) {
try {
const tags = await this.extractWithTreeSitter(code, fname, relFname, lang);
if (tags.length > 0) {
return tags;
}
} catch (error) {
// Tree-sitter 解析失败,回退到正则
}
}
// 回退到正则表达式
return this.extractWithRegex(code, fname, relFname, lang);
}
/**
* 使用 Tree-sitter 提取 tags
*/
private async extractWithTreeSitter(
code: string,
fname: string,
relFname: string,
lang: string
): Promise<Tag[]> {
const parser = await this.getParser(lang);
const query = await this.getQuery(lang);
if (!parser || !query) {
return [];
}
const tree = parser.parse(code);
const captures = query.captures(tree.rootNode);
const tags: Tag[] = [];
const seenKinds = new Set<string>();
for (const capture of captures) {
const { node, name } = capture;
let kind: 'def' | 'ref' | null = null;
if (name.startsWith('name.definition.')) {
kind = 'def';
} else if (name.startsWith('name.reference.')) {
kind = 'ref';
} else {
continue;
}
seenKinds.add(kind);
tags.push({
relFname,
fname,
name: node.text,
kind,
line: node.startPosition.row,
});
}
// 如果只有 def 没有 ref,使用正则回退提取引用
if (seenKinds.has('def') && !seenKinds.has('ref')) {
const refTags = this.extractRefsWithRegex(code, fname, relFname);
tags.push(...refTags);
}
return tags;
}
/**
* 获取或创建解析器
*/
private async getParser(lang: string): Promise<TreeSitterParser | null> {
if (this.parsers.has(lang)) {
return this.parsers.get(lang)!;
}
try {
const language = await this.getLanguage(lang);
if (!language) return null;
const parser = new ParserClass();
parser.setLanguage(language);
this.parsers.set(lang, parser);
return parser;
} catch {
return null;
}
}
/**
* 获取语言定义
*/
private async getLanguage(lang: string): Promise<TreeSitterLanguage | null> {
if (this.languages.has(lang)) {
return this.languages.get(lang)!;
}
try {
// 尝试加载 WASM 语言文件
const wasmPath = this.getWasmPath(lang);
const language = await ParserClass.Language.load(wasmPath);
this.languages.set(lang, language);
return language;
} catch {
return null;
}
}
/**
* 获取 WASM 文件路径
*/
private getWasmPath(lang: string): string {
// tree-sitter WASM 文件的标准位置
const wasmName = `tree-sitter-${lang}.wasm`;
// 尝试多个可能的位置
const possiblePaths = [
path.join(process.cwd(), 'node_modules', 'tree-sitter-wasms', 'out', wasmName),
path.join(process.cwd(), 'node_modules', `tree-sitter-${lang}`, wasmName),
path.join(__dirname, '..', '..', '..', 'wasm', wasmName),
];
// 返回第一个路径(实际使用时需要检查存在性)
return possiblePaths[0];
}
/**
* 获取语言查询
*/
private async getQuery(lang: string): Promise<TreeSitterQuery | null> {
if (this.queries.has(lang)) {
return this.queries.get(lang)!;
}
try {
const language = await this.getLanguage(lang);
if (!language) return null;
const queryPath = path.join(this.queriesDir, `${lang}-tags.scm`);
const queryText = await fs.readFile(queryPath, 'utf-8');
const query = language.query(queryText);
this.queries.set(lang, query);
return query;
} catch {
return null;
}
}
/**
* 使用正则表达式提取 tags (回退方案)
*/
private extractWithRegex(
code: string,
fname: string,
relFname: string,
lang: string
): Tag[] {
const tags: Tag[] = [];
const lines = code.split('\n');
// 根据语言选择正则模式
const patterns = this.getRegexPatterns(lang);
lines.forEach((line, lineNum) => {
for (const pattern of patterns.definitions) {
const match = line.match(pattern.regex);
if (match && match[pattern.nameGroup]) {
tags.push({
relFname,
fname,
name: match[pattern.nameGroup],
kind: 'def',
line: lineNum,
});
}
}
});
// 提取引用
tags.push(...this.extractRefsWithRegex(code, fname, relFname));
return tags;
}
/**
* 使用正则提取引用
*/
private extractRefsWithRegex(code: string, fname: string, relFname: string): Tag[] {
const tags: Tag[] = [];
// 简单的标识符匹配 - 排除关键字
const keywords = new Set([
'if',
'else',
'for',
'while',
'do',
'switch',
'case',
'break',
'continue',
'return',
'function',
'class',
'const',
'let',
'var',
'import',
'export',
'from',
'default',
'async',
'await',
'try',
'catch',
'finally',
'throw',
'new',
'this',
'super',
'extends',
'implements',
'interface',
'type',
'enum',
'public',
'private',
'protected',
'static',
'readonly',
'abstract',
'true',
'false',
'null',
'undefined',
'void',
'never',
'any',
'unknown',
'string',
'number',
'boolean',
'object',
'symbol',
'bigint',
'def',
'class',
'self',
'None',
'True',
'False',
'and',
'or',
'not',
'in',
'is',
'lambda',
'with',
'as',
'pass',
'raise',
'yield',
'global',
'nonlocal',
'assert',
'del',
]);
// 匹配 PascalCase 或 snake_case 标识符(更可能是用户定义的)
const identRegex = /\b([A-Z][a-zA-Z0-9]*|[a-z][a-zA-Z0-9]*_[a-zA-Z0-9_]*)\b/g;
let match;
while ((match = identRegex.exec(code)) !== null) {
const name = match[1];
if (!keywords.has(name) && name.length >= 2) {
tags.push({
relFname,
fname,
name,
kind: 'ref',
line: -1, // 行号未知
});
}
}
return tags;
}
/**
* 获取语言的正则模式
*/
private getRegexPatterns(lang: string): {
definitions: Array<{ regex: RegExp; nameGroup: number }>;
} {
switch (lang) {
case 'typescript':
case 'javascript':
return {
definitions: [
// function name(
{ regex: /(?:export\s+)?(?:async\s+)?function\s+(\w+)\s*[<(]/, nameGroup: 1 },
// class Name
{ regex: /(?:export\s+)?(?:abstract\s+)?class\s+(\w+)/, nameGroup: 1 },
// interface Name
{ regex: /(?:export\s+)?interface\s+(\w+)/, nameGroup: 1 },
// type Name =
{ regex: /(?:export\s+)?type\s+(\w+)\s*[<=]/, nameGroup: 1 },
// const/let/var name = (arrow function or function)
{
regex: /(?:export\s+)?(?:const|let|var)\s+(\w+)\s*=\s*(?:async\s+)?(?:\([^)]*\)|[^=])\s*=>/,
nameGroup: 1,
},
// enum Name
{ regex: /(?:export\s+)?enum\s+(\w+)/, nameGroup: 1 },
// method name(
{ regex: /^\s*(?:async\s+)?(\w+)\s*\([^)]*\)\s*[:{]/, nameGroup: 1 },
],
};
case 'python':
return {
definitions: [
// def name(
{ regex: /^\s*(?:async\s+)?def\s+(\w+)\s*\(/, nameGroup: 1 },
// class Name
{ regex: /^\s*class\s+(\w+)/, nameGroup: 1 },
],
};
default:
return {
definitions: [
// 通用:function/def/class 后面的标识符
{ regex: /(?:function|def|class)\s+(\w+)/, nameGroup: 1 },
],
};
}
}
}
/**
* 创建 Tag 提取器实例
*/
export function createTagExtractor(): TagExtractor {
return new TagExtractor();
}
+5
View File
@@ -0,0 +1,5 @@
/**
* Tags 模块导出
*/
export { TagExtractor, createTagExtractor } from './extractor.js';
@@ -0,0 +1,65 @@
;; JavaScript/JSX 标签查询
;; 定义使用 @name.definition.* 前缀
;; 引用使用 @name.reference.* 前缀
;; ==================== 定义 ====================
;; 函数声明
(function_declaration
name: (identifier) @name.definition.function) @definition.function
;; 生成器函数
(generator_function_declaration
name: (identifier) @name.definition.function) @definition.function
;; 箭头函数赋值 (const/let)
(lexical_declaration
(variable_declarator
name: (identifier) @name.definition.function
value: [(arrow_function) (function)])) @definition.function
;; 箭头函数赋值 (var)
(variable_declaration
(variable_declarator
name: (identifier) @name.definition.function
value: [(arrow_function) (function)])) @definition.function
;; 赋值表达式中的函数
(assignment_expression
left: (identifier) @name.definition.function
right: [(arrow_function) (function)]) @definition.function
(assignment_expression
left: (member_expression
property: (property_identifier) @name.definition.function)
right: [(arrow_function) (function)]) @definition.function
;; 对象方法
(pair
key: (property_identifier) @name.definition.function
value: [(arrow_function) (function)]) @definition.function
;; 方法定义
(method_definition
name: (property_identifier) @name.definition.method) @definition.method
;; 类定义
(class_declaration
name: (identifier) @name.definition.class) @definition.class
(class
name: (identifier) @name.definition.class) @definition.class
;; ==================== 引用 ====================
;; 函数调用
(call_expression
function: (identifier) @name.reference.call) @reference.call
(call_expression
function: (member_expression
property: (property_identifier) @name.reference.call)) @reference.call
;; new 表达式
(new_expression
constructor: (identifier) @name.reference.class) @reference.class
@@ -0,0 +1,37 @@
;; Python 标签查询
;; 定义使用 @name.definition.* 前缀
;; 引用使用 @name.reference.* 前缀
;; ==================== 定义 ====================
;; 函数定义
(function_definition
name: (identifier) @name.definition.function) @definition.function
;; 类定义
(class_definition
name: (identifier) @name.definition.class) @definition.class
;; 装饰器函数 (通常也是定义)
(decorated_definition
(function_definition
name: (identifier) @name.definition.function)) @definition.function
(decorated_definition
(class_definition
name: (identifier) @name.definition.class)) @definition.class
;; ==================== 引用 ====================
;; 函数调用
(call
function: (identifier) @name.reference.call) @reference.call
(call
function: (attribute
attribute: (identifier) @name.reference.call)) @reference.call
;; 类继承
(class_definition
superclasses: (argument_list
(identifier) @name.reference.class)) @reference.class
@@ -0,0 +1,90 @@
;; TypeScript/TSX 标签查询
;; 定义使用 @name.definition.* 前缀
;; 引用使用 @name.reference.* 前缀
;; ==================== 定义 ====================
;; 函数定义
(function_declaration
name: (identifier) @name.definition.function) @definition.function
(function_signature
name: (identifier) @name.definition.function) @definition.function
;; 箭头函数赋值
(lexical_declaration
(variable_declarator
name: (identifier) @name.definition.function
value: (arrow_function))) @definition.function
(variable_declaration
(variable_declarator
name: (identifier) @name.definition.function
value: (arrow_function))) @definition.function
;; 方法定义
(method_definition
name: (property_identifier) @name.definition.method) @definition.method
(method_signature
name: (property_identifier) @name.definition.method) @definition.method
(abstract_method_signature
name: (property_identifier) @name.definition.method) @definition.method
;; 类定义
(class_declaration
name: (type_identifier) @name.definition.class) @definition.class
(abstract_class_declaration
name: (type_identifier) @name.definition.class) @definition.class
;; 接口定义
(interface_declaration
name: (type_identifier) @name.definition.interface) @definition.interface
;; 类型别名
(type_alias_declaration
name: (type_identifier) @name.definition.type) @definition.type
;; 枚举
(enum_declaration
name: (identifier) @name.definition.enum) @definition.enum
;; 模块/命名空间
(module
name: (identifier) @name.definition.module) @definition.module
(internal_module
name: (identifier) @name.definition.module) @definition.module
;; ==================== 引用 ====================
;; 类型注解引用
(type_annotation
(type_identifier) @name.reference.type) @reference.type
;; 类型参数中的引用
(type_arguments
(type_identifier) @name.reference.type) @reference.type
;; extends/implements
(class_heritage
(extends_clause
value: (identifier) @name.reference.class)) @reference.class
(class_heritage
(implements_clause
(type_identifier) @name.reference.interface)) @reference.interface
;; new 表达式
(new_expression
constructor: (identifier) @name.reference.class) @reference.class
;; 函数调用
(call_expression
function: (identifier) @name.reference.call) @reference.call
(call_expression
function: (member_expression
property: (property_identifier) @name.reference.call)) @reference.call
+142
View File
@@ -0,0 +1,142 @@
/**
* RepoMap 类型定义
* 基于 Aider 的 RepoMap 实现
*/
/**
* 代码标签 - 对应 Aider 的 Tag
*/
export interface Tag {
/** 相对路径 */
relFname: string;
/** 绝对路径 */
fname: string;
/** 行号 (0-indexed, -1 表示未知) */
line: number;
/** 符号名称 */
name: string;
/** 定义或引用 */
kind: 'def' | 'ref';
}
/**
* 文件缓存条目
*/
export interface TagCacheEntry {
/** 文件修改时间 (ms) */
mtime: number;
/** 标签数据 */
data: Tag[];
}
/**
* RepoMap 配置
*/
export interface RepoMapConfig {
/** 目标 token 数量 */
mapTokens: number;
/** 无聊天文件时的乘数 */
mapMulNoFiles: number;
/** 最大上下文窗口 */
maxContextWindow: number;
/** 刷新策略 */
refresh: 'auto' | 'always' | 'files' | 'manual';
/** 缓存目录 */
cacheDir: string;
/** 是否详细输出 */
verbose: boolean;
/** 排除的文件模式 */
exclude: string[];
/** 包含的文件模式 */
include: string[];
}
/**
* 默认配置
*/
export const DEFAULT_REPOMAP_CONFIG: RepoMapConfig = {
mapTokens: 2048,
mapMulNoFiles: 8,
maxContextWindow: 128000,
refresh: 'auto',
cacheDir: '.ai-assist/tags-cache',
verbose: false,
exclude: [
'node_modules/**',
'dist/**',
'build/**',
'.git/**',
'*.test.*',
'*.spec.*',
'**/*.d.ts',
],
include: ['**/*.ts', '**/*.tsx', '**/*.js', '**/*.jsx', '**/*.py'],
};
/**
* 图边
*/
export interface GraphEdge {
/** 引用文件 */
from: string;
/** 定义文件 */
to: string;
/** 边权重 */
weight: number;
/** 符号名 */
ident: string;
}
/**
* 排序后的定义
*/
export interface RankedDefinition {
fname: string;
ident: string;
rank: number;
tags: Tag[];
}
/**
* PageRank 算法配置
*/
export interface PageRankOptions {
/** 阻尼系数 (默认 0.85) */
damping?: number;
/** 最大迭代次数 (默认 100) */
maxIterations?: number;
/** 收敛容差 (默认 1e-6) */
tolerance?: number;
/** 个性化向量 */
personalization?: Map<string, number>;
}
/**
* 支持的语言映射
*/
export const LANGUAGE_MAP: Record<string, string> = {
'.ts': 'typescript',
'.tsx': 'typescript',
'.js': 'javascript',
'.jsx': 'javascript',
'.mjs': 'javascript',
'.cjs': 'javascript',
'.py': 'python',
'.go': 'go',
'.rs': 'rust',
'.java': 'java',
'.rb': 'ruby',
'.cpp': 'cpp',
'.cc': 'cpp',
'.c': 'c',
'.h': 'c',
'.hpp': 'cpp',
};
/**
* 获取文件语言
*/
export function getLanguageFromFilename(filename: string): string | null {
const ext = filename.substring(filename.lastIndexOf('.')).toLowerCase();
return LANGUAGE_MAP[ext] || null;
}
+10
View File
@@ -0,0 +1,10 @@
export type {
SessionData,
SessionSummary,
SessionManagerConfig,
Todo,
TodoStatus,
} from './types.js';
export { SessionStorage, sessionStorage } from './storage.js';
export { SessionManager, sessionManager } from './manager.js';
+255
View File
@@ -0,0 +1,255 @@
import type { ModelMessage } from 'ai';
import type { SessionData, Todo, SessionSummary } from './types.js';
import { SessionStorage, sessionStorage } from './storage.js';
/**
* 会话管理器
* 提供高级会话操作接口
*/
export class SessionManager {
private storage: SessionStorage;
private currentSession: SessionData | null = null;
private autoSaveInterval: ReturnType<typeof setInterval> | null = null;
constructor(storage?: SessionStorage) {
this.storage = storage || sessionStorage;
}
/**
* 初始化 - 尝试恢复或创建新会话
*/
async init(workdir: string): Promise<SessionData> {
// 尝试加载当前会话
const existing = await this.storage.loadCurrentSession();
if (existing && existing.workdir === workdir) {
// 同一工作目录,恢复会话
this.currentSession = existing;
} else {
// 不同目录或无会话,归档旧会话并创建新的
if (existing) {
await this.storage.archiveCurrentSession();
}
this.currentSession = this.createNewSession(workdir);
await this.save();
}
// 启动自动保存
this.startAutoSave();
return this.currentSession;
}
/**
* 创建新会话
*/
private createNewSession(workdir: string): SessionData {
return {
id: this.storage.generateSessionId(),
createdAt: new Date().toISOString(),
updatedAt: new Date().toISOString(),
workdir,
messages: [],
discoveredTools: [],
todos: [],
};
}
/**
* 获取当前会话
*/
getSession(): SessionData | null {
return this.currentSession;
}
/**
* 保存当前会话
*/
async save(): Promise<void> {
if (this.currentSession) {
await this.storage.saveCurrentSession(this.currentSession);
}
}
/**
* 添加消息
*/
async addMessage(message: ModelMessage): Promise<void> {
if (!this.currentSession) return;
this.currentSession.messages.push(message);
await this.save();
}
/**
* 批量设置消息(用于同步整个对话历史)
*/
async setMessages(messages: ModelMessage[]): Promise<void> {
if (!this.currentSession) return;
this.currentSession.messages = messages;
await this.save();
}
/**
* 获取对话历史
*/
getMessages(): ModelMessage[] {
return this.currentSession?.messages || [];
}
/**
* 设置已发现的工具
*/
async setDiscoveredTools(tools: string[]): Promise<void> {
if (!this.currentSession) return;
this.currentSession.discoveredTools = tools;
await this.save();
}
/**
* 获取已发现的工具
*/
getDiscoveredTools(): string[] {
return this.currentSession?.discoveredTools || [];
}
/**
* 更新待办事项
*/
async setTodos(todos: Todo[]): Promise<void> {
if (!this.currentSession) return;
this.currentSession.todos = todos;
await this.save();
}
/**
* 获取待办事项
*/
getTodos(): Todo[] {
return this.currentSession?.todos || [];
}
/**
* 清空当前会话并创建新会话
*/
async newSession(workdir?: string): Promise<SessionData> {
// 归档当前会话
if (this.currentSession && this.currentSession.messages.length > 0) {
await this.storage.archiveCurrentSession();
}
// 创建新会话
const newWorkdir = workdir || this.currentSession?.workdir || process.cwd();
this.currentSession = this.createNewSession(newWorkdir);
await this.save();
return this.currentSession;
}
/**
* 创建子会话(用于 Task 工具)
* @param parentId 父会话 ID
* @param agentName 关联的 Agent 名称
* @param title 会话标题
*/
createChildSession(parentId: string, agentName: string, title?: string): SessionData {
const workdir = this.currentSession?.workdir || process.cwd();
const childSession: SessionData = {
id: this.storage.generateSessionId(),
parentId,
agentName,
createdAt: new Date().toISOString(),
updatedAt: new Date().toISOString(),
workdir,
title: title || `子任务 (@${agentName})`,
messages: [],
discoveredTools: [],
todos: [],
};
return childSession;
}
/**
* 保存子会话
*/
async saveChildSession(session: SessionData): Promise<void> {
await this.storage.saveSession(session);
}
/**
* 获取当前会话 ID
*/
getSessionId(): string | undefined {
return this.currentSession?.id;
}
/**
* 恢复指定会话
*/
async restoreSession(sessionId: string): Promise<SessionData | null> {
const session = await this.storage.loadSession(sessionId);
if (!session) return null;
// 归档当前会话
if (this.currentSession && this.currentSession.messages.length > 0) {
await this.storage.archiveCurrentSession();
}
this.currentSession = session;
await this.save();
return session;
}
/**
* 列出历史会话
*/
async listSessions(): Promise<SessionSummary[]> {
return this.storage.listSessions();
}
/**
* 删除历史会话
*/
async deleteSession(sessionId: string): Promise<boolean> {
return this.storage.deleteSession(sessionId);
}
/**
* 启动自动保存(每 30 秒)
*/
private startAutoSave(): void {
if (this.autoSaveInterval) return;
this.autoSaveInterval = setInterval(async () => {
await this.save();
}, 30000);
}
/**
* 停止自动保存
*/
stopAutoSave(): void {
if (this.autoSaveInterval) {
clearInterval(this.autoSaveInterval);
this.autoSaveInterval = null;
}
}
/**
* 关闭管理器(保存并停止自动保存)
*/
async close(): Promise<void> {
this.stopAutoSave();
await this.save();
}
/**
* 清理旧会话
*/
async cleanup(keepCount?: number): Promise<number> {
return this.storage.cleanupOldSessions(keepCount);
}
}
// 导出默认实例
export const sessionManager = new SessionManager();
+215
View File
@@ -0,0 +1,215 @@
import * as fs from 'fs/promises';
import * as path from 'path';
import * as os from 'os';
import type { SessionData, SessionSummary } from './types.js';
/**
* 获取默认存储目录
* 遵循 XDG 规范:~/.local/share/ai-assist/
*/
function getDefaultStorageDir(): string {
const xdgDataHome = process.env.XDG_DATA_HOME;
if (xdgDataHome) {
return path.join(xdgDataHome, 'ai-assist');
}
return path.join(os.homedir(), '.local', 'share', 'ai-assist');
}
/**
* 会话存储类
* 负责会话数据的读写操作
*/
export class SessionStorage {
private storageDir: string;
private sessionsDir: string;
private currentSessionFile: string;
constructor(storageDir?: string) {
this.storageDir = storageDir || getDefaultStorageDir();
this.sessionsDir = path.join(this.storageDir, 'sessions');
this.currentSessionFile = path.join(this.storageDir, 'current-session.json');
}
/**
* 确保存储目录存在
*/
async ensureDir(): Promise<void> {
await fs.mkdir(this.sessionsDir, { recursive: true });
}
/**
* 生成会话 ID
*/
generateSessionId(): string {
const now = new Date();
const timestamp = now.toISOString().slice(0, 10); // YYYY-MM-DD
const random = Math.random().toString(36).substring(2, 8);
return `${timestamp}_${random}`;
}
/**
* 保存当前会话
*/
async saveCurrentSession(session: SessionData): Promise<void> {
await this.ensureDir();
session.updatedAt = new Date().toISOString();
await fs.writeFile(
this.currentSessionFile,
JSON.stringify(session, null, 2),
'utf-8'
);
}
/**
* 加载当前会话
*/
async loadCurrentSession(): Promise<SessionData | null> {
try {
const content = await fs.readFile(this.currentSessionFile, 'utf-8');
return JSON.parse(content) as SessionData;
} catch {
return null;
}
}
/**
* 归档当前会话到历史
*/
async archiveCurrentSession(): Promise<void> {
const current = await this.loadCurrentSession();
if (!current || current.messages.length === 0) {
return;
}
await this.ensureDir();
const archivePath = path.join(this.sessionsDir, `${current.id}.json`);
await fs.writeFile(archivePath, JSON.stringify(current, null, 2), 'utf-8');
}
/**
* 删除当前会话文件
*/
async clearCurrentSession(): Promise<void> {
try {
await fs.unlink(this.currentSessionFile);
} catch {
// 文件不存在,忽略
}
}
/**
* 列出历史会话
*/
async listSessions(): Promise<SessionSummary[]> {
await this.ensureDir();
const files = await fs.readdir(this.sessionsDir);
const summaries: SessionSummary[] = [];
for (const file of files) {
if (!file.endsWith('.json')) continue;
try {
const filePath = path.join(this.sessionsDir, file);
const content = await fs.readFile(filePath, 'utf-8');
const session = JSON.parse(content) as SessionData;
summaries.push({
id: session.id,
title: session.title || this.generateTitle(session),
workdir: session.workdir,
messageCount: session.messages.length,
createdAt: session.createdAt,
updatedAt: session.updatedAt,
});
} catch {
// 跳过无法解析的文件
}
}
// 按更新时间降序排列
return summaries.sort(
(a, b) => new Date(b.updatedAt).getTime() - new Date(a.updatedAt).getTime()
);
}
/**
* 加载指定会话
*/
async loadSession(sessionId: string): Promise<SessionData | null> {
try {
const filePath = path.join(this.sessionsDir, `${sessionId}.json`);
const content = await fs.readFile(filePath, 'utf-8');
return JSON.parse(content) as SessionData;
} catch {
return null;
}
}
/**
* 保存指定会话(用于子会话)
*/
async saveSession(session: SessionData): Promise<void> {
await this.ensureDir();
session.updatedAt = new Date().toISOString();
const filePath = path.join(this.sessionsDir, `${session.id}.json`);
await fs.writeFile(filePath, JSON.stringify(session, null, 2), 'utf-8');
}
/**
* 删除指定会话
*/
async deleteSession(sessionId: string): Promise<boolean> {
try {
const filePath = path.join(this.sessionsDir, `${sessionId}.json`);
await fs.unlink(filePath);
return true;
} catch {
return false;
}
}
/**
* 清理旧会话(保留最近 N 个)
*/
async cleanupOldSessions(keepCount: number = 50): Promise<number> {
const sessions = await this.listSessions();
if (sessions.length <= keepCount) {
return 0;
}
const toDelete = sessions.slice(keepCount);
let deletedCount = 0;
for (const session of toDelete) {
if (await this.deleteSession(session.id)) {
deletedCount++;
}
}
return deletedCount;
}
/**
* 从会话生成标题
*/
private generateTitle(session: SessionData): string {
// 从第一条用户消息生成标题
const firstUserMessage = session.messages.find((m) => m.role === 'user');
if (firstUserMessage && typeof firstUserMessage.content === 'string') {
const content = firstUserMessage.content;
// 取前 50 个字符
return content.length > 50 ? content.substring(0, 50) + '...' : content;
}
return `会话 ${session.id}`;
}
/**
* 获取存储目录路径
*/
getStorageDir(): string {
return this.storageDir;
}
}
// 导出默认实例
export const sessionStorage = new SessionStorage();
+65
View File
@@ -0,0 +1,65 @@
import type { ModelMessage } from 'ai';
/**
* 待办项状态
*/
export type TodoStatus = 'pending' | 'in_progress' | 'completed';
/**
* 待办项
*/
export interface Todo {
id: string;
content: string;
status: TodoStatus;
createdAt: string;
updatedAt: string;
}
/**
* 会话数据(持久化存储格式)
*/
export interface SessionData {
/** 会话 ID */
id: string;
/** 父会话 ID(子会话时存在) */
parentId?: string;
/** 关联的 Agent 名称(子会话时存在) */
agentName?: string;
/** 创建时间 */
createdAt: string;
/** 最后更新时间 */
updatedAt: string;
/** 工作目录 */
workdir: string;
/** 会话标题(可选,从第一条消息生成) */
title?: string;
/** 对话历史 */
messages: ModelMessage[];
/** 已发现的工具 */
discoveredTools: string[];
/** 待办事项 */
todos: Todo[];
}
/**
* 会话摘要(用于列表展示)
*/
export interface SessionSummary {
id: string;
title: string;
workdir: string;
messageCount: number;
createdAt: string;
updatedAt: string;
}
/**
* 会话管理器配置
*/
export interface SessionManagerConfig {
/** 存储目录 */
storageDir: string;
/** 最大历史会话数量 */
maxHistorySessions?: number;
}
+357
View File
@@ -0,0 +1,357 @@
/**
* 内置 Skills
*
* 提供一些常用的预定义 Skills
*/
import type { Skill } from '../types.js';
/**
* 代码审查 Skill
*/
export const codeReviewSkill: Skill = {
name: 'code-review',
displayName: '代码审查',
description: '对代码进行全面审查,检查潜在问题、最佳实践和改进建议',
category: 'development',
promptTemplate: `请对以下代码进行全面审查:
{{code}}
审查重点:
{{#if focus}}
- {{focus}}
{{else}}
- 代码质量和可读性
- 潜在的 bug 和错误
- 性能问题
- 安全隐患
- 最佳实践遵循情况
{{/if}}
请提供:
1. 发现的问题列表(按严重程度排序)
2. 具体的改进建议
3. 代码中做得好的地方`,
parameters: {
code: {
type: 'string',
description: '要审查的代码',
required: true,
},
focus: {
type: 'string',
description: '审查重点(可选)',
required: false,
},
},
keywords: ['review', 'code', 'quality', '审查', '代码', '质量'],
source: 'builtin',
enabled: true,
};
/**
* 代码解释 Skill
*/
export const explainCodeSkill: Skill = {
name: 'explain-code',
displayName: '代码解释',
description: '详细解释代码的功能、逻辑和工作原理',
category: 'development',
promptTemplate: `请详细解释以下代码:
{{code}}
{{#if level}}
解释级别:{{level}}
{{/if}}
请包含:
1. 代码的整体功能
2. 主要逻辑流程
3. 关键部分的详细解释
4. 使用的设计模式或技术(如果有)`,
parameters: {
code: {
type: 'string',
description: '要解释的代码',
required: true,
},
level: {
type: 'string',
description: '解释级别(beginner/intermediate/advanced)',
required: false,
enum: ['beginner', 'intermediate', 'advanced'],
default: 'intermediate',
},
},
keywords: ['explain', 'code', 'understand', '解释', '理解', '说明'],
source: 'builtin',
enabled: true,
};
/**
* 文档生成 Skill
*/
export const generateDocsSkill: Skill = {
name: 'generate-docs',
displayName: '文档生成',
description: '为代码生成文档注释或 README',
category: 'documentation',
promptTemplate: `请为以下代码生成{{type}}:
{{code}}
{{#if style}}
文档风格:{{style}}
{{/if}}
要求:
- 清晰描述功能和用途
- 说明参数和返回值(如果适用)
- 提供使用示例
- 使用规范的格式`,
parameters: {
code: {
type: 'string',
description: '要生成文档的代码',
required: true,
},
type: {
type: 'string',
description: '文档类型',
required: false,
enum: ['JSDoc', 'TSDoc', 'README', '注释', 'API文档'],
default: '文档注释',
},
style: {
type: 'string',
description: '文档风格(简洁/详细)',
required: false,
},
},
keywords: ['docs', 'documentation', 'jsdoc', '文档', '注释', '说明'],
source: 'builtin',
enabled: true,
};
/**
* 单元测试生成 Skill
*/
export const generateTestsSkill: Skill = {
name: 'generate-tests',
displayName: '测试生成',
description: '为代码生成单元测试',
category: 'testing',
promptTemplate: `请为以下代码生成单元测试:
{{code}}
测试框架:{{framework}}
要求:
- 覆盖主要功能路径
- 包含边界条件测试
- 包含错误处理测试
- 使用清晰的测试描述
- 遵循 AAA 模式(Arrange-Act-Assert)`,
parameters: {
code: {
type: 'string',
description: '要测试的代码',
required: true,
},
framework: {
type: 'string',
description: '测试框架',
required: false,
enum: ['vitest', 'jest', 'mocha', 'pytest', 'unittest'],
default: 'vitest',
},
},
keywords: ['test', 'unit', 'testing', '测试', '单元测试', 'vitest', 'jest'],
source: 'builtin',
enabled: true,
};
/**
* 重构建议 Skill
*/
export const refactorSuggestSkill: Skill = {
name: 'refactor-suggest',
displayName: '重构建议',
description: '分析代码并提供重构建议',
category: 'development',
promptTemplate: `请分析以下代码并提供重构建议:
{{code}}
{{#if goal}}
重构目标:{{goal}}
{{/if}}
请提供:
1. 当前代码的问题分析
2. 具体的重构建议
3. 重构后的代码示例
4. 重构的好处说明`,
parameters: {
code: {
type: 'string',
description: '要重构的代码',
required: true,
},
goal: {
type: 'string',
description: '重构目标(如:提高可读性、优化性能、减少重复)',
required: false,
},
},
keywords: ['refactor', 'improve', 'optimize', '重构', '优化', '改进'],
source: 'builtin',
enabled: true,
};
/**
* Bug 修复 Skill
*/
export const fixBugSkill: Skill = {
name: 'fix-bug',
displayName: 'Bug 修复',
description: '分析代码问题并提供修复方案',
category: 'debugging',
promptTemplate: `请分析以下代码中的问题并提供修复方案:
代码:
{{code}}
{{#if error}}
错误信息:
{{error}}
{{/if}}
{{#if context}}
上下文:
{{context}}
{{/if}}
请提供:
1. 问题的根本原因分析
2. 修复方案
3. 修复后的代码
4. 如何避免类似问题的建议`,
parameters: {
code: {
type: 'string',
description: '有问题的代码',
required: true,
},
error: {
type: 'string',
description: '错误信息',
required: false,
},
context: {
type: 'string',
description: '额外的上下文信息',
required: false,
},
},
keywords: ['bug', 'fix', 'debug', 'error', '修复', '错误', '调试'],
source: 'builtin',
enabled: true,
};
/**
* Git Commit 消息生成 Skill
*/
export const gitCommitSkill: Skill = {
name: 'git-commit',
displayName: 'Git Commit',
description: '根据代码变更生成规范的 Git commit 消息',
category: 'git',
promptTemplate: `请根据以下代码变更生成规范的 Git commit 消息:
变更内容:
{{diff}}
{{#if type}}
Commit 类型:{{type}}
{{/if}}
要求:
- 遵循 Conventional Commits 规范
- 第一行不超过 50 个字符
- 清晰描述变更的目的
- 如果有 breaking changes,请说明`,
parameters: {
diff: {
type: 'string',
description: 'Git diff 内容或变更描述',
required: true,
},
type: {
type: 'string',
description: 'Commit 类型',
required: false,
enum: ['feat', 'fix', 'docs', 'style', 'refactor', 'test', 'chore'],
},
},
keywords: ['git', 'commit', 'message', '提交', '消息'],
source: 'builtin',
enabled: true,
};
/**
* API 设计 Skill
*/
export const apiDesignSkill: Skill = {
name: 'api-design',
displayName: 'API 设计',
description: '设计 RESTful API 接口',
category: 'architecture',
promptTemplate: `请为以下需求设计 RESTful API:
需求描述:
{{requirement}}
{{#if constraints}}
约束条件:
{{constraints}}
{{/if}}
请提供:
1. API 端点设计(路径、方法、参数)
2. 请求/响应格式(JSON 示例)
3. 错误处理方案
4. 认证/授权建议(如果需要)`,
parameters: {
requirement: {
type: 'string',
description: 'API 需求描述',
required: true,
},
constraints: {
type: 'string',
description: '设计约束(如:现有系统兼容性、性能要求等)',
required: false,
},
},
keywords: ['api', 'rest', 'design', 'endpoint', 'API', '接口', '设计'],
source: 'builtin',
enabled: true,
};
/**
* 所有内置 Skills
*/
export const builtinSkills: Skill[] = [
codeReviewSkill,
explainCodeSkill,
generateDocsSkill,
generateTestsSkill,
refactorSuggestSkill,
fixBugSkill,
gitCommitSkill,
apiDesignSkill,
];
+29
View File
@@ -0,0 +1,29 @@
/**
* Skills 模块
*
* 提供 Skill 系统的所有功能导出
*/
// 类型
export type {
Skill,
SkillParameter,
SkillContext,
SkillExecutionResult,
SkillFile,
SkillSearchResult,
SkillRegistryConfig,
} from './types.js';
// 加载器
export { SkillLoader, skillLoader } from './loader.js';
// 注册表
export {
SkillRegistry,
getSkillRegistry,
resetSkillRegistry,
} from './registry.js';
// 内置 Skills
export { builtinSkills } from './builtin/index.js';
+201
View File
@@ -0,0 +1,201 @@
/**
* Skill 加载器
*
* 负责从文件系统加载 Skill 定义。
* 支持从以下位置加载:
* 1. 内置 Skills(代码中定义)
* 2. 用户 Skills(~/.config/ai-terminal/skills/)
* 3. 项目 Skills(./.ai-terminal/skills/)
*/
import * as fs from 'fs/promises';
import * as path from 'path';
import * as yaml from 'yaml';
import type { Skill, SkillFile } from './types.js';
/**
* Skill 加载器
*/
export class SkillLoader {
/**
* 从目录加载所有 Skills
*/
async loadFromDirectory(
dir: string,
source: 'user' | 'project'
): Promise<Skill[]> {
const skills: Skill[] = [];
try {
const exists = await fs
.access(dir)
.then(() => true)
.catch(() => false);
if (!exists) {
return skills;
}
const entries = await fs.readdir(dir, { withFileTypes: true });
for (const entry of entries) {
if (entry.isFile()) {
const ext = path.extname(entry.name).toLowerCase();
if (['.yaml', '.yml', '.json', '.md'].includes(ext)) {
const filePath = path.join(dir, entry.name);
try {
const skill = await this.loadFromFile(filePath, source);
if (skill) {
skills.push(skill);
}
} catch (error) {
console.warn(`加载 Skill 文件失败: ${filePath}`, error);
}
}
} else if (entry.isDirectory()) {
// 递归加载子目录
const subDir = path.join(dir, entry.name);
const subSkills = await this.loadFromDirectory(subDir, source);
skills.push(...subSkills);
}
}
} catch (error) {
console.warn(`读取 Skills 目录失败: ${dir}`, error);
}
return skills;
}
/**
* 从单个文件加载 Skill
*/
async loadFromFile(
filePath: string,
source: 'user' | 'project'
): Promise<Skill | null> {
const ext = path.extname(filePath).toLowerCase();
const content = await fs.readFile(filePath, 'utf-8');
let skillData: SkillFile | null = null;
if (ext === '.md') {
// Markdown 格式:从 frontmatter 和内容中解析
skillData = this.parseMarkdownSkill(content, filePath);
} else if (ext === '.yaml' || ext === '.yml') {
// YAML 格式
skillData = yaml.parse(content) as SkillFile;
} else if (ext === '.json') {
// JSON 格式
skillData = JSON.parse(content) as SkillFile;
}
if (!skillData?.skill) {
return null;
}
// 验证必需字段
const { skill } = skillData;
if (!skill.name || !skill.promptTemplate) {
console.warn(`Skill 文件缺少必需字段: ${filePath}`);
return null;
}
return {
...skill,
source,
sourcePath: filePath,
enabled: skill.enabled ?? true,
};
}
/**
* 解析 Markdown 格式的 Skill
*
* 格式示例:
* ```markdown
* ---
* name: code-review
* description: 代码审查
* parameters:
* focus:
* type: string
* description: 审查重点
* ---
*
* # 代码审查
*
* 请审查以下代码,重点关注 {{focus}}:
*
* {{code}}
* ```
*/
private parseMarkdownSkill(
content: string,
filePath: string
): SkillFile | null {
// 解析 frontmatter
const frontmatterMatch = content.match(/^---\n([\s\S]*?)\n---\n([\s\S]*)$/);
if (!frontmatterMatch) {
// 没有 frontmatter,使用文件名作为 name,整个内容作为 promptTemplate
const name = path.basename(filePath, path.extname(filePath));
return {
skill: {
name,
description: `Skill: ${name}`,
promptTemplate: content.trim(),
},
};
}
const [, frontmatterStr, bodyContent] = frontmatterMatch;
try {
const frontmatter = yaml.parse(frontmatterStr) as Partial<Skill>;
// 如果 frontmatter 中没有 promptTemplate,使用 body 内容
const promptTemplate = frontmatter.promptTemplate || bodyContent.trim();
// 从文件名获取默认 name
const defaultName = path.basename(filePath, path.extname(filePath));
return {
skill: {
name: frontmatter.name || defaultName,
displayName: frontmatter.displayName,
description: frontmatter.description || `Skill: ${defaultName}`,
category: frontmatter.category,
promptTemplate,
parameters: frontmatter.parameters,
keywords: frontmatter.keywords,
version: frontmatter.version,
author: frontmatter.author,
enabled: frontmatter.enabled,
},
};
} catch (error) {
console.warn(`解析 Skill frontmatter 失败: ${filePath}`, error);
return null;
}
}
/**
* 获取用户 Skills 目录
*/
getUserSkillsDir(): string {
const home = process.env.HOME || process.env.USERPROFILE || '';
return path.join(home, '.config', 'ai-terminal', 'skills');
}
/**
* 获取项目 Skills 目录
*/
getProjectSkillsDir(workdir: string = process.cwd()): string {
return path.join(workdir, '.ai-terminal', 'skills');
}
}
/**
* 全局 Skill 加载器实例
*/
export const skillLoader = new SkillLoader();
+358
View File
@@ -0,0 +1,358 @@
/**
* Skill 注册表
*
* 管理所有可用的 Skills,支持:
* - 注册/注销 Skills
* - 按名称/分类/关键词查询
* - 渲染 Skill 提示模板
*/
import type {
Skill,
SkillContext,
SkillExecutionResult,
SkillSearchResult,
SkillRegistryConfig,
} from './types.js';
import { skillLoader } from './loader.js';
import { builtinSkills } from './builtin/index.js';
/**
* Skill 注册表
*/
export class SkillRegistry {
private skills = new Map<string, Skill>();
private config: SkillRegistryConfig;
private initialized = false;
constructor(config: SkillRegistryConfig = {}) {
this.config = {
autoLoad: true,
...config,
};
}
/**
* 初始化注册表
*/
async initialize(workdir: string = process.cwd()): Promise<void> {
if (this.initialized) {
return;
}
// 1. 注册内置 Skills
for (const skill of builtinSkills) {
this.register(skill);
}
// 2. 加载用户 Skills
if (this.config.autoLoad) {
const userDir =
this.config.userSkillsDir || skillLoader.getUserSkillsDir();
const userSkills = await skillLoader.loadFromDirectory(userDir, 'user');
for (const skill of userSkills) {
this.register(skill);
}
}
// 3. 加载项目 Skills
if (this.config.autoLoad) {
const projectDir =
this.config.projectSkillsDir || skillLoader.getProjectSkillsDir(workdir);
const projectSkills = await skillLoader.loadFromDirectory(
projectDir,
'project'
);
for (const skill of projectSkills) {
this.register(skill);
}
}
this.initialized = true;
}
/**
* 注册 Skill
*/
register(skill: Skill): void {
// 项目 Skills 优先级最高,可以覆盖同名的内置/用户 Skills
const existing = this.skills.get(skill.name);
if (existing) {
// 优先级: project > user > builtin
const priority = { project: 3, user: 2, builtin: 1 };
if (priority[skill.source] < priority[existing.source]) {
return; // 不覆盖更高优先级的 Skill
}
}
this.skills.set(skill.name, skill);
}
/**
* 注销 Skill
*/
unregister(name: string): boolean {
return this.skills.delete(name);
}
/**
* 获取 Skill
*/
get(name: string): Skill | undefined {
return this.skills.get(name);
}
/**
* 检查 Skill 是否存在
*/
has(name: string): boolean {
return this.skills.has(name);
}
/**
* 获取所有 Skills
*/
getAll(): Skill[] {
return Array.from(this.skills.values());
}
/**
* 获取启用的 Skills
*/
getEnabled(): Skill[] {
return this.getAll().filter((s) => s.enabled !== false);
}
/**
* 按分类获取 Skills
*/
getByCategory(category: string): Skill[] {
return this.getEnabled().filter((s) => s.category === category);
}
/**
* 获取所有分类
*/
getCategories(): string[] {
const categories = new Set<string>();
for (const skill of this.getEnabled()) {
if (skill.category) {
categories.add(skill.category);
}
}
return Array.from(categories).sort();
}
/**
* 搜索 Skills
*/
search(query: string, limit: number = 10): SkillSearchResult[] {
const queryLower = query.toLowerCase();
const results: SkillSearchResult[] = [];
for (const skill of this.getEnabled()) {
let score = 0;
let matchReason = '';
// 精确名称匹配
if (skill.name.toLowerCase() === queryLower) {
score = 100;
matchReason = '名称精确匹配';
}
// 名称前缀匹配
else if (skill.name.toLowerCase().startsWith(queryLower)) {
score = 80;
matchReason = '名称前缀匹配';
}
// 名称包含匹配
else if (skill.name.toLowerCase().includes(queryLower)) {
score = 60;
matchReason = '名称包含匹配';
}
// 描述匹配
else if (skill.description.toLowerCase().includes(queryLower)) {
score = 40;
matchReason = '描述匹配';
}
// 关键词匹配
else if (
skill.keywords?.some((k) => k.toLowerCase().includes(queryLower))
) {
score = 30;
matchReason = '关键词匹配';
}
// 分类匹配
else if (skill.category?.toLowerCase().includes(queryLower)) {
score = 20;
matchReason = '分类匹配';
}
if (score > 0) {
results.push({ skill, score, matchReason });
}
}
// 按分数降序排序
results.sort((a, b) => b.score - a.score);
return results.slice(0, limit);
}
/**
* 渲染 Skill 提示模板
*/
renderPrompt(
skill: Skill,
params: Record<string, unknown>,
context?: SkillContext
): SkillExecutionResult {
try {
// 验证必需参数
if (skill.parameters) {
for (const [name, param] of Object.entries(skill.parameters)) {
if (param.required && !(name in params)) {
return {
success: false,
error: `缺少必需参数: ${name}`,
};
}
}
}
// 构建变量映射
const variables: Record<string, string> = {};
// 添加参数值
for (const [key, value] of Object.entries(params)) {
variables[key] = String(value);
}
// 添加上下文变量
if (context?.variables) {
for (const [key, value] of Object.entries(context.variables)) {
if (!(key in variables)) {
variables[key] = value;
}
}
}
// 添加默认值
if (skill.parameters) {
for (const [name, param] of Object.entries(skill.parameters)) {
if (!(name in variables) && param.default !== undefined) {
variables[name] = String(param.default);
}
}
}
// 渲染模板
let prompt = skill.promptTemplate;
// 替换 {{variable}} 格式的变量
prompt = prompt.replace(/\{\{(\w+)\}\}/g, (match, varName) => {
if (varName in variables) {
return variables[varName];
}
// 保留未匹配的变量(可能是用户意图保留)
return match;
});
return {
success: true,
prompt,
};
} catch (error) {
return {
success: false,
error: `渲染 Skill 提示失败: ${error instanceof Error ? error.message : String(error)}`,
};
}
}
/**
* 执行 Skill(渲染模板并返回提示)
*/
execute(
name: string,
params: Record<string, unknown>,
context?: SkillContext
): SkillExecutionResult {
const skill = this.get(name);
if (!skill) {
return {
success: false,
error: `Skill 不存在: ${name}`,
};
}
if (skill.enabled === false) {
return {
success: false,
error: `Skill 已禁用: ${name}`,
};
}
return this.renderPrompt(skill, params, context);
}
/**
* 重新加载 Skills
*/
async reload(workdir: string = process.cwd()): Promise<void> {
this.skills.clear();
this.initialized = false;
await this.initialize(workdir);
}
/**
* 获取 Skill 统计信息
*/
getStats(): {
total: number;
enabled: number;
bySource: Record<string, number>;
byCategory: Record<string, number>;
} {
const skills = this.getAll();
const enabled = this.getEnabled();
const bySource: Record<string, number> = {};
const byCategory: Record<string, number> = {};
for (const skill of skills) {
bySource[skill.source] = (bySource[skill.source] || 0) + 1;
if (skill.category) {
byCategory[skill.category] = (byCategory[skill.category] || 0) + 1;
}
}
return {
total: skills.length,
enabled: enabled.length,
bySource,
byCategory,
};
}
}
/**
* 全局 Skill 注册表实例
*/
let skillRegistryInstance: SkillRegistry | null = null;
/**
* 获取全局 Skill 注册表
*/
export function getSkillRegistry(): SkillRegistry {
if (!skillRegistryInstance) {
skillRegistryInstance = new SkillRegistry();
}
return skillRegistryInstance;
}
/**
* 重置全局 Skill 注册表(用于测试)
*/
export function resetSkillRegistry(): void {
skillRegistryInstance = null;
}
+109
View File
@@ -0,0 +1,109 @@
/**
* Skill 系统类型定义
*
* Skill 是可复用的提示模板,类似于 Claude Code 的 Skills 功能。
* 与 Agent 不同,Skill 不是独立的执行单元,而是预定义的提示模板,
* 可以被 Agent 调用来执行特定任务。
*/
/**
* Skill 参数定义
*/
export interface SkillParameter {
/** 参数类型 */
type: 'string' | 'number' | 'boolean' | 'array' | 'object';
/** 参数描述 */
description: string;
/** 是否必需 */
required?: boolean;
/** 默认值 */
default?: unknown;
/** 枚举值(仅 string 类型) */
enum?: string[];
}
/**
* Skill 定义
*/
export interface Skill {
/** Skill 唯一标识 */
name: string;
/** Skill 显示名称 */
displayName?: string;
/** Skill 描述 */
description: string;
/** Skill 分类 */
category?: string;
/** 提示模板(支持变量插值 {{variable}}) */
promptTemplate: string;
/** Skill 参数定义 */
parameters?: Record<string, SkillParameter>;
/** 关键词(用于搜索) */
keywords?: string[];
/** 来源(内置/用户定义/项目) */
source: 'builtin' | 'user' | 'project';
/** 来源路径(用户定义或项目 Skill 的文件路径) */
sourcePath?: string;
/** 是否启用 */
enabled?: boolean;
/** 版本 */
version?: string;
/** 作者 */
author?: string;
}
/**
* Skill 执行上下文
*/
export interface SkillContext {
/** 当前工作目录 */
workdir: string;
/** 额外的上下文变量 */
variables?: Record<string, string>;
}
/**
* Skill 执行结果
*/
export interface SkillExecutionResult {
/** 是否成功 */
success: boolean;
/** 渲染后的提示 */
prompt?: string;
/** 错误信息 */
error?: string;
}
/**
* Skill 文件格式(YAML/JSON)
*/
export interface SkillFile {
/** 文件版本 */
version?: string;
/** Skill 定义 */
skill: Omit<Skill, 'source' | 'sourcePath'>;
}
/**
* Skill 搜索结果
*/
export interface SkillSearchResult {
/** Skill */
skill: Skill;
/** 匹配分数 */
score: number;
/** 匹配原因 */
matchReason: string;
}
/**
* Skill 注册表配置
*/
export interface SkillRegistryConfig {
/** 用户 Skills 目录 */
userSkillsDir?: string;
/** 项目 Skills 目录 */
projectSkillsDir?: string;
/** 是否自动加载 */
autoLoad?: boolean;
}
@@ -0,0 +1,89 @@
/**
* 创建检查点工具
*/
import type { ToolResult } from '../../types/index.js';
import type { ToolWithMetadata } from '../types.js';
import { loadDescription } from '../load_description.js';
import { getCheckpointManager } from '../../checkpoint/index.js';
export const checkpointCreateTool: ToolWithMetadata = {
name: 'checkpoint_create',
description: loadDescription('checkpoint_create'),
metadata: {
name: 'checkpoint_create',
category: 'core',
description: '创建一个新的工作区检查点快照',
keywords: [
'checkpoint',
'create',
'snapshot',
'save',
'检查点',
'快照',
'保存',
],
deferLoading: true,
},
parameters: {
name: {
type: 'string',
description: '检查点名称 (可选)',
required: false,
},
description: {
type: 'string',
description: '检查点描述 (可选)',
required: false,
},
},
execute: async (params: Record<string, unknown>): Promise<ToolResult> => {
const name = params.name as string | undefined;
const description = params.description as string | undefined;
try {
const manager = getCheckpointManager();
if (!manager.isEnabled()) {
return {
success: false,
output: '',
error: '检查点系统已禁用',
};
}
await manager.initialize();
const checkpoint = await manager.createCheckpoint({
name,
description,
trigger: 'manual',
});
const lines = [
`✓ 检查点已创建`,
` ID: ${checkpoint.id}`,
` Commit: ${checkpoint.commitHash.slice(0, 8)}`,
];
if (checkpoint.name) {
lines.push(` 名称: ${checkpoint.name}`);
}
if (checkpoint.filesChanged > 0) {
lines.push(` 文件变更: ${checkpoint.filesChanged} 个`);
}
lines.push(` 时间: ${new Date(checkpoint.timestamp).toLocaleString()}`);
return {
success: true,
output: lines.join('\n'),
};
} catch (error) {
return {
success: false,
output: '',
error: error instanceof Error ? error.message : String(error),
};
}
},
};
@@ -0,0 +1,158 @@
/**
* 检查点差异工具
*/
import type { ToolResult } from '../../types/index.js';
import type { ToolWithMetadata } from '../types.js';
import { loadDescription } from '../load_description.js';
import { getCheckpointManager } from '../../checkpoint/index.js';
export const checkpointDiffTool: ToolWithMetadata = {
name: 'checkpoint_diff',
description: loadDescription('checkpoint_diff'),
metadata: {
name: 'checkpoint_diff',
category: 'core',
description: '显示检查点与当前工作区的差异',
keywords: [
'checkpoint',
'diff',
'compare',
'changes',
'检查点',
'差异',
'比较',
'变更',
],
deferLoading: true,
},
parameters: {
checkpoint_id: {
type: 'string',
description: '检查点 ID 或 commit hash (默认为最近的检查点)',
required: false,
},
file: {
type: 'string',
description: '指定文件路径查看详细差异 (可选)',
required: false,
},
},
execute: async (params: Record<string, unknown>): Promise<ToolResult> => {
const checkpointId = params.checkpoint_id as string | undefined;
const file = params.file as string | undefined;
try {
const manager = getCheckpointManager();
if (!manager.isEnabled()) {
return {
success: false,
output: '',
error: '检查点系统已禁用',
};
}
await manager.initialize();
// 获取目标检查点
let targetCheckpoint;
if (checkpointId) {
targetCheckpoint = await manager.getCheckpoint(checkpointId);
if (!targetCheckpoint) {
return {
success: false,
output: '',
error: `找不到检查点: ${checkpointId}`,
};
}
} else {
targetCheckpoint = await manager.getLatestCheckpoint();
if (!targetCheckpoint) {
return {
success: true,
output: '暂无检查点',
};
}
}
// 显示文件详细差异
if (file) {
const fileDiff = await manager.getFileDiff(targetCheckpoint.id, file);
const lines = [
`文件差异: ${file}`,
`检查点: ${targetCheckpoint.commitHash.slice(0, 8)}`,
`变更类型: ${fileDiff.type}`,
'',
];
if (fileDiff.patch) {
lines.push('```diff');
lines.push(fileDiff.patch);
lines.push('```');
} else if (fileDiff.type === 'added') {
lines.push('(新文件)');
} else if (fileDiff.type === 'deleted') {
lines.push('(已删除)');
}
return {
success: true,
output: lines.join('\n'),
};
}
// 显示概要差异
const diff = await manager.getDiff(targetCheckpoint.id);
if (diff.files.length === 0) {
return {
success: true,
output: `检查点 ${targetCheckpoint.commitHash.slice(0, 8)} 与当前工作区相同`,
};
}
const lines = [
`检查点 ${targetCheckpoint.commitHash.slice(0, 8)} 与当前工作区的差异:`,
'',
` +${diff.totalInsertions} 行添加 -${diff.totalDeletions} 行删除`,
'',
'变更的文件:',
];
for (const fileChange of diff.files) {
const symbol =
fileChange.type === 'added'
? '+'
: fileChange.type === 'deleted'
? '-'
: fileChange.type === 'renamed'
? 'R'
: 'M';
let line = ` ${symbol} ${fileChange.path}`;
if (fileChange.oldPath) {
line = ` ${symbol} ${fileChange.oldPath} -> ${fileChange.path}`;
}
if (fileChange.insertions || fileChange.deletions) {
line += ` (+${fileChange.insertions || 0} -${fileChange.deletions || 0})`;
}
lines.push(line);
}
return {
success: true,
output: lines.join('\n'),
};
} catch (error) {
return {
success: false,
output: '',
error: error instanceof Error ? error.message : String(error),
};
}
},
};
@@ -0,0 +1,91 @@
/**
* 列出检查点工具
*/
import type { ToolResult } from '../../types/index.js';
import type { ToolWithMetadata } from '../types.js';
import { loadDescription } from '../load_description.js';
import { getCheckpointManager } from '../../checkpoint/index.js';
export const checkpointListTool: ToolWithMetadata = {
name: 'checkpoint_list',
description: loadDescription('checkpoint_list'),
metadata: {
name: 'checkpoint_list',
category: 'core',
description: '列出所有可用的检查点',
keywords: [
'checkpoint',
'list',
'show',
'history',
'检查点',
'列表',
'历史',
],
deferLoading: true,
},
parameters: {
limit: {
type: 'number',
description: '最多显示的检查点数量 (默认 10)',
required: false,
},
},
execute: async (params: Record<string, unknown>): Promise<ToolResult> => {
const limit = (params.limit as number) || 10;
try {
const manager = getCheckpointManager();
if (!manager.isEnabled()) {
return {
success: false,
output: '',
error: '检查点系统已禁用',
};
}
await manager.initialize();
const checkpoints = await manager.listCheckpoints();
if (checkpoints.length === 0) {
return {
success: true,
output: '暂无检查点',
};
}
const displayCheckpoints = checkpoints.slice(0, limit);
const lines = [`共 ${checkpoints.length} 个检查点:\n`];
for (const cp of displayCheckpoints) {
const date = new Date(cp.timestamp).toLocaleString();
const hash = cp.commitHash.slice(0, 8);
const name = cp.name ? ` "${cp.name}"` : '';
const files = cp.filesChanged > 0 ? ` (${cp.filesChanged} files)` : '';
lines.push(` ${hash}${name}${files}`);
lines.push(` ${cp.description || cp.trigger}`);
lines.push(` ${date}`);
lines.push('');
}
if (checkpoints.length > limit) {
lines.push(` ... 还有 ${checkpoints.length - limit} 个检查点`);
}
return {
success: true,
output: lines.join('\n'),
};
} catch (error) {
return {
success: false,
output: '',
error: error instanceof Error ? error.message : String(error),
};
}
},
};

Some files were not shown because too many files have changed in this diff Show More