sdk_generate/internal/prompts/sdk_generate.go

463 lines
14 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 - 完整优化版(兼容 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 生成完成 ===" 表示任务完成。
请开始工作,一次性生成所有文件。`