// sdk_generate.go - 完整优化版(兼容 DeepSeek GA 版本) package prompts import ( "context" "fmt" "log" "runtime/debug" "sdk-generator/internal/data/impl" "sdk-generator/internal/data/model" "sdk-generator/internal/entitys" "sdk-generator/internal/pkg" "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 } // SDKGeneratorOption 配置选项 type SDKGeneratorOption func(*SDKGenerator) func WithMaxIterations(n int) SDKGeneratorOption { return func(g *SDKGenerator) { g.maxIterations = n } } func WithTimeout(t time.Duration) SDKGeneratorOption { return func(g *SDKGenerator) { g.timeout = t } } func WithModel(model string) SDKGeneratorOption { return func(g *SDKGenerator) { g.model = model } } 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 代码 func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, task *entitys.Task, needImplement string, logImpl *impl.AiGenerateLogImpl) (string, *entitys.Usage, error) { ctx, cancel := context.WithTimeout(ctx, g.timeout) defer cancel() systemPrompt := g.buildSystemPrompt(task.Name, needImplement) messages := []openai.ChatCompletionMessage{ { Role: openai.ChatMessageRoleSystem, Content: systemPrompt, }, { Role: openai.ChatMessageRoleUser, Content: fmt.Sprintf("请根据以下文档生成完整的 SDK 代码,必须包含所有 6 个文件:\n\n**⚠️ 重要规则**:\n1. 调用工具时,**不要输出任何额外文本**(包括句号、逗号、空格等)\n2. 直接调用工具,等待工具返回结果\n3. 工具返回后,**立即开始生成代码**\n4. 不要重复调用已调用过的工具\n\n%s", doc), }, } tools := g.cryptoManager.GetToolsForLLM() iteration := 0 var allResults []string calledTools := make(map[string]bool) useAge := &entitys.Usage{} forceTextGeneration := false guidedToGenerate := 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: nil, Temperature: 0.1, MaxTokens: 16384, } // 如果已经调用了足够的工具,强制切换到文本生成模式 if len(calledTools) >= 2 && iteration > 2 { req.ToolChoice = nil // 兼容所有版本 req.Tools = nil forceTextGeneration = true log.Printf("🛑 强制切换到文本生成模式,已调用工具数: %d", len(calledTools)) } resp, err := g.openaiClient.CreateChatCompletion(ctx, req) if err != nil { return "", useAge, fmt.Errorf("调用 OpenAI 失败: %v", err) } choice := resp.Choices[0] msg := choice.Message logImpl.Add(ctx, &model.AiGenerateLog{ TaskID: task.TasKModel.TaskID, Type: entitys.StatusGenerateCode.String(), RequestContent: pkg.JsonStringIgonErr(req.Messages), ResponseContent: msg.Content, ToolSelect: pkg.JsonStringIgonErr(msg.ToolCalls), }) // 清理无意义内容 cleanedContent := g.cleanContent(msg.Content) if cleanedContent != msg.Content { log.Printf("🧹 清理了 AI 的无意义输出: %q -> %q", msg.Content, cleanedContent) msg.Content = cleanedContent } // 检查是否完成 if strings.Contains(msg.Content, "=== SDK 生成完成 ===") { log.Printf("✅ SDK 生成完成") return g.mergeResults(msg.Content, allResults), useAge, nil } if g.checkAllFilesGenerated(msg.Content) { log.Printf("✅ 所有 6 个文件已生成") return g.mergeResults(msg.Content, allResults), useAge, nil } // 【修复】处理没有工具调用的情况 if len(msg.ToolCalls) == 0 { // 如果有内容输出 if len(msg.Content) > 0 { messages = append(messages, msg) // 如果强制生成模式,直接返回 if forceTextGeneration { log.Printf("✅ 强制文本生成模式,返回结果") return g.mergeResults(msg.Content, allResults), useAge, nil } // 如果生成了部分代码,引导继续 if strings.Contains(msg.Content, "// File:") { messages = append(messages, openai.ChatCompletionMessage{ Role: openai.ChatMessageRoleUser, Content: "请继续生成剩余的文件,完成后输出 '=== SDK 生成完成 ==='", }) } continue } // 没有内容也没有工具调用 log.Printf("⚠️ AI 没有输出内容也没有调用工具") if forceTextGeneration { // 强制模式下,添加更明确的指令 messages = append(messages, openai.ChatCompletionMessage{ Role: openai.ChatMessageRoleUser, Content: "请立即生成完整的 SDK 代码,包含所有 6 个文件。直接输出代码,不要调用工具。", }) } continue } // 【修复】处理工具调用时的无意义内容 if len(msg.ToolCalls) > 0 { content := strings.TrimSpace(msg.Content) if g.isMeaninglessContent(content) { log.Printf("🧹 清除了工具调用时的无意义文本: %q", msg.Content) msg.Content = "" } } // 记录工具调用 var toolNames []string for _, tc := range msg.ToolCalls { toolNames = append(toolNames, tc.Function.Name) calledTools[tc.Function.Name] = true } log.Printf("🔧 AI 调用工具: %v, 迭代: %d, 已调用: %v", toolNames, iteration, getCalledToolsList(calledTools)) messages = append(messages, msg) // 【修复】先执行工具调用,再判断是否完成 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 // 【修复】工具执行完成后,再引导生成代码(避免遗漏工具结果) if len(calledTools) >= 2 && !guidedToGenerate { guidedToGenerate = true log.Printf("✅ 所有加密工具已调用完成(共 %d 个),引导 AI 生成代码", len(calledTools)) messages = append(messages, openai.ChatCompletionMessage{ Role: openai.ChatMessageRoleUser, Content: fmt.Sprintf("所有加密工具已调用完成(共 %d 个),请根据工具返回的结果,现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='。不要再次调用工具。", len(calledTools)), }) forceTextGeneration = true } else if len(allResults) > 0 { // 如果已经有工具结果,提示生成 messages = append(messages, openai.ChatCompletionMessage{ Role: openai.ChatMessageRoleUser, Content: fmt.Sprintf("✅ 已获取 %d 个加密实现,请现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='。不要再次调用工具。", len(allResults)), }) } } } // cleanContent 清理无意义内容 func (g *SDKGenerator) cleanContent(content string) string { trimmed := strings.TrimSpace(content) if trimmed == "" { return "" } if g.isMeaninglessContent(trimmed) { return "" } return content } // isMeaninglessContent 检查内容是否无意义 func (g *SDKGenerator) isMeaninglessContent(content string) bool { punctuations := []string{"。", "、", ".", ",", ",", ";", ";", "!", "!", "?", "?"} for _, p := range punctuations { if content == p { return true } } if len([]rune(content)) <= 3 { for _, r := range content { isPunct := false for _, p := range punctuations { if string(r) == p { isPunct = true break } } if !isPunct { return false } } return true } return false } func getCalledToolsList(calledTools map[string]bool) []string { var list []string for tool := range calledTools { list = append(list, tool) } return list } // processToolCalls 并行处理工具调用 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 } if strings.Contains(code, "SM3Hash") || strings.Contains(code, "GenerateNonce") || strings.Contains(code, "BuildSignString") || strings.Contains(code, "HMAC") { log.Printf("📝 代码中已包含加密实现,跳过合并") 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 ## 工具调用规则(必须严格遵守) - 调用工具时,**只调用工具,不输出任何文本**(包括句号、逗号、标点符号) - 不要在工具调用前后添加任何注释或说明 - 工具调用完成后,**立即**开始生成代码 - 每个工具**只调用一次**,不要重复调用 ## 工作流程 1. 分析文档中的加密需求 → 调用对应的工具获取实现(只调用需要的工具,通常2-3个) 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 生成完成 ===" 表示任务完成。 请开始工作,一次性生成所有文件。`