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" "github.com/gogf/gf/v2/frame/g" ) 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) 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) if err != nil { return nil, err } return &dto.SummaryContractRes{Task: task, RiskSummary: *sum}, nil } func (c *contract) Annotated(ctx context.Context, req *dto.AnnotatedContractReq) (*dto.AnnotatedContractRes, error) { htmlStr, err := service.AnnotationService.AnnotatedHTML(ctx, req.Id) if err != nil { return nil, err } // 直接写响应体(html),中间件检测到已写入则不包装 JSON r := g.RequestFromCtx(ctx) r.Response.Header().Set("Content-Type", "text/html; charset=utf-8") r.Response.Write(htmlStr) return &dto.AnnotatedContractRes{}, 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 }