sdk_generate/internal/biz/generator.go

626 lines
18 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.

package biz
import (
"context"
"encoding/json"
"errors"
"fmt"
"log"
"os"
"path/filepath"
"sdk-generator/internal/config"
"sdk-generator/internal/data/file"
"sdk-generator/internal/data/impl"
"sdk-generator/internal/data/model"
"sdk-generator/internal/entitys"
"sdk-generator/internal/pkg"
"sdk-generator/internal/pkg/call"
"sdk-generator/internal/pkg/extractor"
"sdk-generator/internal/pkg/postprocess"
"sdk-generator/internal/pkg/push"
"sdk-generator/internal/prompts"
"sdk-generator/tmpl/dataTemp"
"sdk-generator/tmpl/errcode"
"strings"
"time"
"github.com/google/uuid"
"xorm.io/builder"
)
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, logImpl *impl.AiGenerateLogImpl) *SDKGeneratorBiz {
// 创建 OpenAI 客户端配置
return &SDKGeneratorBiz{
config: cfg,
docImpl: docImpl,
taskImpl: taskImpl,
logImpl: logImpl,
}
}
func (s *SDKGeneratorBiz) RefineDoc(ctx context.Context, content []byte, callSet *call.LlmCallSet, docDesc string) (*entitys.RefineResponse, error) {
if content == nil || len(content) == 0 {
return nil, errors.New("内容为空")
}
if callSet == nil {
return nil, errors.New("调用参数为空")
}
llmClient := call.NewCallLLM(callSet)
// 1. 分析文档
log.Println("开始分析文档")
res, err := llmClient.DoWithAll(ctx, prompts.BuildRefinePrompt(string(content)))
if err != nil {
return nil, fmt.Errorf("调用大模型失败: %v", err)
}
toolRes, exist := prompts.ExtractDocTypeFromToolCall(res)
if !exist {
return nil, errors.New("提取文档类型失败")
}
uuidStr := uuid.New().String()
s.docImpl.Add(ctx, &model.AiGenerateDoc{
Desc: docDesc,
InstanceID: uuidStr,
Interface: pkg.JsonStringIgonErr(toolRes.Interfaces),
RefinedDoc: toolRes.RefinedDoc,
CreatedAt: time.Now(),
})
return &entitys.RefineResponse{
InstanceId: uuidStr,
RefinedDoc: toolRes.RefinedDoc,
Interfaces: toolRes.Interfaces,
}, err
}
// Generate 异步生成 SDK
func (s *SDKGeneratorBiz) Generate(ctx context.Context, req *entitys.GenerateRequest, llmSet *call.LlmCallSet) (string, error) {
var doc model.AiGenerateDoc
err := s.docImpl.GetByKey(ctx, "instance_id", req.InstanceId, &doc)
if err != nil {
return "", errcode.NotFound("文档不存在")
}
if req.RefinedDoc == "" {
req.RefinedDoc = doc.RefinedDoc
}
llmSet.BaseUrl = "https://ark.cn-beijing.volces.com/api/v3"
// 限制任务数量
taskID := req.InstanceId + time.Now().Format("20060102150405")
task := entitys.Task{
TasKModel: &model.AiGenerateTask{
TaskID: req.InstanceId + "_" + time.Now().Format("20060102150405"),
CodeType: req.CodeType,
Interfaces: req.Interfaces,
GenerateType: req.GenerateType,
CreatedAt: time.Now(),
LlmModel: req.LlmModel,
TaskStatus: entitys.StatusPending.String(),
Desc: req.Desc,
},
RefinedDoc: req.RefinedDoc,
OutputDir: filepath.Join(s.config.OutputDir, taskID),
Name: req.Desc,
CallLLM: call.NewCallLLM(llmSet),
}
if err := s.taskImpl.Add(ctx, task.TasKModel); err != nil {
return "", errors.New("添加任务失败")
}
// 异步处理
//s.processTask(context.Background(), &task)
go func() {
newCtx := context.Background()
defer func() {
if r := recover(); r != nil {
s.failTask(newCtx, &task, fmt.Sprintf("panic: %v", r))
}
}()
s.processTask(newCtx, &task)
}()
return taskID, nil
}
func (s *SDKGeneratorBiz) processTask(ctx context.Context, task *entitys.Task) {
var usAges []*entitys.Usage
defer func() {
var total = &entitys.UsageTask{
Detail: make([]*entitys.Usage, len(usAges)),
}
for k, usage := range usAges {
total.PromptTokens += usage.PromptTokens
total.CompletionTokens += usage.CompletionTokens
total.TotalTokens += usage.TotalTokens
total.Detail[k] = usage
}
task.TasKModel.UseAge = pkg.JsonStringIgonErr(total)
s.updateTask(ctx, task.TasKModel)
}()
s.processTaskStatus(ctx, task, entitys.StatusProcessing)
//0.创建目录
s.processTaskStatus(ctx, task, entitys.StatusMkdir)
err := s.mkdir(ctx, task)
if err != nil {
s.failTask(ctx, task, fmt.Sprintf("[mkdir]: %v", err))
return
}
// 2. 生成代码
s.processTaskStatus(ctx, task, entitys.StatusGenerateCode)
usAge, err := s.generateCode(ctx, task)
if usAge != nil {
usAge.StateName = entitys.StatusGenerateCode.Desc()
usAges = append(usAges, usAge)
}
if err != nil {
s.failTask(ctx, task, fmt.Sprintf("[generateCode]: %v", err))
return
}
// 4. 验证是否存在缺失
s.processTaskStatus(ctx, task, entitys.StatusValid)
usAge, err = s.valid(ctx, task)
if usAge != nil {
usAge.StateName = entitys.StatusValid.Desc()
usAges = append(usAges, usAge)
}
if err != nil {
// 后处理失败不中断流程
s.failTask(ctx, task, fmt.Sprintf("[valid]: %v", err))
return
}
// 3. 处理文件
s.processTaskStatus(ctx, task, entitys.StatusCreatFile)
err = s.creatFile(ctx, task)
if err != nil {
// 后处理失败不中断流程
s.failTask(ctx, task, fmt.Sprintf("[creatFile]: %v", err))
return
}
// 5. 后处理
s.processTaskStatus(ctx, task, entitys.StatusProcess)
err = postprocess.Process(ctx, task.Package)
if err != nil {
//如果运行报错则多轮重试
var failUsage = &entitys.Usage{
StateName: entitys.StatusProcess.Desc(),
}
s.processTaskStatus(ctx, task, entitys.StatusFix)
err = s.fix(ctx, task, err.Error(), 1, failUsage)
if err != nil {
s.failTask(ctx, task, fmt.Sprintf("[fix]: %v", err))
}
usAges = append(usAges, failUsage)
}
// 6. 上传仓库
s.processTaskStatus(ctx, task, entitys.StatusSendToGit)
err = s.sendToGit(ctx, task)
if err != nil {
s.failTask(ctx, task, fmt.Sprintf("[sendToGit]: %v", err))
}
// 6. 发送钉钉消息
//s.processTaskStatus(ctx, task, entitys.StatusMsgSend)
//msg.SendDingTalkAlert(ctx, task.Name, task.TasKModel.RepoURL)
//// 5. 完成
now := time.Now()
task.TasKModel.TaskStatus = entitys.StatusCompleted.String()
task.TasKModel.CompletedAt = &now
s.updateTask(ctx, task.TasKModel)
}
func (s *SDKGeneratorBiz) sendToGit(ctx context.Context, task *entitys.Task) error {
// 1. 获取 Gitea 客户端
giteaClient, err := push.GetClient()
if err != nil || giteaClient == nil {
return fmt.Errorf("[GetClient]: %v", err)
}
// 2. 准备文件列表
var files []push.FileConfig
// 如果有提取的文件列表,推送所有文件
if len(task.Files) > 0 {
for _, file := range task.Files {
files = append(files, push.FileConfig{
Path: file.Path,
Content: file.Content,
Message: fmt.Sprintf("添加文件: %s", file.Path),
Branch: "main",
})
}
} else if task.Generate != "" {
// 如果只有 Resp 内容,作为 main.go 推送
files = append(files, push.FileConfig{
Path: "main.go",
Content: task.Generate,
Message: "Initial commit: SDK code",
Branch: "main",
})
}
//3. 如果有 MarkdownContent添加为 README.md
files = append(files, push.FileConfig{
Path: "README.md",
Content: task.RefinedDoc,
Message: "添加 README 文档",
Branch: "main",
})
// 4. 如果没有文件,返回错误
if len(files) == 0 {
return fmt.Errorf("没有可推送的文件内容")
}
// 5. 生成仓库名称(使用时间戳确保唯一性)
repoName := fmt.Sprintf("%s-%s", task.Name, time.Now().Format("20060102150405"))
// 6. 配置标签(可选)
var tagCfg *push.TagConfig
if task.TasKModel.CompletedAt != nil {
tagCfg = &push.TagConfig{
TagName: fmt.Sprintf("v%s", time.Now().Format("20060102150405")),
Message: fmt.Sprintf("生成完成: %s", task.Name),
}
}
// 7. 创建仓库并推送代码
repoCfg := push.RepoConfig{
Name: repoName,
Description: task.Name,
Private: false, // 公共仓库
AutoInit: false,
Gitignores: "",
License: "",
Readme: "",
}
err = giteaClient.CreateRepoAndPush(repoCfg, files, tagCfg, s.config.OrgName)
if err != nil {
return fmt.Errorf("推送代码到 Gitea 失败: %v", err)
}
task.TasKModel.RepoURL = fmt.Sprintf("%s/%s/%s", s.config.GiteaUrl, s.config.OrgName, repoName)
return nil
}
func (s *SDKGeneratorBiz) mkdir(ctx context.Context, task *entitys.Task) (err error) {
if _, err = os.Stat(task.OutputDir); err == nil {
if err = os.RemoveAll(task.OutputDir); err != nil {
return fmt.Errorf("删除旧目录失败: %v", err)
}
}
if err = os.MkdirAll(task.OutputDir, 0755); err != nil {
return fmt.Errorf("创建输出目录失败: %v", err)
}
task.Package = filepath.Join(task.OutputDir, task.Name)
return err
}
func (s *SDKGeneratorBiz) generateCode(ctx context.Context, task *entitys.Task) (usage *entitys.Usage, err error) {
var NeedImplement string
if len(task.TasKModel.Interfaces) != 0 {
var Interfaces []entitys.Interface
_err := json.Unmarshal([]byte(task.TasKModel.Interfaces), &Interfaces)
if _err == nil {
NeedImplement = "\n ## ⚠️ 接口范围限制(最高优先级,必须严格遵守)\n\n**只实现以下指定的接口,其他接口不要生成:**"
for _, v := range Interfaces {
NeedImplement += fmt.Sprintf("\n\t%s", v.Summary)
}
}
}
var (
finalPrompt strings.Builder
)
finalPrompt.WriteString(task.RefinedDoc)
//finalPrompt.WriteString(NeedImplement)
switch entitys.DocType(task.TasKModel.GenerateType) {
case entitys.DocTypeSdk:
sdkGen := prompts.NewSDKGenerator(
task.CallLLM.Client,
prompts.WithMaxIterations(10),
prompts.WithTimeout(10*time.Minute),
prompts.WithModel(task.TasKModel.LlmModel),
)
task.Generate, usage, err = sdkGen.GenerateSDK(ctx, finalPrompt.String(), task, NeedImplement, s.logImpl)
case entitys.DocTypeServerBoilerplate:
serverGen := prompts.NewServerGenerator(
task.CallLLM.Client,
prompts.WithServerMaxIterations(10),
prompts.WithServerTimeout(10*time.Minute),
prompts.WithServerModel(task.TasKModel.LlmModel),
)
task.Generate, usage, err = serverGen.GenerateServer(ctx, finalPrompt.String(), task.Name, NeedImplement)
}
return
}
func (s *SDKGeneratorBiz) valid(ctx context.Context, task *entitys.Task) (usage *entitys.Usage, err error) {
// ========== 判断输入大小 ==========
// 估算 token 数量1 token ≈ 4 字符(中英文混合)
inputSize := len(task.RefinedDoc) + len(task.Generate)
estimatedTokens := inputSize / 4
log.Printf("valid(): 输入大小: %d 字符, 估算 token: %d", inputSize, estimatedTokens)
// ✅ 如果输入太大(超过 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 {
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)
}
usage = &entitys.Usage{
PromptTokens: useAge.PromptTokens,
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)
// 检查通过
if cleaned == "OK" {
log.Printf("✅ 验证通过,无需修复")
task.Valid = task.Generate
return usage, nil
}
// ========== 步骤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
}
// cleanValidationResult 清理验证结果
func (s *SDKGeneratorBiz) cleanValidationResult(result string) string {
// 1. 去除首尾空白
cleaned := strings.TrimSpace(result)
// 2. 去除 Markdown 代码块标记
cleaned = strings.TrimPrefix(cleaned, "```")
cleaned = strings.TrimSuffix(cleaned, "```")
cleaned = strings.TrimSpace(cleaned)
// 3. 去除可能的引号
cleaned = strings.Trim(cleaned, "\"")
cleaned = strings.Trim(cleaned, "'")
// 4. 只取前两行(防止 OK 后面跟了其他内容)
lines := strings.Split(cleaned, "\n")
if len(lines) > 0 {
cleaned = strings.TrimSpace(lines[0])
}
// 5. 去除可能的尾随标点
cleaned = strings.TrimSuffix(cleaned, ".")
cleaned = strings.TrimSuffix(cleaned, "")
cleaned = strings.TrimSuffix(cleaned, "!")
return cleaned
}
func (s *SDKGeneratorBiz) fix(ctx context.Context, task *entitys.Task, errMsg string, fixCount int, usage *entitys.Usage) (err error) {
log.Printf("代码存在问题,需要修复: %s,当前修复次数:%d", errMsg, fixCount)
if fixCount >= s.config.MaxFixAttempts {
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)
}
usage.PromptTokens += useAge.PromptTokens
usage.CompletionTokens += useAge.CompletionTokens
usage.TotalTokens += useAge.TotalTokens
task.Valid = fix
if err = s.creatFile(ctx, task); err != nil {
// 后处理失败不中断流程
return fmt.Errorf("创建文件失败: %v", err)
}
if err = postprocess.Process(ctx, task.Package); err != nil {
err = s.fix(ctx, task, err.Error(), fixCount+1, usage)
}
return err
}
func (s *SDKGeneratorBiz) creatFile(ctx context.Context, task *entitys.Task) error {
// 3. 提取代码
extract := task.Generate
if task.Valid != "" {
extract = task.Valid
}
files, err := extractor.Extract(extract)
if err != nil {
return fmt.Errorf("提取代码失败: %v", err)
}
if len(files) == 0 {
return fmt.Errorf("未提取到任何代码文件")
}
files = append(files,
extractor.File{
Path: filepath.Join(task.Name, "refine.md"),
Content: task.RefinedDoc,
},
extractor.File{
Path: filepath.Join(task.Name, "generate.md"),
Content: task.Generate,
},
extractor.File{
Path: filepath.Join(task.Name, "valid.md"),
Content: task.Valid,
},
)
switch entitys.DocType(task.TasKModel.GenerateType) {
case entitys.DocTypeSdk:
case entitys.DocTypeServerBoilerplate:
//加上dockerfile
files = append(files, extractor.File{
Path: filepath.Join(task.Name, "Dockerfile"),
Content: file.GetDockerfile(),
})
}
// 3. 写入文件
if err = extractor.WriteFiles(files, task.OutputDir); err != nil {
return fmt.Errorf("写入文件失败: %v", err)
}
task.Files = files
return nil
}
func (s *SDKGeneratorBiz) failTask(ctx context.Context, task *entitys.Task, errMsg string) {
task.TasKModel.TaskStatus = entitys.StatusFailed.String()
task.TasKModel.Error = errMsg
s.updateTask(ctx, task.TasKModel)
}
func (s *SDKGeneratorBiz) processTaskStatus(ctx context.Context, task *entitys.Task, process entitys.TaskStatus) {
s.taskImpl.UpdateByKey(ctx, s.taskImpl.PrimaryKey(), task.TasKModel.ID, &model.AiGenerateTask{TaskStatus: process.String()})
}
func (s *SDKGeneratorBiz) updateTask(ctx context.Context, task *model.AiGenerateTask) error {
return s.taskImpl.UpdateByKey(ctx, s.taskImpl.PrimaryKey(), task.ID, task)
}
// GetTask 获取任务状态
func (s *SDKGeneratorBiz) GetTask(ctx context.Context, taskID string) (*model.AiGenerateTask, bool) {
var task model.AiGenerateTask
err := s.taskImpl.GetByKey(ctx, "task_id", taskID, &task)
if err != nil || task.ID == 0 {
return nil, false
}
return &task, true
}
// ListTasks 列出所有任务
func (s *SDKGeneratorBiz) ListTasks(ctx context.Context, cond *builder.Cond, page, pageSize int) ([]model.AiGenerateTask, int64, error) {
var tasks []model.AiGenerateTask
total, err := s.taskImpl.GetListToStruct(ctx, cond, &dataTemp.ReqPageBo{
Page: page,
Limit: pageSize,
}, &tasks, "updated_at DESC")
if err != nil {
return nil, 0, err
}
return tasks, total.Total, err
}
// DeleteTask 删除任务
func (s *SDKGeneratorBiz) DeleteTask(ctx context.Context, taskID string) error {
err := s.taskImpl.DeleteByKey(ctx, "task_id", taskID)
return err
}
// GetTaskOutputDir 获取任务输出目录
func (s *SDKGeneratorBiz) GetTaskOutputDir(taskID string) string {
return filepath.Join(s.config.OutputDir, taskID)
}
// ListDoc 列出所有任务
func (s *SDKGeneratorBiz) ListDoc(ctx context.Context, cond *builder.Cond, page, pageSize int) ([]model.AiGenerateDoc, int64, error) {
var tasks []model.AiGenerateDoc
total, err := s.docImpl.GetListToStruct(ctx, cond, &dataTemp.ReqPageBo{
Page: page,
Limit: pageSize,
}, &tasks, "updated_at DESC")
if err != nil {
return nil, 0, err
}
return tasks, total.Total, err
}
// ListDoc 列出所有任务
func (s *SDKGeneratorBiz) UpdateDocByInstanceId(ctx context.Context, docModel *model.AiGenerateDoc) error {
return s.docImpl.UpdateByKey(ctx, "instance_id", docModel.InstanceID, docModel)
}
// prefix returns a short prefix of the key for safe logging
func prefix(s string) string {
if s == "" {
return ""
}
if len(s) <= 8 {
return s
}
return s[:8]
}