463 lines
14 KiB
Go
463 lines
14 KiB
Go
// 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 生成完成 ===" 表示任务完成。
|
||
|
||
请开始工作,一次性生成所有文件。`
|