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:
@@ -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,
|
||||
},
|
||||
},
|
||||
};
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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,
|
||||
};
|
||||
@@ -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 };
|
||||
@@ -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,
|
||||
};
|
||||
@@ -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,
|
||||
};
|
||||
@@ -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();
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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,
|
||||
];
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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();
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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,
|
||||
};
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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();
|
||||
@@ -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;
|
||||
});
|
||||
}
|
||||
@@ -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}`;
|
||||
}
|
||||
}
|
||||
@@ -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';
|
||||
}
|
||||
@@ -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 };
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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 });
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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,
|
||||
};
|
||||
@@ -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 };
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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}`;
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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',
|
||||
},
|
||||
};
|
||||
@@ -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 };
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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();
|
||||
@@ -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('');
|
||||
}
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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,
|
||||
}));
|
||||
}
|
||||
@@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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; // 默认启用
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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 };
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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}`));
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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
@@ -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
@@ -0,0 +1,6 @@
|
||||
/**
|
||||
* 缓存模块导出
|
||||
*/
|
||||
|
||||
export { DiskCache, createDiskCache } from './disk-cache.js';
|
||||
export type { CacheEntry } from './disk-cache.js';
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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();
|
||||
}
|
||||
@@ -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
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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';
|
||||
@@ -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();
|
||||
@@ -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();
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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,
|
||||
];
|
||||
@@ -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';
|
||||
@@ -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();
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user