125 lines
3.2 KiB
Go
125 lines
3.2 KiB
Go
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
|
|
}
|