This commit is contained in:
2026-08-10 10:55:50 +08:00
parent 54de49a6de
commit 61d4f1be63
13 changed files with 781 additions and 1289 deletions
+3 -87
View File
@@ -2,18 +2,8 @@ package controller
import (
"context"
"encoding/json"
"io"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"rag-local/common"
"rag-local/kb/dao"
"rag-local/kb/model/dto"
"rag-local/kb/model/entity"
"rag-local/kb/service"
"github.com/gogf/gf/v2/errors/gerror"
@@ -28,56 +18,10 @@ func (c *contract) Upload(ctx context.Context, req *dto.UploadContractReq) (*dto
if req.File == nil {
return nil, gerror.New("请选择合同文件")
}
var dsIds []int64
if req.DatasetIds != "" {
if err := json.Unmarshal([]byte(req.DatasetIds), &dsIds); err != nil {
return nil, gerror.New("dataset_ids 格式错误,应为 JSON 数组")
}
}
if len(dsIds) == 0 {
return nil, gerror.New("请至少选择一个法律语料数据集")
}
f, err := req.File.Open()
id, err := service.AnnotationService.Upload(ctx, req.File, req.DatasetIds)
if err != nil {
return nil, err
}
defer func() { _ = f.Close() }()
data, err := io.ReadAll(f)
if err != nil {
return nil, err
}
if len(data) == 0 {
return nil, gerror.New("文件内容为空")
}
ext := strings.TrimPrefix(strings.ToLower(filepath.Ext(req.File.Filename)), ".")
supported := false
for _, e := range common.SupportedExts() {
if e == ext {
supported = true
break
}
}
if !supported {
return nil, gerror.New("不支持的文件类型,仅支持 txt/md/pdf/docx/doc/html")
}
relDir := filepath.Join("contract", time.Now().Format("20060102"))
relPath := filepath.Join(relDir, common.RandomToken(16)+"."+ext)
absPath := filepath.Join("workspace", relPath)
if err := os.MkdirAll(filepath.Dir(absPath), 0o755); err != nil {
return nil, err
}
if err := os.WriteFile(absPath, data, 0o644); err != nil {
return nil, err
}
ids := make([]string, 0, len(dsIds))
for _, id := range dsIds {
ids = append(ids, strconv.FormatInt(id, 10))
}
id, err := dao.ContractTask.Insert(ctx, req.File.Filename, relPath, strings.Join(ids, ","))
if err != nil {
_ = os.Remove(absPath)
return nil, err
}
return &dto.UploadContractRes{Id: id}, nil
}
@@ -95,43 +39,15 @@ func (c *contract) List(ctx context.Context, req *dto.ListContractTaskReq) (*dto
}
func (c *contract) Detail(ctx context.Context, req *dto.GetContractDetailReq) (*dto.GetContractDetailRes, error) {
task, err := dao.ContractTask.GetOne(ctx, req.Id)
task, clauses, marks, risks, err := service.AnnotationService.Detail(ctx, req.Id)
if err != nil {
return nil, err
}
if task == nil {
return nil, gerror.New("任务不存在")
}
clauses, err := dao.ContractClause.ListByTask(ctx, req.Id)
if err != nil {
return nil, err
}
marks := make(map[int64][]*entity.ContractMark)
risks := make(map[int64][]*entity.ContractRisk)
for _, cl := range clauses {
list, err := dao.ContractMark.ListByClause(ctx, cl.Id)
if err != nil {
return nil, err
}
marks[cl.Id] = list
riskList, err := dao.ContractRisk.ListByClause(ctx, cl.Id)
if err != nil {
return nil, err
}
risks[cl.Id] = riskList
}
return &dto.GetContractDetailRes{Task: task, Clauses: clauses, Marks: marks, Risks: risks}, nil
}
func (c *contract) Summary(ctx context.Context, req *dto.SummaryContractReq) (*dto.SummaryContractRes, error) {
task, err := dao.ContractTask.GetOne(ctx, req.Id)
if err != nil {
return nil, err
}
if task == nil {
return nil, gerror.New("任务不存在")
}
sum, err := service.AnnotationService.Summary(ctx, req.Id)
task, sum, err := service.AnnotationService.Summary(ctx, req.Id)
if err != nil {
return nil, err
}