444 lines
12 KiB
Go
444 lines
12 KiB
Go
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)
|
||
}
|