sdk_generate/internal/prompts/sdk_generate.go

397 lines
11 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

// sdk_generate.go - 完整优化版
package prompts
import (
"context"
"fmt"
"log"
"runtime/debug"
"sdk-generator/internal/entitys"
"strings"
"sync"
"time"
"sdk-generator/internal/prompts/tools/crypt"
"github.com/sashabaranov/go-openai"
)
// SDKGenerator SDK 代码生成器
type SDKGenerator struct {
openaiClient *openai.Client
cryptoManager *crypt.CryptoSkillManager
maxIterations int // 最大迭代次数,防止死循环
timeout time.Duration // 总超时时间
model string // OpenAI 模型
}
// SDKGeneratorOption 配置选项
type SDKGeneratorOption func(*SDKGenerator)
// WithMaxIterations 设置最大迭代次数
func WithMaxIterations(n int) SDKGeneratorOption {
return func(g *SDKGenerator) {
g.maxIterations = n
}
}
// WithTimeout 设置超时时间
func WithTimeout(t time.Duration) SDKGeneratorOption {
return func(g *SDKGenerator) {
g.timeout = t
}
}
// WithModel 设置模型
func WithModel(model string) SDKGeneratorOption {
return func(g *SDKGenerator) {
g.model = model
}
}
// NewSDKGenerator 创建 SDK 生成器
func NewSDKGenerator(client *openai.Client, opts ...SDKGeneratorOption) *SDKGenerator {
g := &SDKGenerator{
openaiClient: client,
cryptoManager: crypt.NewCryptoSkillManager(),
maxIterations: 15,
timeout: 10 * time.Minute,
model: openai.GPT4,
}
for _, opt := range opts {
opt(g)
}
return g
}
// GenerateSDK 生成 SDK 代码 - AI 自主决策流程
func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, sdkName string, needImplement string) (string, *entitys.Usage, error) {
ctx, cancel := context.WithTimeout(ctx, g.timeout)
defer cancel()
systemPrompt := g.buildSystemPrompt(sdkName, needImplement)
messages := []openai.ChatCompletionMessage{
{
Role: openai.ChatMessageRoleSystem,
Content: systemPrompt,
},
{
Role: openai.ChatMessageRoleUser,
Content: fmt.Sprintf("请根据以下文档生成完整的 SDK 代码,必须包含所有 6 个文件:\n\n%s", doc),
},
}
tools := g.cryptoManager.GetToolsForLLM()
iteration := 0
var allResults []string
calledTools := make(map[string]bool)
useAge := &entitys.Usage{}
allFilesGenerated := false
for {
select {
case <-ctx.Done():
return "", useAge, ctx.Err()
default:
}
iteration++
if iteration > g.maxIterations {
stats := g.cryptoManager.GetCalledTools()
log.Printf("调试: 工具调用统计: %v", stats)
return "", useAge, fmt.Errorf("超过最大迭代次数 %d工具调用统计: %v", g.maxIterations, stats)
}
req := openai.ChatCompletionRequest{
Model: g.model,
Messages: messages,
Tools: tools,
ToolChoice: "auto",
Temperature: 0.1,
MaxTokens: 16384, // 增大 token 以生成完整代码
}
resp, err := g.openaiClient.CreateChatCompletion(ctx, req)
if err != nil {
return "", useAge, fmt.Errorf("调用 OpenAI 失败: %v", err)
}
choice := resp.Choices[0]
msg := choice.Message
// 检查是否包含完成标志
if strings.Contains(msg.Content, "=== SDK 生成完成 ===") {
log.Printf("✅ SDK 生成完成")
if len(allResults) > 0 {
return g.mergeResults(msg.Content, allResults), nil, nil
}
return msg.Content, useAge, nil
}
// 检查是否生成了所有文件
if g.checkAllFilesGenerated(msg.Content) {
allFilesGenerated = true
log.Printf("✅ 所有 6 个文件已生成")
if len(allResults) > 0 {
return g.mergeResults(msg.Content, allResults), nil, nil
}
return msg.Content, useAge, nil
}
// 记录工具调用信息
if len(msg.ToolCalls) > 0 {
var toolNames []string
for _, tc := range msg.ToolCalls {
toolNames = append(toolNames, tc.Function.Name)
calledTools[tc.Function.Name] = true
}
log.Printf("🔧 AI 调用工具: %v, 迭代次数: %d", toolNames, iteration)
} else {
log.Printf("📝 AI 回复长度: %d 字符", len(msg.Content))
// 如果 AI 没有调用工具也没有完成,提示它继续
if len(allResults) > 0 && !allFilesGenerated {
messages = append(messages, openai.ChatCompletionMessage{
Role: openai.ChatMessageRoleAssistant,
Content: fmt.Sprintf("已获取 %d 个加密实现,请现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='", len(allResults)),
})
continue
}
}
messages = append(messages, msg)
if len(msg.ToolCalls) == 0 {
if len(allResults) > 0 {
return g.mergeResults(msg.Content, allResults), nil, nil
}
return msg.Content, useAge, nil
}
// 检查是否所有工具都已调用过
allCalled := true
for _, tc := range msg.ToolCalls {
if !calledTools[tc.Function.Name] {
allCalled = false
break
}
}
// 如果 AI 重复调用已调用的工具,强制终止并引导生成代码
if allCalled && iteration > 2 {
log.Printf("⚠️ AI 在重复调用已调用的工具,强制引导生成代码")
messages = append(messages, openai.ChatCompletionMessage{
Role: openai.ChatMessageRoleAssistant,
Content: "所有加密工具已调用完成,请现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='",
})
continue
}
// 并行处理所有工具调用
toolResults, results, err := g.processToolCalls(ctx, msg.ToolCalls)
if err != nil {
return "", useAge, err
}
allResults = append(allResults, results...)
messages = append(messages, toolResults...)
useAge.PromptTokens += resp.Usage.PromptTokens
useAge.CompletionTokens += resp.Usage.CompletionTokens
useAge.TotalTokens += resp.Usage.TotalTokens
// 添加一个明确的提示,告诉 AI 继续
messages = append(messages, openai.ChatCompletionMessage{
Role: openai.ChatMessageRoleAssistant,
Content: fmt.Sprintf("✅ 已获取 %d 个加密实现,请现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='", len(allResults)),
})
}
}
// processToolCalls 并行处理工具调用(带去重和 panic 保护)
func (g *SDKGenerator) processToolCalls(ctx context.Context, toolCalls []openai.ToolCall) ([]openai.ChatCompletionMessage, []string, error) {
// 去重
seen := make(map[string]bool)
var uniqueCalls []openai.ToolCall
for _, tc := range toolCalls {
key := tc.Function.Name + "|" + tc.Function.Arguments
if !seen[key] {
seen[key] = true
uniqueCalls = append(uniqueCalls, tc)
}
}
if len(uniqueCalls) < len(toolCalls) {
log.Printf("🔧 工具调用去重: %d -> %d", len(toolCalls), len(uniqueCalls))
}
var wg sync.WaitGroup
resultChan := make(chan openai.ChatCompletionMessage, len(uniqueCalls))
resultDataChan := make(chan string, len(uniqueCalls))
errChan := make(chan error, len(uniqueCalls))
for _, toolCall := range uniqueCalls {
wg.Add(1)
go func(tc openai.ToolCall) {
defer wg.Done()
defer func() {
if r := recover(); r != nil {
errChan <- fmt.Errorf("工具 %s 执行时发生 panic: %v\n堆栈: %s",
tc.Function.Name, r, string(debug.Stack()))
}
}()
select {
case <-ctx.Done():
errChan <- ctx.Err()
return
default:
}
result, err := g.cryptoManager.ExecuteTool(ctx, tc)
if err != nil {
errChan <- fmt.Errorf("执行工具 %s 失败: %v", tc.Function.Name, err)
return
}
resultChan <- openai.ChatCompletionMessage{
Role: openai.ChatMessageRoleTool,
Content: result,
ToolCallID: tc.ID,
}
resultDataChan <- result
}(toolCall)
}
wg.Wait()
close(resultChan)
close(resultDataChan)
close(errChan)
// 检查错误
for err := range errChan {
if err != nil {
return nil, nil, err
}
}
// 收集结果
var toolResults []openai.ChatCompletionMessage
var results []string
for msg := range resultChan {
toolResults = append(toolResults, msg)
}
for data := range resultDataChan {
results = append(results, data)
}
return toolResults, results, nil
}
// checkAllFilesGenerated 检查是否生成了所有必要文件
func (g *SDKGenerator) checkAllFilesGenerated(content string) bool {
requiredFiles := []string{
"go.mod",
"client.go",
"types.go",
"crypto.go",
"errors.go",
"example_test.go",
}
foundCount := 0
for _, file := range requiredFiles {
if strings.Contains(content, "// File:") && strings.Contains(content, file) {
foundCount++
}
}
return foundCount >= len(requiredFiles)
}
// buildSystemPrompt 构建 System Prompt
func (g *SDKGenerator) buildSystemPrompt(sdkName, needImplement string) string {
if sdkName == "" {
sdkName = "my-sdk"
}
basePrompt := strings.Replace(SdkGeneratePrompt, "{{sdk_name}}", sdkName, -1)
basePrompt = strings.Replace(basePrompt, "{{needImplement}}", needImplement, -1)
return fmt.Sprintf(basePrompt, g.cryptoManager.GetToolDescriptions())
}
// mergeResults 合并加密实现到最终代码
func (g *SDKGenerator) mergeResults(code string, results []string) string {
if len(results) == 0 {
return code
}
var sb strings.Builder
sb.WriteString(code)
sb.WriteString("\n\n## 加密实现\n\n")
sb.WriteString("以下是从加密工具获取的完整实现:\n\n")
for i, result := range results {
sb.WriteString(fmt.Sprintf("### 加密实现 %d\n\n", i+1))
sb.WriteString(result)
sb.WriteString("\n\n")
}
return sb.String()
}
const SdkGeneratePrompt = `你是一个资深的 Go 语言工程师,擅长生成高质量的 SDK 代码。
## 任务
根据用户提供的 API 接口文档,生成一个**完整的** Go SDK 代码工程。
{{needImplement}}
## ⚠️ 重要:输出格式要求(必须严格遵守)
每个文件必须以以下格式输出:
// File: {{sdk_name}}/文件名.go
` + "```go" + `
package {{sdk_name}}
// ... 代码内容 ...
` + "```" + `
## 必须生成的文件列表(缺一不可)
{{sdk_name}}/
├── go.mod # 模块名为 {{sdk_name}}
├── client.go # 客户端主文件,包含 NewClient 和所有 API 方法
├── types.go # 所有请求/响应结构体定义
├── crypto.go # 从工具获取的加密实现
├── errors.go # 错误类型定义
└── example_test.go # 使用示例
## 重要:你有加密工具可用
%s
## 工具调用策略
- 只在文档明确提到特定加密算法时才调用工具
- 工具返回的是完整的实现代码,直接集成到 SDK 的 crypto.go 中
- 每个工具**只调用一次**,不要重复调用
## 工作流程
1. 分析文档中的加密需求 → 调用对应的工具获取实现
2. 将工具返回的代码放到 crypto.go 中
3. **然后立即生成所有其他文件**client.go、types.go、go.mod、errors.go、example_test.go
## 代码规范
- 遵循 Effective Go 规范
- 所有导出类型/方法必须有 godoc 注释
- context.Context 作为方法的第一个参数
## 关键要求
- Go 版本为 1.21
- **必须生成所有 6 个文件,每个文件都必须完整输出**
- 不要引用任何不存在的远程 GitHub 仓库
- 生成的代码必须可直接编译运行
## 完成标志
当你生成了所有 6 个文件后,输出 "=== SDK 生成完成 ===" 表示任务完成。
请开始工作,一次性生成所有文件。`