// server_generator.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" ) // ServerGenerator 服务端代码生成器 type ServerGenerator struct { openaiClient *openai.Client cryptoManager *crypt.CryptoSkillManager maxIterations int timeout time.Duration model string } // ServerGeneratorOption 服务端生成器配置选项 type ServerGeneratorOption func(*ServerGenerator) // WithServerMaxIterations 设置最大迭代次数 func WithServerMaxIterations(n int) ServerGeneratorOption { return func(g *ServerGenerator) { g.maxIterations = n } } // WithServerTimeout 设置超时时间 func WithServerTimeout(t time.Duration) ServerGeneratorOption { return func(g *ServerGenerator) { g.timeout = t } } // WithServerModel 设置模型 func WithServerModel(model string) ServerGeneratorOption { return func(g *ServerGenerator) { g.model = model } } // NewServerGenerator 创建服务端生成器 func NewServerGenerator(client *openai.Client, opts ...ServerGeneratorOption) *ServerGenerator { g := &ServerGenerator{ openaiClient: client, cryptoManager: crypt.NewCryptoSkillManager(), maxIterations: 15, timeout: 10 * time.Minute, model: openai.GPT4, } for _, opt := range opts { opt(g) } return g } // GenerateServer 生成服务端代码 func (g *ServerGenerator) GenerateServer(ctx context.Context, doc string, serverName string, needImplement string) (string, *entitys.Usage, error) { ctx, cancel := context.WithTimeout(ctx, g.timeout) defer cancel() systemPrompt := g.buildSystemPrompt(serverName, needImplement) messages := []openai.ChatCompletionMessage{ { Role: openai.ChatMessageRoleSystem, Content: systemPrompt, }, { Role: openai.ChatMessageRoleUser, Content: fmt.Sprintf("请根据以下 API 对接文档生成完整的服务端代码:\n\n%s", doc), }, } tools := g.cryptoManager.GetToolsForLLM() iteration := 0 useAge := &entitys.Usage{} var allResults []string calledTools := make(map[string]bool) 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, } 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, "=== 服务端生成完成 ===") { log.Printf("✅ 服务端生成完成") if len(allResults) > 0 { return g.mergeResults(msg.Content, allResults), useAge, nil } return msg.Content, useAge, nil } // 检查是否生成了所有文件 if g.checkAllFilesGenerated(msg.Content) { allFilesGenerated = true log.Printf("✅ 所有文件已生成") if len(allResults) > 0 { return g.mergeResults(msg.Content, allResults), useAge, 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)) if len(allResults) > 0 && !allFilesGenerated { messages = append(messages, openai.ChatCompletionMessage{ Role: openai.ChatMessageRoleAssistant, Content: fmt.Sprintf("已获取 %d 个加密实现,请现在生成完整的服务端代码。完成后输出 '=== 服务端生成完成 ==='", len(allResults)), }) continue } } messages = append(messages, msg) if len(msg.ToolCalls) == 0 { if len(allResults) > 0 { return g.mergeResults(msg.Content, allResults), useAge, 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: "所有加密工具已调用完成,请现在生成完整的服务端代码。完成后输出 '=== 服务端生成完成 ==='", }) 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 messages = append(messages, openai.ChatCompletionMessage{ Role: openai.ChatMessageRoleAssistant, Content: fmt.Sprintf("✅ 已获取 %d 个加密实现,请现在生成完整的服务端代码。完成后输出 '=== 服务端生成完成 ==='", len(allResults)), }) } } // processToolCalls 并行处理工具调用(带去重和 panic 保护) func (g *ServerGenerator) 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() // panic 保护 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 *ServerGenerator) checkAllFilesGenerated(content string) bool { requiredFiles := []string{ "go.mod", "main.go", "service.go", "request.go", "response.go", "middleware.go", "crypto.go", } foundCount := 0 for _, file := range requiredFiles { if strings.Contains(content, "// File:") && strings.Contains(content, file) { foundCount++ } } return foundCount >= 5 } // buildSystemPrompt 构建系统提示 func (g *ServerGenerator) buildSystemPrompt(serverName, needImplement string) string { return fmt.Sprintf(`你是一个资深的 Go 语言工程师,擅长根据 API 对接文档生成高质量的服务端骨架代码。 ## 任务 根据用户提供的 API 对接文档,生成一个完整的 Go HTTP 服务端代码工程。 %s ## ⚠️ 重要:输出格式要求(必须严格遵守) 每个文件必须以以下格式输出: // File: {{server_name}}/文件名.go `+"```go"+` package {{server_name}} // ... 代码内容 ... `+"```"+` ## 必须生成的文件列表(缺一不可) %s/ ├── go.mod ├── cmd/ │ └── server/ │ └── main.go # 服务启动入口 ├── internal/ │ ├── service/ │ │ └── service.go #接口实现 │ ├── router/ │ │ └── router.go #接口定义 │ ├── test/ │ │ └── example_test.go #模拟api请求示例最小单元测试 │ ├── biz/ │ │ └── biz.go #业务逻辑实现 │ ├── entities/ │ │ ├── request.go # 请求参数封装结构体 │ │ └── response.go # 响应参数封装结构体 │ ├── middleware/ │ │ └── middleware.go # 中间件实现 │ └── config/ │ └── config.go # 配置文件 └── pkg/ └── crypto/ └── crypto.go # 从加密工具获取的实现 ## 重要:工具列表 - 加密工具: %s ## 工具调用策略 - 只在文档明确提到加密/签名时才调用工具 - 工具返回的是完整的实现代码,直接集成到 pkg/crypto/crypto.go 中 - 每个工具**只调用一次**,不要重复调用 ## 工作流程 1. 分析文档中的加密需求 → 调用对应的工具获取实现 2. 将工具返回的代码放到 pkg/crypto/crypto.go 中 3. **然后立即生成所有其他文件** ## 代码规范 - 使用 Fiber 框架 (github.com/gofiber/fiber/v2) - 遵循 Effective Go 规范 - Handler 中的业务逻辑用 TODO 占位 ## 关键要求 - Go 版本为 1.26 - **必须生成所有文件,每个文件都必须完整输出** - 生成的代码必须可直接编译运行 ## 完成标志 当你生成了所有文件后,输出 "=== 服务端生成完成 ===" 表示任务完成。 请开始工作,一次性生成所有文件。`, needImplement, serverName, g.cryptoManager.GetToolDescriptions(), ) } // mergeResults 合并加密实现到最终代码 func (g *ServerGenerator) 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() }