This commit is contained in:
parent
77f26ec193
commit
8e7066b4ac
|
|
@ -24,7 +24,8 @@ func InitializeApp(configConfig *config.Config, allLogger log.AllLogger) (*serve
|
|||
db, cleanup := utils.NewGormDb(configConfig)
|
||||
aiGenerateDocImpl := impl.NewAiGenerateDocImpl(db)
|
||||
aiGenerateTaskImpl := impl.NewAiGenerateTaskImpl(db)
|
||||
sdkGeneratorBiz := biz.NewSDKGeneratorService(configConfig, aiGenerateDocImpl, aiGenerateTaskImpl)
|
||||
aiGenerateLogImpl := impl.NewAiGenerateLogImpl(db)
|
||||
sdkGeneratorBiz := biz.NewSDKGeneratorService(configConfig, aiGenerateDocImpl, aiGenerateTaskImpl, aiGenerateLogImpl)
|
||||
sdkService := service.NewSDKService(sdkGeneratorBiz)
|
||||
pageService := service.NewPageService()
|
||||
appModule := router.NewAppModule(configConfig, sdkService, pageService)
|
||||
|
|
|
|||
|
|
@ -33,14 +33,16 @@ type SDKGeneratorBiz struct {
|
|||
config *config.Config
|
||||
docImpl *impl.AiGenerateDocImpl
|
||||
taskImpl *impl.AiGenerateTaskImpl
|
||||
logImpl *impl.AiGenerateLogImpl
|
||||
}
|
||||
|
||||
func NewSDKGeneratorService(cfg *config.Config, docImpl *impl.AiGenerateDocImpl, taskImpl *impl.AiGenerateTaskImpl) *SDKGeneratorBiz {
|
||||
func NewSDKGeneratorService(cfg *config.Config, docImpl *impl.AiGenerateDocImpl, taskImpl *impl.AiGenerateTaskImpl, logImpl *impl.AiGenerateLogImpl) *SDKGeneratorBiz {
|
||||
// 创建 OpenAI 客户端配置
|
||||
return &SDKGeneratorBiz{
|
||||
config: cfg,
|
||||
docImpl: docImpl,
|
||||
taskImpl: taskImpl,
|
||||
logImpl: logImpl,
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -184,7 +186,6 @@ func (s *SDKGeneratorBiz) processTask(ctx context.Context, task *entitys.Task) {
|
|||
err = s.creatFile(ctx, task)
|
||||
if err != nil {
|
||||
// 后处理失败不中断流程
|
||||
|
||||
s.failTask(ctx, task, fmt.Sprintf("[creatFile]: %v", err))
|
||||
return
|
||||
}
|
||||
|
|
@ -337,7 +338,7 @@ func (s *SDKGeneratorBiz) generateCode(ctx context.Context, task *entitys.Task)
|
|||
prompts.WithTimeout(10*time.Minute),
|
||||
prompts.WithModel(task.TasKModel.LlmModel),
|
||||
)
|
||||
task.Generate, usage, err = sdkGen.GenerateSDK(ctx, finalPrompt.String(), task.Name, NeedImplement)
|
||||
task.Generate, usage, err = sdkGen.GenerateSDK(ctx, finalPrompt.String(), task, NeedImplement, s.logImpl)
|
||||
case entitys.DocTypeServerBoilerplate:
|
||||
serverGen := prompts.NewServerGenerator(
|
||||
task.CallLLM.Client,
|
||||
|
|
@ -352,12 +353,33 @@ func (s *SDKGeneratorBiz) generateCode(ctx context.Context, task *entitys.Task)
|
|||
}
|
||||
|
||||
func (s *SDKGeneratorBiz) valid(ctx context.Context, task *entitys.Task) (usage *entitys.Usage, err error) {
|
||||
log.Printf("valid(): task.CallLLM=%p, task.CallLLM.Client=%p, task.CallLLM.ApiKey prefix=%s, task.Name=%s",
|
||||
task.CallLLM, task.CallLLM.Client, prefix(task.CallLLM.ApiKey), task.Name)
|
||||
// ========== 判断输入大小 ==========
|
||||
// 估算 token 数量:1 token ≈ 4 字符(中英文混合)
|
||||
inputSize := len(task.RefinedDoc) + len(task.Generate)
|
||||
estimatedTokens := inputSize / 4
|
||||
log.Printf("valid(): 输入大小: %d 字符, 估算 token: %d", inputSize, estimatedTokens)
|
||||
|
||||
validRes, useAge, err := task.CallLLM.Do(ctx, prompts.GetValidatePrompt(task.RefinedDoc, task.Name, task.Generate))
|
||||
// ✅ 如果输入太大(超过 10k tokens),直接跳过
|
||||
if estimatedTokens > 10000 {
|
||||
log.Printf("⚠️ 输入过大(%d tokens),跳过验证,继续流程", estimatedTokens)
|
||||
task.Valid = task.Generate
|
||||
return &entitys.Usage{StateName: entitys.StatusValid.Desc()}, nil
|
||||
}
|
||||
log.Printf("valid(): 步骤1 - 开始检查代码完整性")
|
||||
|
||||
// ========== 步骤1:检查 ==========
|
||||
req := prompts.GetValidatePrompt(task.RefinedDoc, task.Name, task.Generate)
|
||||
|
||||
validCtx, cancel := context.WithTimeout(ctx, 3*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
resp, useAge, err := task.CallLLM.Do(validCtx, req)
|
||||
if err != nil {
|
||||
log.Printf("valid(): error calling LLM: %v", err)
|
||||
if errors.Is(err, context.DeadlineExceeded) {
|
||||
log.Printf("valid(): 检查超时,跳过验证,继续流程")
|
||||
task.Valid = task.Generate
|
||||
return &entitys.Usage{StateName: entitys.StatusValid.Desc()}, nil
|
||||
}
|
||||
return nil, fmt.Errorf("调用大模型失败: %v", err)
|
||||
}
|
||||
|
||||
|
|
@ -366,18 +388,57 @@ func (s *SDKGeneratorBiz) valid(ctx context.Context, task *entitys.Task) (usage
|
|||
CompletionTokens: useAge.CompletionTokens,
|
||||
TotalTokens: useAge.TotalTokens,
|
||||
}
|
||||
s.logImpl.Add(ctx, &model.AiGenerateLog{
|
||||
TaskID: task.TasKModel.TaskID,
|
||||
Type: entitys.StatusValid.String(),
|
||||
RequestContent: pkg.JsonStringIgonErr(req.Messages),
|
||||
ResponseContent: resp,
|
||||
})
|
||||
cleaned := s.cleanValidationResult(resp)
|
||||
|
||||
// ✅ 清理返回结果
|
||||
cleaned := s.cleanValidationResult(validRes)
|
||||
|
||||
// ✅ 判断是否为 OK
|
||||
// 检查通过
|
||||
if cleaned == "OK" {
|
||||
log.Printf("✅ 验证通过,无需修复")
|
||||
task.Valid = task.Generate
|
||||
return usage, nil
|
||||
}
|
||||
|
||||
// 验证不通过,保存修复后的代码
|
||||
task.Valid = validRes
|
||||
// ========== 步骤2:修复 ==========
|
||||
log.Printf("⚠️ 验证发现问题,步骤2 - 开始修复:\n%s", cleaned)
|
||||
|
||||
// 保存问题列表
|
||||
task.ValidIssues = cleaned
|
||||
|
||||
fixReq := prompts.GetFixByIssuesPrompt(task.RefinedDoc, task.Name, task.Generate, cleaned)
|
||||
|
||||
fixCtx, cancel2 := context.WithTimeout(ctx, 5*time.Minute)
|
||||
defer cancel2()
|
||||
|
||||
fixRes, fixUsage, err := task.CallLLM.Do(fixCtx, fixReq)
|
||||
s.logImpl.Add(ctx, &model.AiGenerateLog{
|
||||
TaskID: task.TasKModel.TaskID,
|
||||
Type: entitys.StatusValid.String(),
|
||||
RequestContent: pkg.JsonStringIgonErr(fixReq),
|
||||
ResponseContent: fixRes,
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, context.DeadlineExceeded) {
|
||||
log.Printf("valid(): 修复超时,使用原代码继续")
|
||||
task.Valid = task.Generate
|
||||
return usage, nil
|
||||
}
|
||||
log.Printf("valid(): 修复失败: %v,使用原代码继续", err)
|
||||
task.Valid = task.Generate
|
||||
return usage, nil
|
||||
}
|
||||
|
||||
usage.PromptTokens += fixUsage.PromptTokens
|
||||
usage.CompletionTokens += fixUsage.CompletionTokens
|
||||
usage.TotalTokens += fixUsage.TotalTokens
|
||||
|
||||
task.Valid = fixRes
|
||||
log.Printf("✅ 修复完成")
|
||||
|
||||
return usage, nil
|
||||
}
|
||||
|
||||
|
|
@ -414,6 +475,12 @@ func (s *SDKGeneratorBiz) fix(ctx context.Context, task *entitys.Task, errMsg st
|
|||
return fmt.Errorf("修复失败,已达到最大尝试次数: %v", fixCount)
|
||||
}
|
||||
fix, useAge, err := task.CallLLM.Do(ctx, prompts.FixPrompt(task.Valid, errMsg, task.Name, task.RefinedDoc))
|
||||
s.logImpl.Add(ctx, &model.AiGenerateLog{
|
||||
TaskID: task.TasKModel.TaskID,
|
||||
Type: entitys.StatusFix.String(),
|
||||
RequestContent: errMsg,
|
||||
ResponseContent: fix,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("调用大模型失败: %v", err)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,27 @@
|
|||
package impl
|
||||
|
||||
import (
|
||||
"sdk-generator/internal/data/model"
|
||||
"sdk-generator/tmpl/dataTemp"
|
||||
"sdk-generator/utils"
|
||||
)
|
||||
|
||||
type AiGenerateLogImpl struct {
|
||||
dataTemp.DataTemp
|
||||
db *utils.Db
|
||||
}
|
||||
|
||||
func NewAiGenerateLogImpl(db *utils.Db) *AiGenerateLogImpl {
|
||||
return &AiGenerateLogImpl{
|
||||
DataTemp: *dataTemp.NewDataTemp(db, new(model.AiGenerateLog)),
|
||||
db: db,
|
||||
}
|
||||
}
|
||||
|
||||
func (m *AiGenerateLogImpl) PrimaryKey() string {
|
||||
return "id"
|
||||
}
|
||||
|
||||
func (m *AiGenerateLogImpl) GetTemp() *dataTemp.DataTemp {
|
||||
return &m.DataTemp
|
||||
}
|
||||
|
|
@ -7,4 +7,5 @@ import (
|
|||
var ProviderImpl = wire.NewSet(
|
||||
NewAiGenerateDocImpl,
|
||||
NewAiGenerateTaskImpl,
|
||||
NewAiGenerateLogImpl,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,27 @@
|
|||
// Code generated by gorm.io/gen. DO NOT EDIT.
|
||||
// Code generated by gorm.io/gen. DO NOT EDIT.
|
||||
// Code generated by gorm.io/gen. DO NOT EDIT.
|
||||
|
||||
package model
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
const TableNameAiGenerateLog = "ai_generate_log"
|
||||
|
||||
// AiGenerateLog mapped from table <ai_generate_log>
|
||||
type AiGenerateLog struct {
|
||||
ID int32 `gorm:"column:id;primaryKey;autoIncrement:true" json:"id"`
|
||||
TaskID string `gorm:"column:task_id;not null" json:"task_id"`
|
||||
Type string `gorm:"column:type" json:"type"`
|
||||
RequestContent string `gorm:"column:request_content" json:"request_content"`
|
||||
ResponseContent string `gorm:"column:response_content" json:"response_content"`
|
||||
ToolSelect string `gorm:"column:tool_select" json:"tool_select"`
|
||||
CreatedAt time.Time `gorm:"column:created_at" json:"created_at"`
|
||||
}
|
||||
|
||||
// TableName AiGenerateLog's table name
|
||||
func (*AiGenerateLog) TableName() string {
|
||||
return TableNameAiGenerateLog
|
||||
}
|
||||
|
|
@ -109,6 +109,7 @@ type Task struct {
|
|||
CallLLM *call.CallLLM
|
||||
OutputDir string
|
||||
Valid string
|
||||
ValidIssues string
|
||||
Package string
|
||||
Name string
|
||||
Files []extractor.File
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ import (
|
|||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/sashabaranov/go-openai"
|
||||
)
|
||||
|
|
@ -68,57 +67,57 @@ type loggingRoundTripper struct {
|
|||
|
||||
func (l *loggingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
// Log the URL details for debugging
|
||||
if req != nil && req.URL != nil {
|
||||
log.Printf("AuthRoundTripper: request URL before enforcement: scheme=%s host=%s path=%s full=%s", req.URL.Scheme, req.URL.Host, req.URL.Path, req.URL.String())
|
||||
}
|
||||
//if req != nil && req.URL != nil {
|
||||
// log.Printf("AuthRoundTripper: request URL before enforcement: scheme=%s host=%s path=%s full=%s", req.URL.Scheme, req.URL.Host, req.URL.Path, req.URL.String())
|
||||
//}
|
||||
|
||||
// If base is set and URL has missing scheme/host, resolve against base
|
||||
if l.base != "" && req != nil && req.URL != nil {
|
||||
if req.URL.Scheme == "" || req.URL.Host == "" {
|
||||
if baseURL, err := url.Parse(l.base); err == nil {
|
||||
newURL := baseURL.ResolveReference(req.URL)
|
||||
log.Printf("AuthRoundTripper: fixing request URL: from=%s to=%s", req.URL.String(), newURL.String())
|
||||
//log.Printf("AuthRoundTripper: fixing request URL: from=%s to=%s", req.URL.String(), newURL.String())
|
||||
req.URL = newURL
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Log incoming Authorization header (for debugging)
|
||||
if v := req.Header.Get("Authorization"); v != "" {
|
||||
parts := strings.SplitN(v, " ", 2)
|
||||
if len(parts) == 2 {
|
||||
tok := parts[1]
|
||||
if len(tok) > 8 {
|
||||
log.Printf("AuthRoundTripper: incoming Authorization token prefix: %s", tok[:8])
|
||||
} else {
|
||||
log.Printf("AuthRoundTripper: incoming Authorization token prefix: %s", tok)
|
||||
}
|
||||
}
|
||||
}
|
||||
//if v := req.Header.Get("Authorization"); v != "" {
|
||||
// parts := strings.SplitN(v, " ", 2)
|
||||
// if len(parts) == 2 {
|
||||
// tok := parts[1]
|
||||
// if len(tok) > 8 {
|
||||
// log.Printf("AuthRoundTripper: incoming Authorization token prefix: %s", tok[:8])
|
||||
// } else {
|
||||
// log.Printf("AuthRoundTripper: incoming Authorization token prefix: %s", tok)
|
||||
// }
|
||||
// }
|
||||
//}
|
||||
|
||||
// If we have an API key stored, ensure the Authorization header is set correctly
|
||||
// This defends against middleware/proxies that may have mutated the header
|
||||
if l.apiKey != "" {
|
||||
expectedAuth := "Bearer " + l.apiKey
|
||||
req.Header.Set("Authorization", expectedAuth)
|
||||
log.Printf("AuthRoundTripper: enforced Authorization token prefix: %s", prefix(l.apiKey))
|
||||
//log.Printf("AuthRoundTripper: enforced Authorization token prefix: %s", prefix(l.apiKey))
|
||||
}
|
||||
|
||||
// Log outgoing Authorization header and URL details
|
||||
if req != nil && req.URL != nil {
|
||||
log.Printf("AuthRoundTripper: request URL after enforcement: scheme=%s host=%s path=%s full=%s", req.URL.Scheme, req.URL.Host, req.URL.Path, req.URL.String())
|
||||
}
|
||||
if v := req.Header.Get("Authorization"); v != "" {
|
||||
parts := strings.SplitN(v, " ", 2)
|
||||
if len(parts) == 2 {
|
||||
tok := parts[1]
|
||||
if len(tok) > 8 {
|
||||
log.Printf("AuthRoundTripper: outgoing Authorization token prefix: %s", tok[:8])
|
||||
} else {
|
||||
log.Printf("AuthRoundTripper: outgoing Authorization token prefix: %s", tok)
|
||||
}
|
||||
}
|
||||
//log.Printf("AuthRoundTripper: request URL after enforcement: scheme=%s host=%s path=%s full=%s", req.URL.Scheme, req.URL.Host, req.URL.Path, req.URL.String())
|
||||
}
|
||||
//if v := req.Header.Get("Authorization"); v != "" {
|
||||
// parts := strings.SplitN(v, " ", 2)
|
||||
// if len(parts) == 2 {
|
||||
// tok := parts[1]
|
||||
// if len(tok) > 8 {
|
||||
// log.Printf("AuthRoundTripper: outgoing Authorization token prefix: %s", tok[:8])
|
||||
// } else {
|
||||
// log.Printf("AuthRoundTripper: outgoing Authorization token prefix: %s", tok)
|
||||
// }
|
||||
// }
|
||||
//}
|
||||
return l.rt.RoundTrip(req)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -82,11 +82,6 @@ func JsonStringIgonErr(data interface{}) string {
|
|||
return string(JsonByteIgonErr(data))
|
||||
}
|
||||
|
||||
func JsonByteIgonErr(data interface{}) []byte {
|
||||
dataByte, _ := json.Marshal(data)
|
||||
return dataByte
|
||||
}
|
||||
|
||||
func IntersectionGeneric[T comparable](slice1, slice2 []T) []T {
|
||||
m := make(map[T]bool)
|
||||
result := []T{}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
// sdk_generate.go - 完整优化版
|
||||
// sdk_generate.go - 完整优化版(兼容 DeepSeek GA 版本)
|
||||
package prompts
|
||||
|
||||
import (
|
||||
|
|
@ -6,7 +6,10 @@ import (
|
|||
"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"
|
||||
|
|
@ -20,36 +23,32 @@ import (
|
|||
type SDKGenerator struct {
|
||||
openaiClient *openai.Client
|
||||
cryptoManager *crypt.CryptoSkillManager
|
||||
maxIterations int // 最大迭代次数,防止死循环
|
||||
timeout time.Duration // 总超时时间
|
||||
model string // OpenAI 模型
|
||||
maxIterations int
|
||||
timeout time.Duration
|
||||
model string
|
||||
}
|
||||
|
||||
// SDKGeneratorOption 配置选项
|
||||
type SDKGeneratorOption func(*SDKGenerator)
|
||||
|
||||
// WithMaxIterations 设置最大迭代次数
|
||||
func WithMaxIterations(n int) SDKGeneratorOption {
|
||||
return func(g *SDKGenerator) {
|
||||
g.maxIterations = n
|
||||
}
|
||||
}
|
||||
|
||||
// WithTimeout 设置超时时间
|
||||
func WithTimeout(t time.Duration) SDKGeneratorOption {
|
||||
return func(g *SDKGenerator) {
|
||||
g.timeout = t
|
||||
}
|
||||
}
|
||||
|
||||
// WithModel 设置模型
|
||||
func WithModel(model string) SDKGeneratorOption {
|
||||
return func(g *SDKGenerator) {
|
||||
g.model = model
|
||||
}
|
||||
}
|
||||
|
||||
// NewSDKGenerator 创建 SDK 生成器
|
||||
func NewSDKGenerator(client *openai.Client, opts ...SDKGeneratorOption) *SDKGenerator {
|
||||
g := &SDKGenerator{
|
||||
openaiClient: client,
|
||||
|
|
@ -58,20 +57,18 @@ func NewSDKGenerator(client *openai.Client, opts ...SDKGeneratorOption) *SDKGene
|
|||
timeout: 10 * time.Minute,
|
||||
model: openai.GPT4,
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
opt(g)
|
||||
}
|
||||
|
||||
return g
|
||||
}
|
||||
|
||||
// GenerateSDK 生成 SDK 代码 - AI 自主决策流程
|
||||
func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, sdkName string, needImplement string) (string, *entitys.Usage, error) {
|
||||
// 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(sdkName, needImplement)
|
||||
systemPrompt := g.buildSystemPrompt(task.Name, needImplement)
|
||||
|
||||
messages := []openai.ChatCompletionMessage{
|
||||
{
|
||||
|
|
@ -80,7 +77,7 @@ func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, sdkName stri
|
|||
},
|
||||
{
|
||||
Role: openai.ChatMessageRoleUser,
|
||||
Content: fmt.Sprintf("请根据以下文档生成完整的 SDK 代码,必须包含所有 6 个文件:\n\n%s", doc),
|
||||
Content: fmt.Sprintf("请根据以下文档生成完整的 SDK 代码,必须包含所有 6 个文件:\n\n**⚠️ 重要规则**:\n1. 调用工具时,**不要输出任何额外文本**(包括句号、逗号、空格等)\n2. 直接调用工具,等待工具返回结果\n3. 工具返回后,**立即开始生成代码**\n4. 不要重复调用已调用过的工具\n\n%s", doc),
|
||||
},
|
||||
}
|
||||
|
||||
|
|
@ -89,7 +86,8 @@ func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, sdkName stri
|
|||
var allResults []string
|
||||
calledTools := make(map[string]bool)
|
||||
useAge := &entitys.Usage{}
|
||||
allFilesGenerated := false
|
||||
forceTextGeneration := false
|
||||
guidedToGenerate := false // 是否已经引导过生成代码
|
||||
|
||||
for {
|
||||
select {
|
||||
|
|
@ -109,9 +107,17 @@ func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, sdkName stri
|
|||
Model: g.model,
|
||||
Messages: messages,
|
||||
Tools: tools,
|
||||
ToolChoice: "auto",
|
||||
ToolChoice: nil,
|
||||
Temperature: 0.1,
|
||||
MaxTokens: 16384, // 增大 token 以生成完整代码
|
||||
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)
|
||||
|
|
@ -122,74 +128,83 @@ func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, sdkName stri
|
|||
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 生成完成")
|
||||
if len(allResults) > 0 {
|
||||
return g.mergeResults(msg.Content, allResults), nil, nil
|
||||
}
|
||||
return msg.Content, useAge, nil
|
||||
return g.mergeResults(msg.Content, allResults), useAge, nil
|
||||
}
|
||||
|
||||
// 检查是否生成了所有文件
|
||||
if g.checkAllFilesGenerated(msg.Content) {
|
||||
allFilesGenerated = true
|
||||
log.Printf("✅ 所有 6 个文件已生成")
|
||||
if len(allResults) > 0 {
|
||||
return g.mergeResults(msg.Content, allResults), nil, nil
|
||||
}
|
||||
return msg.Content, useAge, nil
|
||||
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", toolNames, iteration)
|
||||
} else {
|
||||
log.Printf("📝 AI 回复长度: %d 字符", len(msg.Content))
|
||||
// 如果 AI 没有调用工具也没有完成,提示它继续
|
||||
if len(allResults) > 0 && !allFilesGenerated {
|
||||
messages = append(messages, openai.ChatCompletionMessage{
|
||||
Role: openai.ChatMessageRoleAssistant,
|
||||
Content: fmt.Sprintf("已获取 %d 个加密实现,请现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='", len(allResults)),
|
||||
})
|
||||
continue
|
||||
}
|
||||
}
|
||||
log.Printf("🔧 AI 调用工具: %v, 迭代: %d, 已调用: %v", toolNames, iteration, getCalledToolsList(calledTools))
|
||||
|
||||
messages = append(messages, msg)
|
||||
|
||||
if len(msg.ToolCalls) == 0 {
|
||||
if len(allResults) > 0 {
|
||||
return g.mergeResults(msg.Content, allResults), nil, 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: "所有加密工具已调用完成,请现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='",
|
||||
})
|
||||
continue
|
||||
}
|
||||
|
||||
// 并行处理所有工具调用
|
||||
// 【修复】先执行工具调用,再判断是否完成
|
||||
toolResults, results, err := g.processToolCalls(ctx, msg.ToolCalls)
|
||||
if err != nil {
|
||||
return "", useAge, err
|
||||
|
|
@ -200,17 +215,74 @@ func (g *SDKGenerator) GenerateSDK(ctx context.Context, doc string, sdkName stri
|
|||
useAge.PromptTokens += resp.Usage.PromptTokens
|
||||
useAge.CompletionTokens += resp.Usage.CompletionTokens
|
||||
useAge.TotalTokens += resp.Usage.TotalTokens
|
||||
// 添加一个明确的提示,告诉 AI 继续
|
||||
|
||||
// 【修复】工具执行完成后,再引导生成代码(避免遗漏工具结果)
|
||||
if len(calledTools) >= 2 && !guidedToGenerate {
|
||||
guidedToGenerate = true
|
||||
log.Printf("✅ 所有加密工具已调用完成(共 %d 个),引导 AI 生成代码", len(calledTools))
|
||||
messages = append(messages, openai.ChatCompletionMessage{
|
||||
Role: openai.ChatMessageRoleAssistant,
|
||||
Content: fmt.Sprintf("✅ 已获取 %d 个加密实现,请现在生成完整的 SDK 代码,包含所有 6 个文件。完成后输出 '=== SDK 生成完成 ==='", len(allResults)),
|
||||
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)),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// processToolCalls 并行处理工具调用(带去重和 panic 保护)
|
||||
// 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 {
|
||||
|
|
@ -234,27 +306,23 @@ func (g *SDKGenerator) processToolCalls(ctx context.Context, toolCalls []openai.
|
|||
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,
|
||||
|
|
@ -269,14 +337,12 @@ func (g *SDKGenerator) processToolCalls(ctx context.Context, toolCalls []openai.
|
|||
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 {
|
||||
|
|
@ -299,14 +365,12 @@ func (g *SDKGenerator) checkAllFilesGenerated(content string) bool {
|
|||
"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)
|
||||
}
|
||||
|
||||
|
|
@ -320,23 +384,25 @@ func (g *SDKGenerator) buildSystemPrompt(sdkName, needImplement string) string {
|
|||
return fmt.Sprintf(basePrompt, g.cryptoManager.GetToolDescriptions())
|
||||
}
|
||||
|
||||
// mergeResults 合并加密实现到最终代码
|
||||
// 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()
|
||||
}
|
||||
|
||||
|
|
@ -345,7 +411,6 @@ const SdkGeneratePrompt = `你是一个资深的 Go 语言工程师,擅长生
|
|||
## 任务
|
||||
根据用户提供的 API 接口文档,生成一个**完整的** Go SDK 代码工程。
|
||||
|
||||
|
||||
{{needImplement}}
|
||||
|
||||
## ⚠️ 重要:输出格式要求(必须严格遵守)
|
||||
|
|
@ -369,13 +434,14 @@ package {{sdk_name}}
|
|||
## 重要:你有加密工具可用
|
||||
%s
|
||||
|
||||
## 工具调用策略
|
||||
- 只在文档明确提到特定加密算法时才调用工具
|
||||
- 工具返回的是完整的实现代码,直接集成到 SDK 的 crypto.go 中
|
||||
## 工具调用规则(必须严格遵守)
|
||||
- 调用工具时,**只调用工具,不输出任何文本**(包括句号、逗号、标点符号)
|
||||
- 不要在工具调用前后添加任何注释或说明
|
||||
- 工具调用完成后,**立即**开始生成代码
|
||||
- 每个工具**只调用一次**,不要重复调用
|
||||
|
||||
## 工作流程
|
||||
1. 分析文档中的加密需求 → 调用对应的工具获取实现
|
||||
1. 分析文档中的加密需求 → 调用对应的工具获取实现(只调用需要的工具,通常2-3个)
|
||||
2. 将工具返回的代码放到 crypto.go 中
|
||||
3. **然后立即生成所有其他文件**:client.go、types.go、go.mod、errors.go、example_test.go
|
||||
|
||||
|
|
|
|||
|
|
@ -5,55 +5,139 @@ import (
|
|||
"github.com/sashabaranov/go-openai"
|
||||
)
|
||||
|
||||
// GetValidatePrompt 验证代码完整性,直接返回修复后的代码
|
||||
// GetValidatePrompt 步骤1:只检查,输出缺失列表
|
||||
func GetValidatePrompt(refineDoc, sdkName, codeContent string) openai.ChatCompletionRequest {
|
||||
return openai.ChatCompletionRequest{
|
||||
Messages: []openai.ChatCompletionMessage{
|
||||
{
|
||||
Role: openai.ChatMessageRoleSystem,
|
||||
Content: "你是代码审查专家。检查生成的 SDK 代码是否完整实现了文档中的所有内容。",
|
||||
Content: "你是代码审查专家。检查 SDK 代码是否完整实现了文档中的所有内容。只输出检查结论,不要生成代码。",
|
||||
},
|
||||
{
|
||||
Role: openai.ChatMessageRoleUser,
|
||||
Content: ValidPrompt(refineDoc, sdkName, codeContent),
|
||||
Content: ValidateOnlyPrompt(refineDoc, sdkName, codeContent),
|
||||
},
|
||||
},
|
||||
Temperature: 0.1, // 更低,保证确定性输出
|
||||
TopP: 0.9, // 核采样,平衡多样性和质量
|
||||
MaxTokens: 65536, // 加大到 64k,确保足够
|
||||
FrequencyPenalty: 0.0, // 代码生成不需要
|
||||
PresencePenalty: 0.0, // 代码生成不需要
|
||||
Stop: nil, // 不设置,让模型完整输出
|
||||
//ToolChoice: nil,
|
||||
Temperature: 0.1,
|
||||
TopP: 0.9,
|
||||
MaxTokens: 4096, // 稍微大一点,因为缺失列表可能不止几行
|
||||
FrequencyPenalty: 0.0,
|
||||
PresencePenalty: 0.0,
|
||||
Stop: nil,
|
||||
}
|
||||
}
|
||||
|
||||
func ValidPrompt(codeContent, sdkName, refineDoc string) string {
|
||||
// GetFixByIssuesPrompt 步骤2:根据缺失列表修复代码
|
||||
func GetFixByIssuesPrompt(refineDoc, sdkName, codeContent, issues string) openai.ChatCompletionRequest {
|
||||
return openai.ChatCompletionRequest{
|
||||
Messages: []openai.ChatCompletionMessage{
|
||||
{
|
||||
Role: openai.ChatMessageRoleSystem,
|
||||
Content: "你是代码生成专家。根据问题列表修复代码,输出全量代码。",
|
||||
},
|
||||
{
|
||||
Role: openai.ChatMessageRoleUser,
|
||||
Content: FixByIssuesPrompt(refineDoc, sdkName, codeContent, issues),
|
||||
},
|
||||
},
|
||||
//ToolChoice: nil,
|
||||
Temperature: 0.1,
|
||||
TopP: 0.9,
|
||||
MaxTokens: 65536,
|
||||
FrequencyPenalty: 0.0,
|
||||
PresencePenalty: 0.0,
|
||||
Stop: nil,
|
||||
}
|
||||
}
|
||||
|
||||
return `你是代码审查专家。检查生成的 SDK 代码是否完整实现了文档中的所有内容。
|
||||
// ========== 通用 Prompt ==========
|
||||
|
||||
## 精炼文档
|
||||
func ValidateOnlyPrompt(refineDoc, sdkName, codeContent string) string {
|
||||
return `你是代码审查专家。检查 SDK 代码是否完整实现了文档中的所有内容。
|
||||
|
||||
## 精炼文档(接口定义)
|
||||
` + refineDoc + `
|
||||
|
||||
## 当前代码
|
||||
` + codeContent + `
|
||||
|
||||
## 检查清单
|
||||
## 检查项(根据文档内容自动适配)
|
||||
|
||||
1. 文档中的所有接口是否都已实现?
|
||||
2. 每个接口的请求参数是否完整(字段名、类型、必填)?
|
||||
3. 每个接口的响应字段是否完整?
|
||||
4. 认证方式是否已实现?
|
||||
5. 加密方式是否已实现?
|
||||
6. 签名方式是否已实现?
|
||||
7. 错误码是否已定义?
|
||||
请根据文档中的接口定义,逐项检查以下内容:
|
||||
|
||||
## 输出规则(严格遵守)
|
||||
1. **接口/方法完整性**
|
||||
- 文档中定义的所有 API 接口/方法,代码中是否都有对应的实现?
|
||||
- 方法名、参数、返回值是否与文档一致?
|
||||
|
||||
2. **数据结构完整性**
|
||||
- 文档中定义的所有请求/响应结构体,代码中是否都已定义?
|
||||
- 结构体字段名、类型、必填/可选是否与文档一致?
|
||||
- 枚举值、常量是否正确定义?
|
||||
|
||||
3. **认证与安全**
|
||||
- 文档中的认证方式(签名、加密、token等)是否已实现?
|
||||
- 签名算法、加密算法是否正确?
|
||||
|
||||
4. **错误处理**
|
||||
- 文档中的错误码是否已定义?
|
||||
- 错误类型、错误信息是否完整?
|
||||
|
||||
5. **客户端初始化**
|
||||
- 是否有 NewClient 方法?
|
||||
- 是否支持配置(超时、重试等)?
|
||||
|
||||
6. **其他文档要求**
|
||||
- 文档中提到的其他功能(日志、监控、中间件等)是否已实现?
|
||||
|
||||
## ⚠️ 输出规则(必须严格遵守)
|
||||
|
||||
**情况一:审查通过,没有任何问题**
|
||||
只输出两个字符:` + "`OK`" + `
|
||||
只输出:OK
|
||||
|
||||
**情况二:存在问题**
|
||||
输出修复后的**全量代码**,每个文件用以下格式:
|
||||
输出缺失项列表,按以下格式:
|
||||
|
||||
### 缺失接口
|
||||
- InterfaceName: 文档定义了但代码中没有实现
|
||||
|
||||
### 缺失字段
|
||||
- StructName.FieldName: 文档定义了但结构体中缺少
|
||||
|
||||
### 缺失认证/加密
|
||||
- 具体描述缺失的认证或加密逻辑
|
||||
|
||||
### 其他缺失
|
||||
- 其他文档要求但代码中缺失的内容
|
||||
|
||||
❌ 绝对不要输出修复后的代码!
|
||||
❌ 绝对不要输出完整的文件内容!
|
||||
✅ 只输出 OK 或上述格式的缺失列表!
|
||||
`
|
||||
}
|
||||
|
||||
func FixByIssuesPrompt(refineDoc, sdkName, codeContent, issues string) string {
|
||||
return `根据问题列表修复 SDK 代码。
|
||||
|
||||
## 精炼文档(供参考)
|
||||
` + refineDoc + `
|
||||
|
||||
## 当前代码
|
||||
` + codeContent + `
|
||||
|
||||
## 需要修复的问题
|
||||
` + issues + `
|
||||
|
||||
## 修复要求
|
||||
1. 根据问题列表,逐项修复代码
|
||||
2. 补充缺失的接口、方法、结构体、字段
|
||||
3. 补充缺失的认证、加密、签名逻辑
|
||||
4. 补充缺失的错误码定义
|
||||
5. 保持原有代码风格不变
|
||||
6. 不要删除现有代码(除非是重复定义)
|
||||
|
||||
## 输出规则
|
||||
输出修复后的**完整代码**,每个文件用以下格式:
|
||||
|
||||
// File: ` + sdkName + `/文件名.go
|
||||
` + "```go" + `
|
||||
|
|
@ -61,13 +145,7 @@ package ` + sdkName + `
|
|||
// ... 代码内容 ...
|
||||
` + "```" + `
|
||||
|
||||
|
||||
|
||||
## 重要提醒
|
||||
|
||||
1. 如果存在问题,必须输出**所有文件**的完整代码,不能只输出修改的部分
|
||||
2. 保持原有代码风格不变
|
||||
3. 补全所有遗漏的接口、参数、加密/签名方式
|
||||
4. 不要加任何描述性注释
|
||||
❌ 不要只输出修改的部分!
|
||||
✅ 必须输出所有文件的完整代码!
|
||||
`
|
||||
}
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
Loading…
Reference in New Issue