1
This commit is contained in:
@@ -0,0 +1,124 @@
|
||||
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"
|
||||
)
|
||||
|
||||
type contract struct{}
|
||||
|
||||
var Contract = &contract{}
|
||||
|
||||
func (c *contract) Upload(ctx context.Context, req *dto.UploadContractReq) (*dto.UploadContractRes, error) {
|
||||
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()
|
||||
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
|
||||
}
|
||||
|
||||
func (c *contract) List(ctx context.Context, req *dto.ListContractTaskReq) (*dto.ListContractTaskRes, error) {
|
||||
list, total, err := service.AnnotationService.List(ctx, req.Page, req.PageSize)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.ListContractTaskRes{
|
||||
List: list,
|
||||
Total: total,
|
||||
Page: req.Page,
|
||||
PageSize: req.PageSize,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (c *contract) Detail(ctx context.Context, req *dto.GetContractDetailReq) (*dto.GetContractDetailRes, error) {
|
||||
task, err := dao.ContractTask.GetOne(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)
|
||||
for _, cl := range clauses {
|
||||
list, err := dao.ContractMark.ListByClause(ctx, cl.Id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
marks[cl.Id] = list
|
||||
}
|
||||
return &dto.GetContractDetailRes{Task: task, Clauses: clauses, Marks: marks}, nil
|
||||
}
|
||||
|
||||
func (c *contract) Delete(ctx context.Context, req *dto.DeleteContractReq) (*dto.DeleteContractRes, error) {
|
||||
if err := service.AnnotationService.Delete(ctx, req.Id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &dto.DeleteContractRes{}, nil
|
||||
}
|
||||
Reference in New Issue
Block a user