sdk_generate/internal/service/generator.go

444 lines
12 KiB
Go
Raw Permalink 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 service
import (
"context"
"fmt"
"log"
"os"
"path/filepath"
"sdk-generator/internal/call"
"sdk-generator/internal/config"
"sdk-generator/internal/extractor"
"sdk-generator/internal/models"
"sdk-generator/internal/msg"
"sdk-generator/internal/postprocess"
"sdk-generator/internal/prompts"
"sdk-generator/internal/push"
"strings"
"sync"
"time"
"github.com/google/uuid"
)
type SDKGeneratorService struct {
config *config.Config
tasks map[string]*models.Task
tasksMu sync.RWMutex
}
func NewSDKGeneratorService(cfg *config.Config) *SDKGeneratorService {
// 创建 OpenAI 客户端配置
return &SDKGeneratorService{
config: cfg,
tasks: make(map[string]*models.Task),
}
}
// Generate 异步生成 SDK
func (s *SDKGeneratorService) Generate(ctx context.Context, content []byte, sdkName string, llmSet *call.LlmCallSet) (string, error) {
// 限制任务数量
s.tasksMu.RLock()
if len(s.tasks) >= s.config.MaxTasks {
s.tasksMu.RUnlock()
return "", fmt.Errorf("任务队列已满,请稍后重试")
}
s.tasksMu.RUnlock()
taskID := uuid.New().String()
task := &models.Task{
ID: taskID,
Status: models.StatusPending,
SdkName: sdkName,
OutputDir: filepath.Join(s.config.OutputDir, taskID),
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
CallLLM: call.NewCallLLM(llmSet),
MarkdownContent: string(content),
}
s.tasksMu.Lock()
s.tasks[taskID] = task
s.tasksMu.Unlock()
// 异步处理
go func() {
defer func() {
if r := recover(); r != nil {
s.failTask(task, fmt.Sprintf("panic: %v", r))
}
}()
s.processTask(context.Background(), task, string(content))
}()
return taskID, nil
}
func (s *SDKGeneratorService) processTask(ctx context.Context, task *models.Task, markdownContent string) {
s.processTaskStatus(task, models.StatusProcessing)
//0.创建目录
s.processTaskStatus(task, models.StatusMkdir)
err := s.mkdir(ctx, task)
if err != nil {
s.failTask(task, fmt.Sprintf("[mkdir]: %v", err))
return
}
// 1. 分析文档
s.processTaskStatus(task, models.StatusAnaMd)
err = s.anaMd(ctx, task)
if err != nil {
s.failTask(task, fmt.Sprintf("[anaMd]: %v", err))
return
}
// 2. 生成代码
s.processTaskStatus(task, models.StatusGenerateCode)
err = s.generateCode(ctx, task)
if err != nil {
s.failTask(task, fmt.Sprintf("[generateCode]: %v", err))
return
}
// 4. 验证是否存在缺失
s.processTaskStatus(task, models.StatusValid)
err = s.valid(ctx, task)
if err != nil {
// 后处理失败不中断流程
s.failTask(task, fmt.Sprintf("[valid]: %v", err))
return
}
// 3. 处理文件
s.processTaskStatus(task, models.StatusCreatFile)
err = s.creatFile(ctx, task)
if err != nil {
// 后处理失败不中断流程
s.failTask(task, fmt.Sprintf("[creatFile]: %v", err))
return
}
// 5. 后处理
s.processTaskStatus(task, models.StatusProcess)
err = postprocess.Process(ctx, task.SdkPackage)
if err != nil {
//如果运行报错则多轮重试
s.processTaskStatus(task, models.StatusFix)
err = s.fix(ctx, task, err.Error(), 1)
}
// 6. 上传仓库
s.processTaskStatus(task, models.StatusSendToGit)
err = s.sendToGit(ctx, task)
if err != nil {
s.failTask(task, fmt.Sprintf("[sendToGit]: %v", err))
}
// 6. 发送钉钉消息
s.processTaskStatus(task, models.StatusMsgSend)
msg.SendDingTalkAlert(ctx, task.SdkName, task.RepoURL)
// 5. 完成
now := time.Now()
task.Status = models.StatusCompleted
task.UpdatedAt = now
task.CompletedAt = &now
s.updateTask(task)
}
func (s *SDKGeneratorService) sendToGit(ctx context.Context, task *models.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.Resp != "" {
// 如果只有 Resp 内容,作为 main.go 推送
files = append(files, push.FileConfig{
Path: "main.go",
Content: task.Resp,
Message: "Initial commit: SDK code",
Branch: "main",
})
}
//3. 如果有 MarkdownContent添加为 README.md
if task.MarkdownContent != "" {
files = append(files, push.FileConfig{
Path: "README.md",
Content: task.MarkdownContent,
Message: "添加 README 文档",
Branch: "main",
})
}
// 4. 如果没有文件,返回错误
if len(files) == 0 {
return fmt.Errorf("没有可推送的文件内容")
}
// 5. 生成仓库名称(使用时间戳确保唯一性)
repoName := fmt.Sprintf("%s-%s", task.SdkName, task.CreatedAt.Format("20060102-150405"))
// 6. 配置标签(可选)
var tagCfg *push.TagConfig
if task.CompletedAt != nil {
tagCfg = &push.TagConfig{
TagName: fmt.Sprintf("v%s", task.CompletedAt.Format("20060102.150405")),
Message: fmt.Sprintf("SDK 生成完成: %s", task.SdkName),
}
}
// 7. 创建仓库并推送代码
repoCfg := push.RepoConfig{
Name: repoName,
Description: task.SdkName,
Private: false, // 公共仓库
AutoInit: false,
Gitignores: "",
License: "",
Readme: "",
}
err = giteaClient.CreateRepoAndPush(repoCfg, files, tagCfg, s.config.OrgName)
if err != nil {
s.failTask(task, fmt.Sprintf("推送代码到 Gitea 失败: %v", err))
return fmt.Errorf("推送代码到 Gitea 失败: %v", err)
}
task.RepoURL = fmt.Sprintf("%s/%s/%s", s.config.GiteaUrl, s.config.OrgName, repoName)
return nil
}
func (s *SDKGeneratorService) mkdir(ctx context.Context, task *models.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.SdkPackage = filepath.Join(task.OutputDir, task.SdkName)
return err
}
func (s *SDKGeneratorService) generateCode(ctx context.Context, task *models.Task) (err error) {
promptSetSdtName := strings.Replace(prompts.GenSDKPrompt, "{{sdk_name}}", task.SdkName, -1)
doc := task.Refine
if len(doc) == 0 {
doc = task.MarkdownContent
}
prompt := fmt.Sprintf(promptSetSdtName, doc)
task.Resp, err = task.CallLLM.Do(ctx, prompts.SDKGeneratorPrompt(prompt))
if err != nil {
s.failTask(task, fmt.Sprintf("调用大模型失败: %v", err))
return
}
return err
}
func (s *SDKGeneratorService) anaMd(ctx context.Context, task *models.Task) (err error) {
task.Refine, err = task.CallLLM.Do(ctx, prompts.BuildRefinePrompt(task.MarkdownContent))
if err != nil {
return fmt.Errorf("调用大模型失败: %v", err)
}
return err
}
func (s *SDKGeneratorService) valid(ctx context.Context, task *models.Task) (err error) {
validRes, err := task.CallLLM.Do(ctx, prompts.GetValidatePrompt(task.Refine, task.SdkName, task.Resp))
if err != nil {
return fmt.Errorf("调用大模型失败: %v", err)
}
// ✅ 清理返回结果
cleaned := s.cleanValidationResult(validRes)
// ✅ 判断是否为 OK
if cleaned == "OK" {
task.Valid = "OK"
return nil
}
// 验证不通过,保存修复后的代码
task.Valid = validRes
return nil
}
// cleanValidationResult 清理验证结果
func (s *SDKGeneratorService) 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 *SDKGeneratorService) fix(ctx context.Context, task *models.Task, errMsg string, fixCount int) (err error) {
log.Printf("代码存在问题,需要修复: %s,当前修复次数:%d", errMsg, fixCount)
if fixCount >= s.config.MaxFixAttempts {
return fmt.Errorf("修复失败,已达到最大尝试次数: %v", fixCount)
}
fix, err := task.CallLLM.Do(ctx, prompts.FixPrompt(task.Resp, errMsg, task.SdkName, task.Refine))
if err != nil {
return fmt.Errorf("调用大模型失败: %v", err)
}
if len(task.Valid) > 0 && task.Valid != "OK" {
task.Valid = fix
} else {
task.Resp = fix
}
if err = s.creatFile(ctx, task); err != nil {
// 后处理失败不中断流程
return fmt.Errorf("创建文件失败: %v", err)
}
if err = postprocess.Process(ctx, task.SdkPackage); err != nil {
err = s.fix(ctx, task, err.Error(), fixCount+1)
}
return err
}
func (s *SDKGeneratorService) creatFile(ctx context.Context, task *models.Task) error {
// 3. 提取代码
extract := task.Resp
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.SdkName, "doc.md"),
Content: task.MarkdownContent,
},
extractor.File{
Path: filepath.Join(task.SdkName, "refine.md"),
Content: task.Refine,
},
extractor.File{
Path: filepath.Join(task.SdkName, "generate.md"),
Content: task.Resp,
},
extractor.File{
Path: filepath.Join(task.SdkName, "valid.md"),
Content: task.Valid,
},
)
// 3. 写入文件
if err = extractor.WriteFiles(files, task.OutputDir); err != nil {
return fmt.Errorf("写入文件失败: %v", err)
}
task.Files = files
return nil
}
func (s *SDKGeneratorService) failTask(task *models.Task, errMsg string) {
task.Status = models.StatusFailed
task.Error = errMsg
task.UpdatedAt = time.Now()
s.updateTask(task)
}
func (s *SDKGeneratorService) processTaskStatus(task *models.Task, process models.TaskStatus) {
task.Status = process
task.UpdatedAt = time.Now()
s.tasksMu.Lock()
defer s.tasksMu.Unlock()
s.tasks[task.ID] = task
}
func (s *SDKGeneratorService) updateTask(task *models.Task) {
s.tasksMu.Lock()
defer s.tasksMu.Unlock()
s.tasks[task.ID] = task
}
// GetTask 获取任务状态
func (s *SDKGeneratorService) GetTask(taskID string) (*models.Task, bool) {
s.tasksMu.RLock()
defer s.tasksMu.RUnlock()
task, ok := s.tasks[taskID]
return task, ok
}
// ListTasks 列出所有任务
func (s *SDKGeneratorService) ListTasks() []*models.Task {
s.tasksMu.RLock()
defer s.tasksMu.RUnlock()
tasks := make([]*models.Task, 0, len(s.tasks))
for _, task := range s.tasks {
tasks = append(tasks, task)
}
return tasks
}
// DeleteTask 删除任务
func (s *SDKGeneratorService) DeleteTask(taskID string) error {
s.tasksMu.Lock()
defer s.tasksMu.Unlock()
task, ok := s.tasks[taskID]
if !ok {
return fmt.Errorf("任务不存在")
}
// 删除输出目录
if err := os.RemoveAll(task.OutputDir); err != nil {
return fmt.Errorf("删除输出目录失败: %w", err)
}
delete(s.tasks, taskID)
return nil
}
// GetTaskOutputDir 获取任务输出目录
func (s *SDKGeneratorService) GetTaskOutputDir(taskID string) string {
return filepath.Join(s.config.OutputDir, taskID)
}