sdk_generate/internal/biz/generator.go

559 lines
16 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
}
func NewSDKGeneratorService(cfg *config.Config, docImpl *impl.AiGenerateDocImpl, taskImpl *impl.AiGenerateTaskImpl) *SDKGeneratorBiz {
// 创建 OpenAI 客户端配置
return &SDKGeneratorBiz{
config: cfg,
docImpl: docImpl,
taskImpl: taskImpl,
}
}
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.Name, NeedImplement)
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) {
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)
validRes, useAge, err := task.CallLLM.Do(ctx, prompts.GetValidatePrompt(task.RefinedDoc, task.Name, task.Generate))
if err != nil {
log.Printf("valid(): error calling LLM: %v", err)
return nil, fmt.Errorf("调用大模型失败: %v", err)
}
usage = &entitys.Usage{
PromptTokens: useAge.PromptTokens,
CompletionTokens: useAge.CompletionTokens,
TotalTokens: useAge.TotalTokens,
}
// ✅ 清理返回结果
cleaned := s.cleanValidationResult(validRes)
// ✅ 判断是否为 OK
if cleaned == "OK" {
task.Valid = task.Generate
return usage, nil
}
// 验证不通过,保存修复后的代码
task.Valid = validRes
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))
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]
}