sdk_generate/internal/prompts/server_genrate.go

411 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.

// 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()
}